2"""aquin diff sae — base vs checkpoint SAE activation diff (via aquin diff sae)."""
6from pathlib
import Path
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):
18def _has_flag(args: list[str], name: str) -> bool:
23 """Return canonical model id."""
26 return resolve_model_id(model_id)
30 """Match cmd_tool: shared model + local HTTP server for web popovers."""
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 {}
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"),
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 {}
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:])),
82 sync_cli_result(ctx, tool_name, args, payload, card=card)
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]:
95 f
" f{row['feature_idx']:>5} base={row['base_act']:.4f} "
96 f
"ft={row['ft_act']:.4f} Δ={row['delta']:+.4f}"
103 locked = (get_active_model_id()
or "").strip()
105 print(
"Error: no model loaded.")
106 print(
" Run: aquin load --model <id>")
108 if model_id
and model_id != locked:
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.")
117 return resolve_model_id(locked)
123 """Session model, explicit --model, or capture-dir manifest (activations-only)."""
127 locked = (get_active_model_id()
or "").strip()
133 return resolve_model_id(model_id)
138 root = Path(activations).expanduser()
139 manifest = read_manifest(root)
142 for child
in root.glob(
"**/manifest.json"):
143 manifest = read_manifest(child.parent)
144 if manifest
and manifest.get(
"model_id"):
146 mid = (manifest
or {}).get(
"model_id")
if manifest
else None
147 if isinstance(mid, str)
and mid.strip():
149 return resolve_model_id(mid.strip())
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.")
160 reject_legacy_output_flags(args)
166 name =
_parse_flag(args,
"--name")
or Path(checkpoint
or "checkpoint").stem
170 print(
"Usage: aquin diff sae [--model <id>] --checkpoint <path> [--prompts <json|jsonl>]")
171 print(
" [--layer N] [--sae <path>] [--name <label>] [--save <json>]")
183 except ValueError
as e:
187 ckpt = Path(checkpoint)
188 if not ckpt.exists():
189 print(f
"Checkpoint not found: {ckpt}")
192 layer = int(layer_s)
if layer_s
else None
193 prompts = load_prompts(prompts_path)
194 ctx = _build_tool_ctx(model_id=mid)
197 require_active_session(ctx, label=
"aquin diff sae")
200 "checkpoint": str(ckpt),
202 "prompts": prompts_path
or "",
207 print(f
"[diff sae] llm base={mid} target={ckpt.name} prompts={len(prompts)}")
208 payload = run_sae_diff(
211 target_checkpoint=str(ckpt),
212 checkpoint_name=name,
216 except Exception
as e:
222 Path(save_path).write_text(json.dumps(payload, indent=2), encoding=
"utf-8")
223 print(f
"\n[diff sae] wrote {save_path}")
228 raw = (spec
or "").strip().lower()
229 if raw
in (
"all",
"*"):
230 return list(range(n_layers))
232 for part
in raw.split(
","):
237 if layer < 0
or layer >= n_layers:
238 raise ValueError(f
"Layer {layer} out of range for model (0..{n_layers - 1})")
241 raise ValueError(
"--layers must be 'all' or comma-separated indices, e.g. --layers 0,2,8,15")
250 checkpoint: str |
None =
None,
253 if "{layer}" in name:
254 return name.replace(
"{layer}", str(layer))
256 return f
"{name}-l{layer}"
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"
270 checkpoint: str |
None,
272 activations: str |
None,
275 balance_group: str |
None,
276 save_path: str |
None,
277 max_steps: int |
None =
None,
278 max_epochs: int |
None =
None,
282 out = Path(save_path)
if save_path
else default_user_sae_path(mid, tag, layer)
286 "checkpoint": checkpoint
or "",
290 "activations": activations
or "",
292 "group": balance_group
or "",
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}")
303 checkpoint_path=checkpoint,
306 activations_dir=activations,
308 balance_group=balance_group,
310 max_epochs=max_epochs,
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}")
321 reject_legacy_output_flags(args)
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
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>]")
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")
347 print(
" aquin sae train --layers all --name 'lfm-instruct-l{layer}-v1'")
350 if layer_s
and layers_s:
351 print(
"Error: pass --layer or --layers, not both.")
354 if save_path
and layers_s:
355 print(
"Error: --save is only valid with a single --layer.")
367 except ValueError
as e:
371 n_layers = int(get_config(mid)[
"n_layers"])
375 except ValueError
as e:
379 layers = [int(layer_s)]
381 ctx = _build_tool_ctx(model_id=mid)
385 if not activations
or (get_active_model_id()
or "").strip():
386 require_active_session(ctx, label=
"aquin sae train")
388 if checkpoint
and not Path(checkpoint).exists():
389 print(f
"Checkpoint not found: {checkpoint}")
392 if activations
and not Path(activations).expanduser().exists():
393 print(f
"Activations directory not found: {activations}")
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.")
403 multi = len(layers) > 1
405 print(f
"[sae train] training {len(layers)} layers: {layers}")
407 for i, layer
in enumerate(layers):
409 print(f
"\n[sae train] layer {layer} ({i + 1}/{len(layers)})")
417 checkpoint=checkpoint,
419 activations=activations,
422 balance_group=balance_group,
425 max_epochs=max_epochs,
427 except Exception
as e:
428 print(f
"Error on layer {layer}: {e}")
432 print(f
"\n[sae train] all {len(layers)} layers complete.")
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))
452 skip_interp =
_has_flag(args,
"--skip-interp")
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]")
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.")
464 print(
" aquin sae catalog-metrics --user-run 'lfm-230m-l{layer}-v1' --layers all")
467 if user_run
and not layers_s:
468 print(
"Error: --user-run requires --layers (e.g. --layers all)")
470 if sae_path_s
and not layer_s
and not user_run:
471 print(
"Error: --sae requires --layer")
473 if user_run
and sae_path_s:
474 print(
"Error: pass --sae + --layer, or --user-run + --layers, not both.")
486 mid = resolve_model_id(model_id)
487 except ValueError
as e:
491 n_layers = int(get_config(mid)[
"n_layers"])
492 n_features = int(n_features_s)
if n_features_s
else 8
497 except ValueError
as e:
502 jobs = [(int(layer_s), Path(sae_path_s).expanduser().resolve())]
504 for i, (layer, sae_path)
in enumerate(jobs):
506 print(f
"\n[sae catalog-metrics] layer {layer} ({i + 1}/{len(jobs)})")
508 compute_and_write_catalog_sidecars(
512 skip_interp=skip_interp,
513 n_features=n_features,
515 except Exception
as e:
516 print(f
"Error on layer {layer}: {e}")
520 print(f
"\n[sae catalog-metrics] all {len(jobs)} layers complete.")
526 rows = list_user_saes()
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
536 sae_root = (Path.home() /
".aquin" /
"sae").resolve()
538 resolved = path.expanduser().resolve()
541 if not resolved.is_file():
544 resolved.relative_to(sae_root)
547 return resolved.name.startswith(
"sae_layer")
and resolved.suffix ==
".pt"
554 path = Path(raw).expanduser()
556 return path.resolve()
557 if default
is not None:
559 f
"[sae align] {path} is not an SAE dictionary under ~/.aquin/sae; using {default}",
563 return path
if path.is_file()
else None
567 reject_legacy_output_flags(args)
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")
587 ctx = _build_tool_ctx()
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 ""}
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})")
602 mapping = align_sae_decoders(
603 a, b, max_features=int(max_f)
if max_f
else 256,
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}")
609 low = sorted(mapping, key=
lambda x: x[
"cosine"])[:5]
610 print(
"lowest cosine matches:")
612 print(f
" A[{row['feature_a']}] <-> B[{row['feature_b']}] cos={row['cosine']:.4f}")
614 result = {
"pairs": mapping,
"mean_cosine": round(mean_cos, 6),
"n_pairs": len(mapping)}
616 Path(save_path).write_text(json.dumps(result, indent=2), encoding=
"utf-8")
617 print(f
"[sae align] wrote {save_path}")
621def cmd_sae(args: list[str]) ->
None:
623 "features",
"contrastive",
"interp",
"browser",
"graph",
624 "circuit",
"steer",
"absorption",
"polysemy",
"faithfulness",
627 if not args
or args[0]
in (
"-h",
"--help",
"help"):
628 print(
"Checkpoint SAE — train / align / publish on fine-tuned weights.")
630 print(
"Prerequisite: aquin load --model <id> && aquin load sae <model-l{n}>")
631 print(
" Example SAE id: llama-3.2-1b-l8")
633 print(
"Usage: aquin sae <train|align|catalog-metrics|publish> ...")
635 print(
" train --model <id> --layer <n> [--checkpoint <path>] [--quick] [--corpus <file>]")
636 print(
" [--activations <dir>] [--name <tag>] [--save <path>]")
638 print(
" align --sae-a <path> --sae-b <path> [--save <json>] [--max-features <n>]")
640 print(
" catalog-metrics / publish (public catalog pipeline)")
641 print(
" For base vs checkpoint feature deltas: aquin diff sae --checkpoint <path>")
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)")
648 print(
"Docs: https://aquin.app/docs/checkpoint-sae")
654 if sub
in _OLD_EMBED_SUBS:
655 if sub
in (
"steer",
"interp"):
656 print(f
"Use: aquin {sub}", file=sys.stderr)
658 print(f
"Use: aquin check {sub}", file=sys.stderr)
665 elif sub
in (
"catalog-metrics",
"catalog_metrics"):
667 elif sub ==
"publish":
670 publish_public_sae(rest)
672 print(f
"Unknown sae subcommand: {sub}")
673 print(
"Run: aquin sae help")
Path|None _resolve_align_sae_path(str|None raw)
str _resolve_sae_model(str model_id)
bool _is_dictionary_sae_path(Path path)
None cmd_sae_catalog_metrics(list[str] args)
str _resolve_train_model_id(str|None model_id, str|None activations)
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)
Path _resolve_user_sae_path(str model_id, str run_name, int layer)
str _resolve_train_name(str|None name, int layer, *, bool multi, str|None checkpoint=None)
bool _has_flag(list[str] args, str name)
None cmd_sae_align(list[str] args)
list[int] _parse_train_layers(str spec, int n_layers)
dict|None _build_sae_card(str tool_name, dict payload, dict|None sync_args=None)
None _ensure_compute_env()
None _sync_sae_result(dict ctx, str tool_name, dict args, dict payload)
str|None _parse_flag(list[str] args, str name)
None cmd_sae_diff(list[str] args)
None cmd_sae(list[str] args)
None _print_sae_diff(dict payload)
None cmd_sae_train(list[str] args)
str _require_session_model(str|None model_id)
Path|None _default_user_sae_checkpoint()