AQIT 0.1.0
Loading...
Searching...
No Matches
user_sae.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""Local user-trained SAE registry and active binding for inspect / steer."""
3
4from __future__ import annotations
5
6import json
7from pathlib import Path
8from typing import Any
9
10USER_SAE_ROOT = Path.home() / ".aquin" / "sae" / "user"
11ACTIVE_SAE_PATH = Path.home() / ".aquin" / "active_user_sae.json"
12
13
14def _read_meta(sae_path: Path) -> dict[str, Any]:
15 meta_path = sae_path.with_suffix(".meta.json")
16 if not meta_path.exists():
17 return {}
18 try:
19 return json.loads(meta_path.read_text(encoding="utf-8"))
20 except Exception:
21 return {}
22
23
24def _layer_from_filename(path: Path) -> int | None:
25 stem = path.stem
26 if stem.startswith("sae_layer"):
27 try:
28 return int(stem.replace("sae_layer", ""))
29 except ValueError:
30 return None
31 return None
32
33
34def user_sae_dirs_for_model(model_id: str, *, embedding: bool = False) -> list[Path]:
35 from aquin.compute.model_loader import resolve_model_id
36
37 del embedding # LLM-only registry
38 slug = resolve_model_id(model_id)
39 roots = [USER_SAE_ROOT / slug]
40 return [r for r in roots if r.is_dir()]
41
42
44 model_id: str,
45 name: str,
46 layer: int | None = None,
47 *,
48 embedding: bool | None = None,
49) -> Path:
50 """Resolve ~/.aquin/sae/user/<model>/<name>/sae_layer{L}.pt"""
51 from aquin.compute.model_loader import resolve_model_id
52
53 del embedding
54 safe = name.replace("/", "--").replace(" ", "_")
55 roots = user_sae_dirs_for_model(model_id)
56
57 candidates: list[Path] = []
58 for root in roots:
59 run_dir = root / safe
60 if not run_dir.is_dir():
61 continue
62 if layer is not None:
63 candidates.append(run_dir / f"sae_layer{layer}.pt")
64 else:
65 found = sorted(run_dir.glob("sae_layer*.pt"))
66 if len(found) > 1:
67 layers = ", ".join(str(_layer_from_filename(p)) for p in found)
68 raise ValueError(
69 f"User SAE '{name}' has multiple layers ({layers}). Pass --layer <n>."
70 )
71 candidates.extend(found)
72
73 for path in candidates:
74 if path.is_file():
75 return path
76
77 slug = resolve_model_id(model_id)
78 expected = USER_SAE_ROOT / slug / safe
79 if layer is not None:
80 expected_file = expected / f"sae_layer{layer}.pt"
81 else:
82 expected_file = expected / "sae_layer<N>.pt"
83
84 hints: list[str] = [
85 f"Session model resolves to: {slug}",
86 f"Expected path: {expected_file}",
87 ]
88 all_runs = list_user_saes(slug)
89 if all_runs:
90 hints.append("On disk:")
91 for row in all_runs[:8]:
92 hints.append(f" --user {row['name']} --layer {row.get('layer', '?')} ({row['model_id']})")
93 else:
94 hints.append(f"No user SAEs under {USER_SAE_ROOT}/")
95 hints.append("Train: aquin sae train --layer <n> --name my-run --quick")
96
97 layer_hint = f" --layer {layer}" if layer is not None else ""
98 raise FileNotFoundError(
99 f"No user SAE '{name}' for this model{layer_hint}.\n" + "\n".join(hints)
100 )
101
102
103def resolve_path_sae(path: str | Path, *, model_id: str | None = None, layer: int | None = None) -> tuple[Path, str, int]:
104 """Resolve explicit .pt path; infer model + layer from meta or args."""
105 sae_path = Path(path).expanduser().resolve()
106 if not sae_path.is_file():
107 raise FileNotFoundError(f"SAE file not found: {sae_path}")
108
109 from aquin.compute.torch_io import is_valid_sae_checkpoint_path, load_checkpoint
110
111 if not is_valid_sae_checkpoint_path(sae_path):
112 detail: str | None = None
113 try:
114 load_checkpoint(sae_path, map_location="cpu")
115 detail = "the file loaded, but it does not contain SAE weight tensors"
116 except Exception as exc:
117 detail = str(exc).strip() or exc.__class__.__name__
118 raise ValueError(
119 "Invalid SAE checkpoint for --path.\n"
120 f" file: {sae_path}\n"
121 " expected: a non-empty PyTorch/safetensors SAE checkpoint containing encoder/decoder weights\n"
122 f" detail: {detail}\n"
123 "Use a real SAE weights file, or re-download/re-export the checkpoint and try again."
124 )
125
126 meta = _read_meta(sae_path)
127 resolved_layer = layer
128 if resolved_layer is None:
129 resolved_layer = meta.get("layer")
130 if resolved_layer is None:
131 resolved_layer = _layer_from_filename(sae_path)
132 if resolved_layer is None:
133 raise ValueError(f"Could not infer layer for {sae_path}. Pass --layer <n>.")
134
135 resolved_model = model_id or meta.get("model_id")
136 if not resolved_model:
137 raise ValueError(f"Could not infer model for {sae_path}. Pass --model <id> or use a .meta.json.")
138
139 from aquin.compute.model_loader import resolve_model_id
140
141 resolved_model = resolve_model_id(resolved_model)
142 return sae_path, resolved_model, int(resolved_layer)
143
144
145def list_user_saes(model_id: str | None = None) -> list[dict[str, Any]]:
146 """Scan ~/.aquin/sae/user for trained SAE checkpoints."""
147 rows: list[dict[str, Any]] = []
148 if not USER_SAE_ROOT.is_dir():
149 return rows
150
151 for model_dir in sorted(USER_SAE_ROOT.iterdir()):
152 if not model_dir.is_dir():
153 continue
154 dir_name = model_dir.name
155 if dir_name.startswith("embed-"):
156 continue
157 slug = dir_name
158 if model_id:
159 try:
160 from aquin.compute.model_loader import resolve_model_id
161
162 want = resolve_model_id(model_id)
163 if slug != want and dir_name != model_id:
164 continue
165 except ValueError:
166 if slug != model_id and dir_name != model_id:
167 continue
168
169 for run_dir in sorted(model_dir.iterdir()):
170 if not run_dir.is_dir():
171 continue
172 for sae_file in sorted(run_dir.glob("sae_layer*.pt")):
173 layer = _layer_from_filename(sae_file)
174 meta = _read_meta(sae_file)
175 rows.append(
176 {
177 "name": run_dir.name,
178 "model_id": meta.get("model_id") or slug,
179 "layer": layer if layer is not None else meta.get("layer"),
180 "path": str(sae_file),
181 "embedding": False,
182 "steps": meta.get("steps"),
183 }
184 )
185 return rows
186
187
188def load_active_binding() -> dict[str, Any] | None:
189 if not ACTIVE_SAE_PATH.exists():
190 return None
191 try:
192 data = json.loads(ACTIVE_SAE_PATH.read_text(encoding="utf-8"))
193 return data if isinstance(data, dict) else None
194 except Exception:
195 return None
196
197
199 *,
200 model_id: str,
201 layer: int,
202 path: str | Path,
203 name: str | None = None,
204 embedding: bool = False,
205) -> dict[str, Any]:
206 binding = {
207 "model_id": model_id,
208 "layer": int(layer),
209 "path": str(Path(path).resolve()),
210 "name": name,
211 "embedding": False,
212 }
213 del embedding
214 ACTIVE_SAE_PATH.parent.mkdir(parents=True, exist_ok=True)
215 ACTIVE_SAE_PATH.write_text(json.dumps(binding, indent=2), encoding="utf-8")
216 return binding
217
218
219def clear_active_binding() -> None:
220 if ACTIVE_SAE_PATH.exists():
221 ACTIVE_SAE_PATH.unlink(missing_ok=True)
222
224def get_active_user_sae_path(model_id: str, layer: int) -> Path | None:
225 """Return active user SAE path when it matches model + layer."""
226 binding = load_active_binding()
227 if not binding:
228 return None
229 try:
230 bound_layer = int(binding["layer"])
231 except (KeyError, TypeError, ValueError):
232 return None
233 if bound_layer != int(layer):
234 return None
235
236 bound_model = binding.get("model_id")
237 if not bound_model:
238 return None
239
240 from aquin.compute.model_loader import resolve_model_id
241
242 try:
243 if resolve_model_id(model_id) != resolve_model_id(str(bound_model)):
244 return None
245 except ValueError:
246 if str(bound_model).lower() != str(model_id).lower():
247 return None
248
249 path = Path(str(binding.get("path", "")))
250 return path if path.is_file() else None
251
252
254 *,
255 model_id: str,
256 layer: int,
257 path: Path,
258 name: str | None = None,
259) -> dict[str, Any]:
260 """Register active user SAE and warm the in-process cache."""
261 from aquin.compute.model_loader import resolve_model_id
262 from aquin.compute import feature_analysis
263
264 short = resolve_model_id(model_id)
265
266 binding = save_active_binding(
267 model_id=short,
268 layer=layer,
269 path=path,
270 name=name,
271 embedding=False,
272 )
273
274 feature_analysis._sae_cache.pop((short, layer), None)
275 feature_analysis.load_sae(short, layer)
276
277 return binding
list[Path] user_sae_dirs_for_model(str model_id, *, bool embedding=False)
Definition user_sae.py:38
Path|None get_active_user_sae_path(str model_id, int layer)
Definition user_sae.py:228
Path resolve_user_sae_path(str model_id, str name, int|None layer=None, *, bool|None embedding=None)
Definition user_sae.py:53
dict[str, Any]|None load_active_binding()
Definition user_sae.py:192
dict[str, Any] activate_user_sae(*, str model_id, int layer, Path path, str|None name=None)
Definition user_sae.py:263
list[dict[str, Any]] list_user_saes(str|None model_id=None)
Definition user_sae.py:149
int|None _layer_from_filename(Path path)
Definition user_sae.py:28
tuple[Path, str, int] resolve_path_sae(str|Path path, *, str|None model_id=None, int|None layer=None)
Definition user_sae.py:107
dict[str, Any] _read_meta(Path sae_path)
Definition user_sae.py:18
dict[str, Any] save_active_binding(*, str model_id, int layer, str|Path path, str|None name=None, bool embedding=False)
Definition user_sae.py:209