AQIT 0.1.0
Loading...
Searching...
No Matches
sae_cli.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""aquin diff sae — base vs checkpoint SAE activation diff (via aquin diff sae)."""
3
4import json
5import sys
6from pathlib import Path
7
8from aquin.cli_flags import reject_legacy_output_flags
9
10
11def _parse_flag(args: list[str], name: str) -> str | None:
12 for i, a in enumerate(args):
13 if a == name and i + 1 < len(args):
14 return args[i + 1]
15 return None
16
17
18def _has_flag(args: list[str], name: str) -> bool:
19 return name in args
20
21
22def _resolve_sae_model(model_id: str) -> str:
23 """Return canonical model id."""
24 from aquin.compute.model_loader import resolve_model_id
25
26 return resolve_model_id(model_id)
27
28
29def _ensure_compute_env() -> None:
30 """Match cmd_tool: shared model + local HTTP server for web popovers."""
31 from aquin.compute.loader_shim import apply as _shim_apply
32 from aquin.engine.local_server import start as _start_local_server
34 _shim_apply()
35 _start_local_server()
36
37
38def _build_sae_card(tool_name: str, payload: dict, sync_args: dict | None = None) -> dict | None:
39 if tool_name == "run_sae_diff":
40 return {"type": "saeDiff", "data": payload}
41 if tool_name == "run_sae_train":
42 args = sync_args or {}
43 return {
44 "type": "saeTrain",
45 "data": {
46 "modelId": args.get("model_id") or payload.get("model_id"),
47 "layer": args.get("layer") if args.get("layer") is not None else payload.get("layer"),
48 "quick": bool(args.get("quick", payload.get("quick"))),
49 "name": args.get("name") or payload.get("name"),
50 "outputPath": payload.get("output_path") or args.get("save"),
51 "checkpoint": args.get("checkpoint") or payload.get("checkpoint"),
52 "status": payload.get("status", "done"),
53 },
54 }
55 if tool_name == "run_sae_align":
56 pairs = payload.get("pairs") or []
57 sorted_pairs = sorted(pairs, key=lambda x: float(x.get("cosine", 0)))
58 args = sync_args or {}
59 return {
60 "type": "saeAlign",
61 "data": {
62 "meanCosine": payload.get("mean_cosine"),
63 "nPairs": payload.get("n_pairs"),
64 "saeA": args.get("sae_a") or payload.get("sae_a"),
65 "saeB": args.get("sae_b") or payload.get("sae_b"),
66 "lowestPairs": sorted_pairs[:8],
67 "highestPairs": list(reversed(sorted_pairs[-8:])),
68 },
69 }
70 return None
71
72
74 ctx: dict,
75 tool_name: str,
76 args: dict,
77 payload: dict,
78) -> None:
79 from aquin.engine.sync_dispatch import sync_cli_result
80
81 card = _build_sae_card(tool_name, payload, args)
82 sync_cli_result(ctx, tool_name, args, payload, card=card)
83
84
85def _print_sae_diff(payload: dict) -> None:
86 print(f"model : {payload.get('baseModelId')}")
87 print(f"checkpoint : {payload.get('ftCheckpointName')}")
88 print(f"layer : {payload.get('layer')}")
89 print(f"changed : {payload.get('nChanged')}/{payload.get('nFeatures')}")
90 print(f"mean |delta| : {payload.get('meanAbsDelta')}")
91 print(f"max |delta| : {payload.get('maxAbsDelta')}")
92 print("\ntop feature deltas:")
93 for row in (payload.get("featureDeltas") or [])[:15]:
94 print(
95 f" f{row['feature_idx']:>5} base={row['base_act']:.4f} "
96 f"ft={row['ft_act']:.4f} Δ={row['delta']:+.4f}"
97 )
98
99
100def _require_session_model(model_id: str | None) -> str:
101 from aquin.compute.model_loader import get_active_model_id, resolve_model_id
102
103 locked = (get_active_model_id() or "").strip()
104 if not locked:
105 print("Error: no model loaded.")
106 print(" Run: aquin load --model <id>")
107 sys.exit(1)
108 if model_id and model_id != locked:
109 try:
110 if resolve_model_id(model_id) != resolve_model_id(locked):
111 print(f"Error: loaded model is '{locked}'.")
112 print(" --model must match the loaded model or be omitted.")
113 sys.exit(1)
114 except ValueError:
115 pass
116 try:
117 return resolve_model_id(locked)
118 except ValueError:
119 return locked
120
121
122def _resolve_train_model_id(model_id: str | None, activations: str | None) -> str:
123 """Session model, explicit --model, or capture-dir manifest (activations-only)."""
124 from aquin.compute.activation_store import read_manifest
125 from aquin.compute.model_loader import get_active_model_id, resolve_model_id
127 locked = (get_active_model_id() or "").strip()
128 if locked:
129 return _require_session_model(model_id)
130
131 if model_id:
132 try:
133 return resolve_model_id(model_id)
134 except ValueError:
135 raise
136
137 if activations:
138 root = Path(activations).expanduser()
139 manifest = read_manifest(root)
140 if not manifest:
141 # capture dirs nest training acts under _acts_layerN after resolve
142 for child in root.glob("**/manifest.json"):
143 manifest = read_manifest(child.parent)
144 if manifest and manifest.get("model_id"):
145 break
146 mid = (manifest or {}).get("model_id") if manifest else None
147 if isinstance(mid, str) and mid.strip():
148 try:
149 return resolve_model_id(mid.strip())
150 except ValueError:
151 pass
152
153 print("Error: no model loaded.")
154 print(" Run: aquin load --model <id>")
155 print(" Or pass --model <id> / train from --activations with a capture manifest.")
156 sys.exit(1)
157
158
159def cmd_sae_diff(args: list[str]) -> None:
160 reject_legacy_output_flags(args)
161 model_id = _parse_flag(args, "--model")
162 checkpoint = _parse_flag(args, "--checkpoint")
163 prompts_path = _parse_flag(args, "--prompts")
164 layer_s = _parse_flag(args, "--layer")
165 sae_path = _parse_flag(args, "--sae")
166 name = _parse_flag(args, "--name") or Path(checkpoint or "checkpoint").stem
167 save_path = _parse_flag(args, "--save")
168
169 if not checkpoint:
170 print("Usage: aquin diff sae [--model <id>] --checkpoint <path> [--prompts <json|jsonl>]")
171 print(" [--layer N] [--sae <path>] [--name <label>] [--save <json>]")
172 sys.exit(1)
173
174 model_id = _require_session_model(model_id)
175
177
178 from aquin.cli import _build_tool_ctx
179 from aquin.compute.sae_diff import load_prompts, run_sae_diff
180
181 try:
182 mid = _resolve_sae_model(model_id)
183 except ValueError as e:
184 print(f"Error: {e}")
185 sys.exit(1)
186
187 ckpt = Path(checkpoint)
188 if not ckpt.exists():
189 print(f"Checkpoint not found: {ckpt}")
190 sys.exit(1)
191
192 layer = int(layer_s) if layer_s else None
193 prompts = load_prompts(prompts_path)
194 ctx = _build_tool_ctx(model_id=mid)
195 from aquin.engine.sync_dispatch import require_active_session
196
197 require_active_session(ctx, label="aquin diff sae")
198 sync_args = {
199 "model_id": mid,
200 "checkpoint": str(ckpt),
201 "name": name,
202 "prompts": prompts_path or "",
203 "model_mode": "llm",
204 }
205
206 try:
207 print(f"[diff sae] llm base={mid} target={ckpt.name} prompts={len(prompts)}")
208 payload = run_sae_diff(
209 mid,
210 prompts,
211 target_checkpoint=str(ckpt),
212 checkpoint_name=name,
213 layer=layer,
214 sae_path=sae_path,
215 )
216 except Exception as e:
217 print(f"Error: {e}")
218 sys.exit(1)
219
220 _print_sae_diff(payload)
221 if save_path:
222 Path(save_path).write_text(json.dumps(payload, indent=2), encoding="utf-8")
223 print(f"\n[diff sae] wrote {save_path}")
224 _sync_sae_result(ctx, "run_sae_diff", sync_args, payload)
225
226
227def _parse_train_layers(spec: str, n_layers: int) -> list[int]:
228 raw = (spec or "").strip().lower()
229 if raw in ("all", "*"):
230 return list(range(n_layers))
231 out: list[int] = []
232 for part in raw.split(","):
233 part = part.strip()
234 if not part:
235 continue
236 layer = int(part)
237 if layer < 0 or layer >= n_layers:
238 raise ValueError(f"Layer {layer} out of range for model (0..{n_layers - 1})")
239 out.append(layer)
240 if not out:
241 raise ValueError("--layers must be 'all' or comma-separated indices, e.g. --layers 0,2,8,15")
242 return out
243
244
246 name: str | None,
247 layer: int,
248 *,
249 multi: bool,
250 checkpoint: str | None = None,
251) -> str:
252 if name:
253 if "{layer}" in name:
254 return name.replace("{layer}", str(layer))
255 if multi:
256 return f"{name}-l{layer}"
257 return name
258 if checkpoint:
259 stem = Path(checkpoint).stem
260 return f"{stem}-l{layer}" if multi else stem
261 return f"base-l{layer}" if multi else "base"
262
263
265 *,
266 mid: str,
267 layer: int,
268 tag: str,
269 ctx: dict,
270 checkpoint: str | None,
271 corpus: str | None,
272 activations: str | None,
273 quick: bool,
274 balance: bool,
275 balance_group: str | None,
276 save_path: str | None,
277 max_steps: int | None = None,
278 max_epochs: int | None = None,
279) -> None:
280 from aquin.compute.sae_train import default_user_sae_path, train_sae
281
282 out = Path(save_path) if save_path else default_user_sae_path(mid, tag, layer)
283 sync_args = {
284 "model_id": mid,
285 "layer": layer,
286 "checkpoint": checkpoint or "",
287 "quick": quick,
288 "name": tag,
289 "output": str(out),
290 "activations": activations or "",
291 "balance": balance,
292 "group": balance_group or "",
293 "model_mode": "llm",
294 }
295 acts_note = f" activations={activations}" if activations else ""
296 bal_note = f" balance=True group={balance_group}" if balance and balance_group else (" balance=True" if balance else "")
297 print(f"[sae train] llm model={mid} layer={layer} quick={quick}{acts_note}{bal_note} -> {out}")
298
299 train_sae(
300 mid,
301 layer,
302 out,
303 checkpoint_path=checkpoint,
304 corpus_path=corpus,
305 quick=quick,
306 activations_dir=activations,
307 balance=balance,
308 balance_group=balance_group,
309 max_steps=max_steps,
310 max_epochs=max_epochs,
311 )
312
313 payload = {**sync_args, "status": "done", "output_path": str(out)}
314 print(f"[sae train] done: {out}")
315 print(f"[sae train] load for inspect: aquin load sae --user {tag} --layer {layer}")
316 print(f"[sae train] diff vs public: aquin diff sae --checkpoint <path> | aquin sae align --sae-a <public> --sae-b {out}")
317 _sync_sae_result(ctx, "run_sae_train", sync_args, payload)
318
319
320def cmd_sae_train(args: list[str]) -> None:
321 reject_legacy_output_flags(args)
322 model_id = _parse_flag(args, "--model")
323 checkpoint = _parse_flag(args, "--checkpoint")
324 layer_s = _parse_flag(args, "--layer")
325 layers_s = _parse_flag(args, "--layers")
326 corpus = _parse_flag(args, "--corpus")
327 activations = _parse_flag(args, "--activations")
328 name = _parse_flag(args, "--name")
329 save_path = _parse_flag(args, "--save")
330 quick = _has_flag(args, "--quick")
331 balance = _has_flag(args, "--balance")
332 balance_group = _parse_flag(args, "--group")
333 max_steps_s = _parse_flag(args, "--max-steps")
334 max_epochs_s = _parse_flag(args, "--max-epochs")
335 max_steps = int(max_steps_s) if max_steps_s else None
336 max_epochs = int(max_epochs_s) if max_epochs_s else None
337
338 if not layer_s and not layers_s:
339 print("Usage: aquin sae train [--model <id>] --layer <n> | --layers <all|0,1,2,...>")
340 print(" [--checkpoint <path>] [--corpus <json|jsonl>] [--activations <dir>]")
341 print(" [--quick] [--max-steps N] [--max-epochs N] [--balance] [--group <field>]")
342 print(" [--name <tag>] [--save <path>]")
343 print("")
344 print(" Full train (default): 2M tokens, up to 50k steps / 200 epochs per layer")
345 print(" --name 'lfm-instruct-l{layer}-v1' use {layer} when training multiple layers")
346 print("")
347 print(" aquin sae train --layers all --name 'lfm-instruct-l{layer}-v1'")
348 sys.exit(1)
349
350 if layer_s and layers_s:
351 print("Error: pass --layer or --layers, not both.")
352 sys.exit(1)
353
354 if save_path and layers_s:
355 print("Error: --save is only valid with a single --layer.")
356 sys.exit(1)
357
358 model_id = _resolve_train_model_id(model_id, activations)
359
361
362 from aquin.cli import _build_tool_ctx
363 from aquin.compute.model_loader import get_active_model_id, get_config
364
365 try:
366 mid = _resolve_sae_model(model_id)
367 except ValueError as e:
368 print(f"Error: {e}")
369 sys.exit(1)
370
371 n_layers = int(get_config(mid)["n_layers"])
372 if layers_s:
373 try:
374 layers = _parse_train_layers(layers_s, n_layers)
375 except ValueError as e:
376 print(f"Error: {e}")
377 sys.exit(1)
378 else:
379 layers = [int(layer_s)]
380
381 ctx = _build_tool_ctx(model_id=mid)
382 from aquin.engine.sync_dispatch import require_active_session
383
384 # Activations-only train never needs a resident GPU model.
385 if not activations or (get_active_model_id() or "").strip():
386 require_active_session(ctx, label="aquin sae train")
387
388 if checkpoint and not Path(checkpoint).exists():
389 print(f"Checkpoint not found: {checkpoint}")
390 sys.exit(1)
391
392 if activations and not Path(activations).expanduser().exists():
393 print(f"Activations directory not found: {activations}")
394 sys.exit(1)
395
396 if activations and corpus:
397 print("Note: --corpus is ignored when --activations is set (no forward passes).")
398 if balance_group and not balance:
399 print("Note: --group is ignored unless --balance is set.")
400 if balance and not activations:
401 print("Note: --balance currently applies when training from metadata-tagged capture dirs via --activations.")
402
403 multi = len(layers) > 1
404 if multi:
405 print(f"[sae train] training {len(layers)} layers: {layers}")
406
407 for i, layer in enumerate(layers):
408 if multi:
409 print(f"\n[sae train] layer {layer} ({i + 1}/{len(layers)})")
410 tag = _resolve_train_name(name, layer, multi=multi, checkpoint=checkpoint)
411 try:
413 mid=mid,
414 layer=layer,
415 tag=tag,
416 ctx=ctx,
417 checkpoint=checkpoint,
418 corpus=corpus,
419 activations=activations,
420 quick=quick,
421 balance=balance,
422 balance_group=balance_group,
423 save_path=save_path,
424 max_steps=max_steps,
425 max_epochs=max_epochs,
426 )
427 except Exception as e:
428 print(f"Error on layer {layer}: {e}")
429 sys.exit(1)
430
431 if multi:
432 print(f"\n[sae train] all {len(layers)} layers complete.")
433
434
435def _resolve_user_sae_path(model_id: str, run_name: str, layer: int) -> Path:
436 from aquin.compute.model_loader import resolve_model_id
437 from aquin.compute.user_sae import USER_SAE_ROOT
438
439 slug = resolve_model_id(model_id)
440 path = USER_SAE_ROOT / slug / run_name.format(layer=layer) / f"sae_layer{layer}.pt"
441 if not path.is_file():
442 raise FileNotFoundError(str(path))
443 return path
444
445
446def cmd_sae_catalog_metrics(args: list[str]) -> None:
447 model_id = _parse_flag(args, "--model")
448 sae_path_s = _parse_flag(args, "--sae")
449 layer_s = _parse_flag(args, "--layer")
450 user_run = _parse_flag(args, "--user-run")
451 layers_s = _parse_flag(args, "--layers")
452 skip_interp = _has_flag(args, "--skip-interp")
453 n_features_s = _parse_flag(args, "--n-features")
454
455 if not sae_path_s and not user_run:
456 print("Usage: aquin sae catalog-metrics [--model <id>]")
457 print(" --sae <path> --layer <n>")
458 print(" | --user-run <name> --layers <all|0,1,2,...>")
459 print(" [--skip-interp] [--n-features N]")
460 print("")
461 print(" Writes sae_layerN.umap.json + sae_layerN.metrics.json next to each checkpoint.")
462 print(" Then publish the SAE files from this directory.")
463 print("")
464 print(" aquin sae catalog-metrics --user-run 'lfm-230m-l{layer}-v1' --layers all")
465 sys.exit(1)
466
467 if user_run and not layers_s:
468 print("Error: --user-run requires --layers (e.g. --layers all)")
469 sys.exit(1)
470 if sae_path_s and not layer_s and not user_run:
471 print("Error: --sae requires --layer")
472 sys.exit(1)
473 if user_run and sae_path_s:
474 print("Error: pass --sae + --layer, or --user-run + --layers, not both.")
475 sys.exit(1)
476
477 model_id = _require_session_model(model_id)
479
480 from aquin.compute.model_loader import get_config
481 from aquin.compute.sae_catalog_metrics import compute_and_write_catalog_sidecars
482
483 try:
484 from aquin.compute.model_loader import resolve_model_id
485
486 mid = resolve_model_id(model_id)
487 except ValueError as e:
488 print(f"Error: {e}")
489 sys.exit(1)
490
491 n_layers = int(get_config(mid)["n_layers"])
492 n_features = int(n_features_s) if n_features_s else 8
493
494 if user_run:
495 try:
496 layers = _parse_train_layers(layers_s or "", n_layers)
497 except ValueError as e:
498 print(f"Error: {e}")
499 sys.exit(1)
500 jobs = [(layer, _resolve_user_sae_path(mid, user_run, layer)) for layer in layers]
501 else:
502 jobs = [(int(layer_s), Path(sae_path_s).expanduser().resolve())]
503
504 for i, (layer, sae_path) in enumerate(jobs):
505 if len(jobs) > 1:
506 print(f"\n[sae catalog-metrics] layer {layer} ({i + 1}/{len(jobs)})")
507 try:
508 compute_and_write_catalog_sidecars(
509 model_id=mid,
510 layer=layer,
511 sae_path=sae_path,
512 skip_interp=skip_interp,
513 n_features=n_features,
514 )
515 except Exception as e:
516 print(f"Error on layer {layer}: {e}")
517 sys.exit(1)
518
519 if len(jobs) > 1:
520 print(f"\n[sae catalog-metrics] all {len(jobs)} layers complete.")
521
522
523def _default_user_sae_checkpoint() -> Path | None:
524 from aquin.compute.user_sae import list_user_saes
525
526 rows = list_user_saes()
527 if not rows:
528 return None
529 smoke = [r for r in rows if str(r.get("name") or "") == "chat-smoke"]
530 pick = smoke[0] if smoke else rows[0]
531 path = pick.get("path")
532 return Path(str(path)).expanduser() if path else None
533
534
535def _is_dictionary_sae_path(path: Path) -> bool:
536 sae_root = (Path.home() / ".aquin" / "sae").resolve()
537 try:
538 resolved = path.expanduser().resolve()
539 except OSError:
540 return False
541 if not resolved.is_file():
542 return False
543 try:
544 resolved.relative_to(sae_root)
545 except ValueError:
546 return False
547 return resolved.name.startswith("sae_layer") and resolved.suffix == ".pt"
548
549
550def _resolve_align_sae_path(raw: str | None) -> Path | None:
552 if not raw:
553 return default
554 path = Path(raw).expanduser()
556 return path.resolve()
557 if default is not None:
558 print(
559 f"[sae align] {path} is not an SAE dictionary under ~/.aquin/sae; using {default}",
560 flush=True,
561 )
562 return default
563 return path if path.is_file() else None
564
565
566def cmd_sae_align(args: list[str]) -> None:
567 reject_legacy_output_flags(args)
568 sae_a = _parse_flag(args, "--sae-a") or _parse_flag(args, "--sae_a")
569 sae_b = _parse_flag(args, "--sae-b") or _parse_flag(args, "--sae_b")
570 save_path = _parse_flag(args, "--save")
571 max_f = _parse_flag(args, "--max-features") or _parse_flag(args, "--max_features")
572
573 path_a = _resolve_align_sae_path(sae_a)
574 path_b = _resolve_align_sae_path(sae_b)
575 if path_a is None or path_b is None:
576 print("Usage: aquin sae align --sae-a <path> --sae-b <path> [--save <map.json>] [--max-features N]")
577 print(" Dictionaries live under ~/.aquin/sae/user/<model>/<name>/sae_layerN.pt")
578 sys.exit(1)
579
581
582 from aquin.cli import _build_tool_ctx
583 from aquin.compute.sae import SparseAutoencoder
584 from aquin.compute.sae_diff import align_sae_decoders
585 import torch
586
587 ctx = _build_tool_ctx()
588 from aquin.engine.sync_dispatch import require_active_session
589
590 require_active_session(ctx, label="aquin sae align")
591 sync_args = {"sae_a": str(path_a), "sae_b": str(path_b), "save": save_path or ""}
592
593 from aquin.compute.device import resolve_compute_device
594 device = resolve_compute_device()
595 a = SparseAutoencoder.load(path_a, device=device)
596 b = SparseAutoencoder.load(path_b, device=device)
597 if not torch.isfinite(a.W_dec.data).all() or not torch.isfinite(b.W_dec.data).all():
598 print("[sae align] decoder has NaNs; matching encoder directions instead", flush=True)
599 if a.n_features != b.n_features or a.d_model != b.d_model:
600 print(f"Warning: shape mismatch A=({a.n_features},{a.d_model}) B=({b.n_features},{b.d_model})")
601
602 mapping = align_sae_decoders(
603 a, b, max_features=int(max_f) if max_f else 256,
604 )
605 cosines = [m["cosine"] for m in mapping]
606 mean_cos = sum(cosines) / len(cosines) if cosines else 0.0
607 print(f"aligned {len(mapping)} feature pairs mean cosine={mean_cos:.4f}")
608
609 low = sorted(mapping, key=lambda x: x["cosine"])[:5]
610 print("lowest cosine matches:")
611 for row in low:
612 print(f" A[{row['feature_a']}] <-> B[{row['feature_b']}] cos={row['cosine']:.4f}")
613
614 result = {"pairs": mapping, "mean_cosine": round(mean_cos, 6), "n_pairs": len(mapping)}
615 if save_path:
616 Path(save_path).write_text(json.dumps(result, indent=2), encoding="utf-8")
617 print(f"[sae align] wrote {save_path}")
618 _sync_sae_result(ctx, "run_sae_align", sync_args, result)
619
620
621def cmd_sae(args: list[str]) -> None:
622 _OLD_EMBED_SUBS = {
623 "features", "contrastive", "interp", "browser", "graph",
624 "circuit", "steer", "absorption", "polysemy", "faithfulness",
626
627 if not args or args[0] in ("-h", "--help", "help"):
628 print("Checkpoint SAE — train / align / publish on fine-tuned weights.")
629 print("")
630 print("Prerequisite: aquin load --model <id> && aquin load sae <model-l{n}>")
631 print(" Example SAE id: llama-3.2-1b-l8")
632 print("")
633 print("Usage: aquin sae <train|align|catalog-metrics|publish> ...")
634 print("")
635 print(" train --model <id> --layer <n> [--checkpoint <path>] [--quick] [--corpus <file>]")
636 print(" [--activations <dir>] [--name <tag>] [--save <path>]")
637 print("")
638 print(" align --sae-a <path> --sae-b <path> [--save <json>] [--max-features <n>]")
639 print("")
640 print(" catalog-metrics / publish (public catalog pipeline)")
641 print(" For base vs checkpoint feature deltas: aquin diff sae --checkpoint <path>")
642 print("")
643 print("Embedding SAE inspection:")
644 print(" aquin check features | contrastive | browser | graph | circuit")
645 print(" aquin check absorption | polysemy | faithfulness")
646 print(" aquin interp | steer | decomp (bare verbs)")
647 print("")
648 print("Docs: https://aquin.app/docs/checkpoint-sae")
649 return
650
651 sub = args[0]
652 rest = args[1:]
653
654 if sub in _OLD_EMBED_SUBS:
655 if sub in ("steer", "interp"):
656 print(f"Use: aquin {sub}", file=sys.stderr)
657 else:
658 print(f"Use: aquin check {sub}", file=sys.stderr)
659 sys.exit(1)
660
661 if sub == "train":
662 cmd_sae_train(rest)
663 elif sub == "align":
664 cmd_sae_align(rest)
665 elif sub in ("catalog-metrics", "catalog_metrics"):
667 elif sub == "publish":
668 from aquin.publish_public_sae import publish_public_sae
669
670 publish_public_sae(rest)
671 else:
672 print(f"Unknown sae subcommand: {sub}")
673 print("Run: aquin sae help")
674 sys.exit(1)
Path|None _resolve_align_sae_path(str|None raw)
Definition sae_cli.py:554
str _resolve_sae_model(str model_id)
Definition sae_cli.py:26
bool _is_dictionary_sae_path(Path path)
Definition sae_cli.py:539
None cmd_sae_catalog_metrics(list[str] args)
Definition sae_cli.py:450
str _resolve_train_model_id(str|None model_id, str|None activations)
Definition sae_cli.py:126
None _train_one_layer(*, str mid, int layer, str tag, dict ctx, str|None checkpoint, str|None corpus, str|None activations, bool quick, bool balance, str|None balance_group, str|None save_path, int|None max_steps=None, int|None max_epochs=None)
Definition sae_cli.py:283
Path _resolve_user_sae_path(str model_id, str run_name, int layer)
Definition sae_cli.py:439
str _resolve_train_name(str|None name, int layer, *, bool multi, str|None checkpoint=None)
Definition sae_cli.py:255
bool _has_flag(list[str] args, str name)
Definition sae_cli.py:22
None cmd_sae_align(list[str] args)
Definition sae_cli.py:570
list[int] _parse_train_layers(str spec, int n_layers)
Definition sae_cli.py:231
dict|None _build_sae_card(str tool_name, dict payload, dict|None sync_args=None)
Definition sae_cli.py:42
None _ensure_compute_env()
Definition sae_cli.py:33
None _sync_sae_result(dict ctx, str tool_name, dict args, dict payload)
Definition sae_cli.py:82
str|None _parse_flag(list[str] args, str name)
Definition sae_cli.py:15
None cmd_sae_diff(list[str] args)
Definition sae_cli.py:163
None cmd_sae(list[str] args)
Definition sae_cli.py:625
None _print_sae_diff(dict payload)
Definition sae_cli.py:89
None cmd_sae_train(list[str] args)
Definition sae_cli.py:324
str _require_session_model(str|None model_id)
Definition sae_cli.py:104
Path|None _default_user_sae_checkpoint()
Definition sae_cli.py:527