AQIT 0.1.0
Loading...
Searching...
No Matches
sae.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6"""SAE — load / train / align / stats / features."""
7
8from __future__ import annotations
9
10from typing import Any
11
12from aquin.sdk._runtime import invoke
13
14
15def info(sae_id: str, **kwargs: Any) -> dict[str, Any]:
16 from aquin.cli import cmd_info
17
18 args = [sae_id]
19 if kwargs.get("json"):
20 args.append("--json")
21 cmd_info(["sae", *args])
22 return {"ok": True, "sae_id": sae_id}
23
24
25def load(sae_id: str | None = None, **kwargs: Any) -> dict[str, Any]:
26 from aquin.cli import cmd_load
27
28 argv: list[str] = ["sae"]
29 if sae_id:
30 argv.append(sae_id)
31 if kwargs.get("user"):
32 argv.extend(["--user", str(kwargs["user"])])
33 if kwargs.get("path"):
34 argv.extend(["--path", str(kwargs["path"])])
35 if kwargs.get("layer") is not None:
36 argv.extend(["--layer", str(kwargs["layer"])])
37 cmd_load(argv)
38 return {"ok": True}
39
40
41def train(**kwargs: Any) -> dict[str, Any]:
42 from aquin.sae_cli import cmd_sae_train
43
44 argv: list[str] = []
45 for key, flag in (
46 ("model", "--model"),
47 ("layer", "--layer"),
48 ("layers", "--layers"),
49 ("name", "--name"),
50 ("activations", "--activations"),
51 ("checkpoint", "--checkpoint"),
52 ):
53 if kwargs.get(key) is not None:
54 argv.extend([flag, str(kwargs[key])])
55 if kwargs.get("quick"):
56 argv.append("--quick")
57 if kwargs.get("balance"):
58 argv.append("--balance")
59 cmd_sae_train(argv)
60 return {"ok": True}
61
62
63def align(**kwargs: Any) -> dict[str, Any]:
64 from aquin.sae_cli import cmd_sae_align
65
66 argv: list[str] = []
67 for key, flag in (("sae_a", "--sae-a"), ("sae_b", "--sae-b"), ("sae-a", "--sae-a"), ("sae-b", "--sae-b")):
68 if kwargs.get(key) is not None:
69 argv.extend([flag, str(kwargs[key])])
70 cmd_sae_align(argv)
71 return {"ok": True}
72
73
74def stats(**kwargs: Any) -> dict[str, Any]:
75 return invoke("run_sae_stats", kwargs, command="sae-stats")
76
77
78def feature_locate(**kwargs: Any) -> dict[str, Any]:
79 return invoke("run_find_feature", kwargs, command="feature locate")
80
81
82def feature_logit(**kwargs: Any) -> dict[str, Any]:
83 return invoke("get_feature_logits", kwargs, command="feature logit")
84
85
86def feature_neighbor(**kwargs: Any) -> dict[str, Any]:
87 return invoke("get_feature_neighbors", kwargs, command="feature neighbor")
dict[str, Any] align(**Any kwargs)
Definition sae.py:67
dict[str, Any] feature_logit(**Any kwargs)
Definition sae.py:86
dict[str, Any] feature_locate(**Any kwargs)
Definition sae.py:82
dict[str, Any] stats(**Any kwargs)
Definition sae.py:78
dict[str, Any] load(str|None sae_id=None, **Any kwargs)
Definition sae.py:29
dict[str, Any] info(str sae_id, **Any kwargs)
Definition sae.py:19
dict[str, Any] feature_neighbor(**Any kwargs)
Definition sae.py:90