AQIT 0.1.0
Loading...
Searching...
No Matches
capture_cli.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""aquin capture-activations"""
3
4from __future__ import annotations
5
6import sys
7from pathlib import Path
8
9DEFAULT_PROBE_COUNT = 6
10
11
12def _parse_flag(args: list[str], name: str) -> str | None:
13 for i, a in enumerate(args):
14 if a == name and i + 1 < len(args):
15 return args[i + 1]
16 return None
17
18
19def _has_flag(args: list[str], name: str) -> bool:
20 return name in args
21
22
23def _parse_count(args: list[str]) -> int:
24 raw = _parse_flag(args, "--count")
25 if raw is None:
26 return DEFAULT_PROBE_COUNT
27 try:
28 return max(1, min(int(raw), 64))
29 except ValueError:
30 print("Error: --count must be an integer between 1 and 64")
31 sys.exit(1)
32
33
34def _ensure_compute_env() -> None:
35 from aquin.compute.loader_shim import apply as _shim_apply
36 from aquin.engine.local_server import start as _start_local_server
37
38 _shim_apply()
39 _start_local_server()
40
41
42def _build_capture_card(payload: dict, sync_args: dict) -> dict:
43 probes_source = payload.get("probes_source") or sync_args.get("probes_source")
44 return {
45 "type": "activationCapture",
46 "data": {
47 "modelId": payload.get("model_id") or sync_args.get("model_id"),
48 "modelMode": payload.get("model_mode"),
49 "nProbes": payload.get("n_probes"),
50 "layers": payload.get("layers") or [],
51 "position": payload.get("position") or sync_args.get("position"),
52 "encodeSae": bool(payload.get("encode_sae") or sync_args.get("encode_sae")),
53 "dModel": payload.get("d_model"),
54 "outputDir": payload.get("output_dir") or sync_args.get("dir"),
55 "manifestPath": payload.get("manifest_path"),
56 "checkpoint": payload.get("checkpoint") or sync_args.get("checkpoint") or None,
57 "captureName": payload.get("capture_id") or sync_args.get("name"),
58 "probeSamples": payload.get("probe_samples") or [],
59 "saeFeatures": payload.get("sae_features"),
60 "probesSource": probes_source,
61 "topic": payload.get("topic") or sync_args.get("topic"),
62 "status": payload.get("status", "done"),
63 },
64 }
65
66
67def _require_loaded_model_id() -> str:
68 from aquin.compute.model_loader import get_active_model_id, resolve_model_id
69
70 active = (get_active_model_id() or "").strip()
71 if not active:
72 print("Error: no model loaded. Run: aquin load --model <id>")
73 sys.exit(1)
74 try:
75 return resolve_model_id(active)
76 except ValueError as e:
77 print(f"Error: {e}")
78 sys.exit(1)
79
80def cmd_capture_activations(args: list[str]) -> None:
81 if _has_flag(args, "--help") or _has_flag(args, "-h"):
83 return
85 requested_model = _parse_flag(args, "--model")
86 if requested_model is not None and _parse_flag(args, "--checkpoint") is None:
87 print("Error: capture-activations uses the loaded session model only.")
88 print(" Pass --model <id> only when using --checkpoint <path>.")
89 sys.exit(1)
90
91 prompts_path = _parse_flag(args, "--prompts")
92 out_dir_arg = _parse_flag(args, "--dir") or _parse_flag(args, "--output")
93 layers_spec = _parse_flag(args, "--layers")
94 checkpoint = _parse_flag(args, "--checkpoint")
95 name = _parse_flag(args, "--name")
96 topic = _parse_flag(args, "--topic")
97 position = _parse_flag(args, "--position") or "last"
98 granularity = _parse_flag(args, "--granularity") or "prompt"
99 balance_group = _parse_flag(args, "--group")
100 sae_layer_s = _parse_flag(args, "--sae-layer")
101 encode_sae = _has_flag(args, "--encode-sae")
102 balance = _has_flag(args, "--balance")
103 count = _parse_count(args)
104
105 if not out_dir_arg:
106 print("Error: --dir <path> is required (alias: --output <path>).", file=sys.stderr)
107 print(" Example: aquin capture-activations --dir ./captures/my-run", file=sys.stderr)
108 sys.exit(1)
109
110 if position not in ("last", "mean"):
111 print("Error: --position must be last or mean")
112 sys.exit(1)
113 if granularity not in ("prompt", "token"):
114 print("Error: --granularity must be prompt or token")
115 sys.exit(1)
116
118
119 from aquin.cli import _build_tool_ctx
121 llm_layer_count,
122 parse_layers,
123 resolve_capture_model_id,
124 resolve_probes_for_capture,
125 run_capture_activations,
126 )
127 from aquin.compute.model_loader import resolve_model_id
128 from aquin.engine.sync_dispatch import require_active_session, sync_cli_result
129
130 try:
131 if requested_model:
132 mid = resolve_model_id(requested_model)
133 model_mode = "llm"
134 else:
135 mid, model_mode = resolve_capture_model_id(_require_loaded_model_id())
136 except ValueError as e:
137 print(f"Error: {e}")
138 sys.exit(1)
139
140 if checkpoint and not Path(checkpoint).exists():
141 print(f"Checkpoint not found: {checkpoint}")
142 sys.exit(1)
143
144 out_dir = Path(out_dir_arg)
145 out_dir.mkdir(parents=True, exist_ok=True)
146
147 try:
148 probes, probe_meta = resolve_probes_for_capture(
149 model_id=mid,
150 model_mode=model_mode,
151 prompts_path=prompts_path,
152 count=count,
153 topic=topic,
154 balance=balance,
155 balance_group=balance_group,
156 output_dir=out_dir if not prompts_path else None,
157 )
158 except (FileNotFoundError, ValueError) as e:
159 print(f"Error: {e}")
160 sys.exit(1)
161
162 try:
163 n_layers = llm_layer_count(mid, checkpoint)
164 layer_list = parse_layers(layers_spec, n_layers)
165 except ValueError as e:
166 print(f"Error: {e}")
167 sys.exit(1)
168
169 sae_layer = int(sae_layer_s) if sae_layer_s else None
170 ckpt_name = name or (Path(checkpoint).stem if checkpoint else "base")
171
172 ctx = _build_tool_ctx(model_id=mid)
173 require_active_session(ctx, label="aquin capture-activations")
174 sync_args = {
175 "model_id": mid,
176 "model_mode": model_mode,
177 "prompts": probe_meta.get("prompts_path") or probe_meta.get("generated_probes_path") or "",
178 "probes_source": probe_meta.get("probes_source"),
179 "topic": probe_meta.get("topic") or topic,
180 "count": count,
181 "dir": str(out_dir),
182 "layers": layers_spec or "all",
183 "checkpoint": checkpoint or "",
184 "name": ckpt_name,
185 "position": position,
186 "granularity": granularity,
187 "encode_sae": encode_sae,
188 "balance": balance,
189 "group": balance_group or "",
190 }
191
192 print(
193 f"[capture] mode={model_mode} model={mid} probes={len(probes)} "
194 f"source={probe_meta.get('probes_source')} layers={layer_list} "
195 f"position={position} granularity={granularity} balance={balance}"
196 f"{f' group={balance_group}' if balance_group else ''} → {out_dir}"
197 )
198 try:
199 payload = run_capture_activations(
200 mid,
201 probes,
202 out_dir,
203 model_mode=model_mode,
204 layers=layer_list,
205 checkpoint_path=checkpoint,
206 checkpoint_name=ckpt_name if checkpoint else None,
207 position=position, # type: ignore[arg-type]
208 granularity=granularity, # type: ignore[arg-type]
209 encode_sae=encode_sae,
210 sae_layer=sae_layer,
211 capture_name=ckpt_name,
212 manifest_extras=probe_meta,
213 )
214 except Exception as e:
215 print(f"Error: {e}")
216 sys.exit(1)
217
218 print(f"[capture] wrote {payload['manifest_path']}")
219 if payload.get("metadata_path"):
220 print(f"[capture] wrote {payload['metadata_path']}")
221 for layer, rel in payload.get("activation_files", {}).items():
222 print(
223 f"[capture] layer {layer}: {rel} "
224 f"shape=({payload['n_probes']}, {payload.get('d_model', 'd_model')})"
225 )
226 if payload.get("sae_features"):
227 sf = payload["sae_features"]
228 print(f"[capture] sae layer {sf['layer']}: n_features={sf['n_features']}")
229
230 card = _build_capture_card(payload, sync_args)
231 sync_cli_result(ctx, "run_capture_activations", sync_args, payload, card=card)
232
233
234def _print_help() -> None:
235 print("Batch activation capture for labeled probe sets (LLM).")
236 print("")
237 print("Prerequisite: aquin load --model <id>")
238 print("")
239 print("Usage: aquin capture-activations --dir <path>")
240 print(" (alias: --output <path>)")
241 print(" [--prompts <json|jsonl>] (omit to auto-generate probes)")
242 print(" [--count N] [--topic <text>] (default count=6; LLM generates, embed uses templates)")
243 print(" [--layers all|8,15] [--checkpoint <path>] (LLM only)")
244 print(" [--position last|mean] [--granularity prompt|token] [--encode-sae] [--sae-layer N] [--name <label>] [--balance] [--group <field>]")
245 print("")
246 print(" Without --prompts: loaded LLM writes probe lines.")
247 print(" Generated probes saved to <dir>/probes.jsonl.")
248 print(" Always writes root manifest.json + metadata.json + summary.jsonl.")
249 print(" --granularity token saves per-token activations + token_spans.jsonl.")
250 print(" --balance samples evenly across metadata groups (label, stressor, lang, group) when present.")
251 print(" --group chooses the metadata field used with --balance.")
252 print(" Syncs activationCapture card.")
253 print("")
254 print("Docs: https://aquin.app/docs/sae-training")
str _require_loaded_model_id()
dict _build_capture_card(dict payload, dict sync_args)
bool _has_flag(list[str] args, str name)
None _ensure_compute_env()
str|None _parse_flag(list[str] args, str name)
None cmd_capture_activations(list[str] args)
int _parse_count(list[str] args)