AQIT 0.1.0
Loading...
Searching...
No Matches
find_feature_cli.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""aquin feature locate — rank SAE features for honest vs deceptive probes."""
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 _parse_int_flag(args: list[str], name: str, default: int) -> int:
23 raw = _parse_flag(args, name)
24 if raw is None:
25 return default
26 try:
27 return int(raw)
28 except ValueError:
29 print(f"Error: {name} must be an integer")
30 sys.exit(1)
31
32
33def _ensure_compute_env() -> None:
34 from aquin.compute.loader_shim import apply as _shim_apply
35 from aquin.engine.local_server import start as _start_local_server
36
37 _shim_apply()
38 _start_local_server()
39
40
41def _require_loaded_model_id() -> str:
42 from aquin.compute.model_loader import get_active_model_id, resolve_model_id
43
44 active = (get_active_model_id() or "").strip()
45 if not active:
46 print("Error: no model loaded. Run: aquin load --model <id>")
47 sys.exit(1)
48 try:
49 return resolve_model_id(active)
50 except ValueError as e:
51 print(f"Error: {e}")
52 sys.exit(1)
53
54def _build_find_feature_card(payload: dict, sync_args: dict) -> dict:
55 return {
56 "type": "findFeature",
57 "data": {
58 "modelId": payload.get("model_id") or sync_args.get("model_id"),
59 "layer": payload.get("layer"),
60 "scorer": payload.get("scorer"),
61 "nHonest": payload.get("n_honest"),
62 "nDeceptive": payload.get("n_deceptive"),
63 "promptsPath": payload.get("prompts_path") or sync_args.get("prompts"),
64 "checkpoint": payload.get("checkpoint") or sync_args.get("checkpoint"),
65 "chosenFeatureIdx": payload.get("chosen_feature_idx"),
66 "chosenDelta": payload.get("chosen_delta"),
67 "direction": payload.get("direction"),
68 "conditioning": payload.get("conditioning"),
69 "behavior": payload.get("behavior"),
70 "warning": payload.get("warning"),
71 "persistedKey": payload.get("persisted_key"),
72 "experimentPath": payload.get("experiment_path"),
73 "rankings": payload.get("rankings") or [],
74 "status": payload.get("status", "done"),
75 },
76 }
77
78
79def _print_find_feature(payload: dict) -> None:
80 print(f"model : {payload.get('model_id')}")
81 print(f"layer : {payload.get('layer')}")
82 print(f"scorer : {payload.get('scorer')}")
83 if payload.get("direction"):
84 print(f"direction : {payload.get('direction')}")
85 if payload.get("conditioning"):
86 print(f"condition : {payload.get('conditioning')}")
87 behavior = payload.get("behavior")
88 if isinstance(behavior, dict):
89 print(
90 f"behavior : {behavior.get('n_truthful', '?')} truthful · "
91 f"{behavior.get('n_deceptive', '?')} deceptive · "
92 f"{behavior.get('n_ambiguous', '?')} ambiguous "
93 f"({behavior.get('n_generated', '?')} generated)"
94 )
95 print(f"probes : {payload.get('n_honest')} honest · {payload.get('n_deceptive')} deceptive")
96 if payload.get("prompts_path"):
97 print(f"prompts : {payload.get('prompts_path')}")
98 chosen = payload.get("chosen_feature_idx")
99 if chosen is not None:
100 print(f"\nchosen : feature {chosen} Δ={payload.get('chosen_delta')}")
101 if payload.get("persisted_key"):
102 print(f"persisted : {payload.get('persisted_key')} → {payload.get('experiment_path')}")
103 if payload.get("warning"):
104 print(f"\nwarning : {payload.get('warning')}")
105 rankings = payload.get("rankings") or []
106 if rankings:
107 direction = payload.get("direction") or "both"
108 rank_label = {
109 "deceptive": "top features by Δ (deceptive > honest)",
110 "honest": "top features by |Δ| (honest > deceptive)",
111 "both": "top features by |Δ| (deceptive − honest)",
112 }.get(direction, "top features")
113 print(f"\n{rank_label} ({min(len(rankings), 10)} shown):")
114 for i, row in enumerate(rankings[:10], 1):
115 interp = row.get("interp_score")
116 extra = f" interp={interp:.2f}" if interp is not None else ""
117 print(
118 f" {i:2}. f{row['feature_idx']:<5} "
119 f"honest={row['honest_mean']:.4f} deceptive={row['deceptive_mean']:.4f} "
120 f"Δ={row['delta']:+.4f}{extra}"
121 )
122
123
124def cmd_find_feature(args: list[str]) -> None:
125 if _has_flag(args, "--help") or _has_flag(args, "-h"):
127 return
129 if _parse_flag(args, "--model") is not None:
130 print("Error: feature locate uses the loaded session model only.")
131 print(" Run: aquin load --model <id>")
132 sys.exit(1)
133
134 reject_legacy_output_flags(args)
135 scorer = _parse_flag(args, "--scorer") or "deception"
136 direction = _parse_flag(args, "--direction") or "both"
137 conditioning = _parse_flag(args, "--conditioning") or "behavior"
138 prompts = _parse_flag(args, "--prompts")
139 layer_s = _parse_flag(args, "--layer")
140 checkpoint = _parse_flag(args, "--checkpoint")
141 persist = _parse_flag(args, "--persist")
142 save_path = _parse_flag(args, "--save")
143 top_k = _parse_int_flag(args, "--top", 20)
144 benchmark_top = _parse_int_flag(args, "--benchmark-top", 0)
145 want_umap = _has_flag(args, "--umap")
146
147 if checkpoint and not Path(checkpoint).exists():
148 print(f"Checkpoint not found: {checkpoint}")
149 sys.exit(1)
150
152
153 from aquin.cli import _build_tool_ctx
154 from aquin.engine.sync_dispatch import require_active_session, sync_cli_result
155
157 layer = int(layer_s) if layer_s else None
158 ctx = _build_tool_ctx(model_id=mid)
159 require_active_session(ctx, label="aquin feature locate")
160
161 sync_args = {
162 "model_id": mid,
163 "scorer": scorer,
164 "direction": direction,
165 "conditioning": conditioning,
166 "prompts": prompts or "",
167 "layer": layer,
168 "checkpoint": checkpoint or "",
169 "top_k": top_k,
170 "persist": persist or "",
171 }
172
173 openai_client = None
174 if benchmark_top > 0:
175 try:
176 from aquin.compute.openai_client import get_openai_client
177
178 openai_client = get_openai_client(ctx)
179 except Exception:
180 print("[feature locate] --benchmark-top ignored (OpenAI client unavailable)", flush=True)
181 benchmark_top = 0
182
183 session_id = ctx.get("session_id") or ctx.get("state", {}).get("session_id")
184
185 try:
186 from aquin.compute.find_feature import run_find_feature
187
188 print(f"[feature locate] model={mid} scorer={scorer} direction={direction} conditioning={conditioning} top={top_k}")
189 payload = run_find_feature(
190 mid,
191 scorer=scorer,
192 prompts_path=prompts,
193 layer=layer,
194 checkpoint_path=checkpoint,
195 top_k=top_k,
196 direction=direction,
197 conditioning=conditioning,
198 benchmark_top=benchmark_top,
199 persist_key=persist,
200 session_id=str(session_id) if session_id else None,
201 openai_client=openai_client,
202 )
203 except Exception as e:
204 print(f"Error: {e}")
205 sys.exit(1)
206
207 _print_find_feature(payload)
208 if save_path:
209 Path(save_path).write_text(json.dumps(payload, indent=2), encoding="utf-8")
210 print(f"\n[feature locate] wrote {save_path}")
211
212 card = _build_find_feature_card(payload, sync_args)
213 sync_cli_result(ctx, "run_find_feature", sync_args, payload, card=card)
214
215 if want_umap:
216 from aquin.cli import _run_umap_followup
217
218 _run_umap_followup(
219 ctx,
220 result=payload,
221 tool_args=sync_args,
222 ensure_model=mid,
223 layer=layer,
224 )
225
226
227def _print_help() -> None:
228 print("Rank SAE features that separate honest vs deceptive probes (LLM).")
229 print("")
230 print("Prerequisite: aquin load --model <id>")
231 print(" aquin load sae <model-l{n}>")
232 print("")
233 print("Usage: aquin feature locate [--scorer deception] [--prompts <json|jsonl>]")
234 print(" [--layer N] [--checkpoint <path>] [--top N] [--direction both|deceptive|honest]")
235 print(" [--conditioning behavior|prompt] [--benchmark-top K] [--persist <key>] [--save <json>]")
236 print(" [--umap]")
237 print("")
238 print(" --prompts is required (honest/deceptive JSON or JSONL)")
239 print(" --conditioning behavior (default): generate completions, classify output, bucket by behavior")
240 print(" --conditioning prompt: legacy static encoding on probe text only")
241 print(" --direction both (default): rank by |Δ|; deceptive: only Δ>0; honest: only Δ<0")
242 print(" --persist writes chosen feature to ~/.aquin/experiments/<model>.json + session memory")
243 print(" --umap loads SAE UMAP projection after ranking (web explorer)")
244 print(" Syncs findFeature card to the web orchestrator.")
245 print("")
246 print("Docs: https://aquin.app/docs/deception")
None cmd_find_feature(list[str] args)
bool _has_flag(list[str] args, str name)
str|None _parse_flag(list[str] args, str name)
dict _build_find_feature_card(dict payload, dict sync_args)
int _parse_int_flag(list[str] args, str name, int default)
None _print_find_feature(dict payload)