AQIT 0.1.0
Loading...
Searching...
No Matches
steer_vector.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""LAT steering-vector extract, persist, and load for SAE decoder directions."""
3from __future__ import annotations
4
5import json
6from datetime import datetime, timezone
7from pathlib import Path
8from typing import Any
9
10import torch
11
12SCHEMA_VERSION = 1
13KIND = "sae_decoder_steer"
14
15
16def _now_iso() -> str:
17 return datetime.now(timezone.utc).isoformat()
18
19
20def _experiment_path(model_id: str) -> Path:
21 slug = model_id.replace("/", "--")
22 return Path.home() / ".aquin" / "experiments" / f"{slug}.json"
23
25def lookup_probe_record(model_id: str, probe_id: str | None = None) -> dict[str, Any] | None:
26 """Load persisted find-feature record from ~/.aquin/experiments/<model>.json."""
27 path = _experiment_path(model_id)
28 if not path.is_file():
29 return None
30 try:
31 data = json.loads(path.read_text(encoding="utf-8"))
32 except Exception:
33 return None
34 if not isinstance(data, dict):
35 return None
36 key = probe_id or "deception_feature"
37 rec = data.get(key)
38 return rec if isinstance(rec, dict) else None
39
40
42 model_id: str,
43 feature_idx: int,
44 *,
45 layer: int | None = None,
46 apply_norm_std: bool = True,
47) -> tuple[torch.Tensor, int, dict[str, Any]]:
48 """Return unit-ready decoder direction (d_model,) and resolved layer."""
49 from aquin.compute.feature_analysis import load_norm, load_sae
50 from aquin.compute.model_loader import get_config
51
52 cfg = get_config(model_id)
53 resolved_layer = int(layer if layer is not None else cfg["sae_layer"])
54 try:
55 sae = load_sae(model_id, resolved_layer)
56 except FileNotFoundError:
57 from aquin.compute.model_loader import get_available_sae_layers, get_sae_layer
58
59 available = get_available_sae_layers(model_id)
60 fallback = int(get_sae_layer(model_id))
61 if available:
62 resolved_layer = fallback if fallback in available else int(available[0])
63 sae = load_sae(model_id, resolved_layer)
64 else:
65 raise
66 if feature_idx < 0 or feature_idx >= sae.W_dec.shape[0]:
67 raise ValueError(f"feature_idx {feature_idx} out of range (0..{sae.W_dec.shape[0] - 1})")
68
69 feat_dir = sae.W_dec[feature_idx].detach().float().clone()
70 if not torch.isfinite(feat_dir).all():
71 enc = sae.W_enc.data.detach().float()
72 if enc.ndim == 2 and enc.shape[1] > feature_idx:
73 feat_dir = enc[:, feature_idx].clone()
74 feat_dir = torch.nan_to_num(feat_dir, nan=0.0, posinf=0.0, neginf=0.0)
75 norm_stats: dict[str, Any] | None = None
76 norm = load_norm(model_id, resolved_layer) if apply_norm_std else None
77 if norm is not None and norm.get("std") is not None:
78 # Keep std on feat_dir's device — SAE may be CUDA while a naive .cpu()
79 # std would raise "cuda:0 and cpu" on the multiply below.
80 std = norm["std"].detach().float().to(device=feat_dir.device)
81 feat_dir = feat_dir * std
82 mean = norm.get("mean")
83 norm_stats = {
84 "mean": mean.detach().float().cpu().tolist() if mean is not None else None,
85 "std": std.detach().cpu().tolist(),
86 "applied_to_vector": True,
87 }
88
89 meta = {
90 "d_model": int(feat_dir.shape[-1]),
91 "vector_l2_norm": round(float(feat_dir.norm().item()), 6),
92 "norm": norm_stats,
93 }
94 return feat_dir, resolved_layer, meta
95
96
98 model_id: str,
99 feature_idx: int,
100 *,
101 layer: int | None = None,
102 feature_label: str | None = None,
103 probe_id: str | None = None,
104) -> dict[str, Any]:
105 direction, resolved_layer, vec_meta = resolve_feature_direction(
106 model_id, feature_idx, layer=layer,
107 )
108 probe_record = lookup_probe_record(model_id, probe_id)
109 label = feature_label or f"F{feature_idx}"
110
111 return {
112 "schema_version": SCHEMA_VERSION,
113 "kind": KIND,
114 "model_id": model_id,
115 "layer": resolved_layer,
116 "feature_idx": int(feature_idx),
117 "feature_label": label or f"F{feature_idx}",
118 "probe_id": probe_id or ("deception_feature" if probe_record else None),
119 "probe_record": probe_record,
120 "d_model": vec_meta["d_model"],
121 "vector_l2_norm": vec_meta["vector_l2_norm"],
122 "norm": vec_meta.get("norm"),
123 "vector": direction.cpu().tolist(),
124 "created_at": _now_iso(),
125 "source": "steer --save",
126 }
127
128
129def write_steer_vector_file(payload: dict[str, Any], output: str | Path) -> Path:
130 out = Path(output).expanduser().resolve()
131 out.parent.mkdir(parents=True, exist_ok=True)
132 out.write_text(json.dumps(payload, indent=2), encoding="utf-8")
133 return out
134
135
136def load_steer_vector_file(path: str | Path) -> dict[str, Any]:
137 p = Path(path).expanduser().resolve()
138 if not p.is_file():
139 raise FileNotFoundError(f"Steering vector file not found: {p}")
140 try:
141 data = json.loads(p.read_text(encoding="utf-8"))
142 except json.JSONDecodeError as exc:
143 raise ValueError(f"Invalid steering vector JSON: {p}") from exc
144 if not isinstance(data, dict):
145 raise ValueError(f"Steering vector file must be a JSON object: {p}")
146 if data.get("kind") != KIND:
147 raise ValueError(f"Unsupported steering vector kind: {data.get('kind')!r}")
148 if "vector" not in data or not isinstance(data["vector"], list) or not data["vector"]:
149 raise ValueError("Steering vector file missing non-empty 'vector' array")
150 return data
151
152
154 data: dict[str, Any],
155 *,
156 device: torch.device | str,
157 dtype: torch.dtype,
158) -> torch.Tensor:
159 vec = torch.tensor(data["vector"], device=device, dtype=dtype)
160 if vec.ndim != 1:
161 raise ValueError("Steering vector must be a 1-D array")
162 return vec
163
164
166 model_id: str,
167 feature_idx: int,
168 output: str | Path,
169 *,
170 layer: int | None = None,
171 feature_label: str | None = None,
172 probe_id: str | None = None,
173) -> dict[str, Any]:
175 model_id,
176 feature_idx,
177 layer=layer,
178 feature_label=feature_label,
179 probe_id=probe_id,
180 )
181 out_path = write_steer_vector_file(payload, output)
182 return {
183 "status": "done",
184 "output_path": str(out_path),
185 "model_id": payload["model_id"],
186 "layer": payload["layer"],
187 "feature_idx": payload["feature_idx"],
188 "feature_label": payload["feature_label"],
189 "probe_id": payload.get("probe_id"),
190 "d_model": payload["d_model"],
191 "vector_l2_norm": payload["vector_l2_norm"],
192 "norm": payload.get("norm"),
193 "source_model_id": payload["model_id"],
194 "source_layer": payload["layer"],
195 "source_feature_idx": payload["feature_idx"],
196 }
197
198
200 model: Any,
201 prompt: str,
202 *,
203 steer_vec: torch.Tensor,
204 steer_strength: float,
205 layer: int,
206 max_new_tokens: int,
207) -> str:
208 from aquin.compute.causal_trace import _format_prompt
209
210 def _steer_hook(value, hook):
211 vec = steer_vec.to(device=value.device, dtype=value.dtype)
212 return value + steer_strength * vec.unsqueeze(0).unsqueeze(0)
213
214 hook_name = f"blocks.{layer}.hook_resid_post"
215 fmt_prompt = _format_prompt(model, prompt)
216 tokens = model.to_tokens(fmt_prompt)
217 steered_tokens: list[int] = []
218 cur = tokens
219 for _ in range(max_new_tokens):
220 with torch.no_grad():
221 out = model.run_with_hooks(cur, fwd_hooks=[(hook_name, _steer_hook)])
222 next_id = int(out[0, -1].argmax().item())
223 if next_id == model.tokenizer.eos_token_id:
224 break
225 steered_tokens.append(next_id)
226 cur = torch.cat([cur, torch.tensor([[next_id]], device=cur.device)], dim=1)
227 return model.tokenizer.decode(steered_tokens, skip_special_tokens=True)
228
229
231 *,
232 model_id: str,
233 prompt: str | None = None,
234 steer_strength: float,
235 layer: int | None = None,
236 feature_idx: int | None = None,
237 vector_path: str | Path | None = None,
238 vector_data: dict[str, Any] | None = None,
239 feature_label: str | None = None,
240 max_new_tokens: int = 80,
241 ctx: dict | None = None,
242 args: dict | None = None,
243) -> dict[str, Any]:
244 """Steer generation using a live feature index or a saved vector file."""
245 from aquin.compute.causal_trace import run_chat
246 from aquin.compute.feature_analysis import format_feature_ref, resolve_feature_label
247 from aquin.compute.model_loader import get_config, get_loaded_model, load_model, resolve_model_id
248
249 args = args or {}
250 do_eval = bool(args.get("eval"))
251 explicit_prompt = args.get("prompt")
252 if explicit_prompt is not None:
253 explicit_prompt = str(explicit_prompt).strip() or None
254
255 prompt_text = (prompt or "").strip() if prompt else ""
256 if not prompt_text and explicit_prompt:
257 prompt_text = explicit_prompt
258
259 # --eval without an explicit --prompt: skip single-demo generation; suite is the output.
260 run_single = bool(prompt_text) and (not do_eval or explicit_prompt is not None)
261 if do_eval and not explicit_prompt:
262 run_single = False
263 prompt_text = ""
264
265 model_id = resolve_model_id(model_id)
266 model = get_loaded_model()
267 if model is None:
268 model = load_model(model_id)
269
270 cfg = get_config(model_id)
271 resolved_layer = layer
272 steer_vec: torch.Tensor | None = None
273 vec_meta: dict[str, Any] | None = None
274
275 if vector_path or vector_data:
276 vec_meta = vector_data or load_steer_vector_file(str(vector_path))
277 steer_vec = steer_vector_tensor(
278 vec_meta,
279 device=model.W_E.device,
280 dtype=model.W_E.dtype,
281 )
282 if resolved_layer is None:
283 resolved_layer = int(vec_meta.get("layer", cfg["sae_layer"]))
284 if feature_idx is None:
285 feature_idx = int(vec_meta.get("feature_idx", 0))
286 if not feature_label:
287 feature_label = str(vec_meta.get("feature_label") or f"F{feature_idx}")
288 if steer_vec.shape[-1] != model.cfg.d_model:
289 return {
290 "error": (
291 f"Vector d_model={steer_vec.shape[-1]} does not match model d_model={model.cfg.d_model}. "
292 "Use a vector extracted from a compatible checkpoint."
293 ),
294 }
295 elif feature_idx is not None:
296 direction, resolved_layer, _ = resolve_feature_direction(
297 model_id, int(feature_idx), layer=layer,
298 )
299 steer_vec = direction.to(device=model.W_E.device, dtype=model.W_E.dtype)
300 if not feature_label:
301 # CLI paths (sweep/steer) pass no session ctx — keep a cheap display name.
302 # Causal re-labeling is for UI sessions that provide ctx.
303 if ctx:
304 feature_label = resolve_feature_label(
305 int(feature_idx),
306 ctx=ctx,
307 args={**args, "model_id": model_id},
308 layer=resolved_layer,
309 )
310 else:
311 feature_label = f"F{feature_idx}"
312 else:
313 return {"error": "Provide feature_idx or vector (path to saved LAT file)."}
314
315 resolved_layer = int(resolved_layer if resolved_layer is not None else cfg["sae_layer"])
316 assert steer_vec is not None
317
318 baseline = ""
319 steered_response = ""
320 if run_single:
321 try:
322 baseline = run_chat(
323 prompt_text, model_id=model_id, max_new_tokens=max_new_tokens, temperature=0.0,
324 )
325 except Exception as e:
326 return {"error": f"Baseline generation failed: {e}"}
327 try:
328 steered_response = _generate_steered_response(
329 model,
330 prompt_text,
331 steer_vec=steer_vec,
332 steer_strength=steer_strength,
333 layer=resolved_layer,
334 max_new_tokens=max_new_tokens,
335 )
336 except Exception as e:
337 steered_response = f"[steer error: {e}]"
338
339 out: dict[str, Any] = {
340 "feature_idx": int(feature_idx or 0),
341 "feature_label": feature_label or f"F{feature_idx}",
342 "feature_ref": format_feature_ref(int(feature_idx or 0), feature_label or f"F{feature_idx}"),
343 "steer_strength": steer_strength,
344 "layer": resolved_layer,
345 "prompt": prompt_text,
346 "baseline_response": baseline,
347 "steered_response": steered_response,
348 "demo": run_single,
349 }
350 if vec_meta:
351 out["vector_path"] = str(vector_path) if vector_path else None
352 out["vector_source_model_id"] = vec_meta.get("model_id")
353 out["vector_source_layer"] = vec_meta.get("layer")
354 out["vector_source_feature_idx"] = vec_meta.get("feature_idx")
355
356 if do_eval:
357 from aquin.compute.steer_eval import run_steer_probe_eval
358
359 eval_tokens = int(args.get("eval_max_new_tokens") or min(max_new_tokens, 64))
360 threshold = float(args.get("threshold") if args.get("threshold") is not None else 0.5)
361
362 def _baseline(p: str) -> str:
363 return run_chat(p, model_id=model_id, max_new_tokens=eval_tokens, temperature=0.0)
364
365 def _steered(p: str) -> str:
367 model,
368 p,
369 steer_vec=steer_vec,
370 steer_strength=steer_strength,
371 layer=resolved_layer,
372 max_new_tokens=eval_tokens,
373 )
374
375 eval_result = run_steer_probe_eval(
376 generate_baseline=_baseline,
377 generate_steered=_steered,
378 prompts=args.get("prompts") or args.get("prompts_path"),
379 reference_answers=args.get("reference_answers"),
380 threshold=threshold,
381 max_probes=int(args.get("max_probes") or 50),
382 )
383 if eval_result.get("error"):
384 return {"error": eval_result["error"], **{k: v for k, v in out.items() if k != "error"}}
385 out["eval"] = eval_result
386 # Seed demo fields from first probe when --eval ran without --prompt
387 if not run_single and eval_result.get("probes"):
388 first = eval_result["probes"][0]
389 out["prompt"] = first.get("prompt") or ""
390 out["baseline_response"] = (first.get("baseline") or {}).get("response") or ""
391 out["steered_response"] = (first.get("steered") or {}).get("response") or ""
392
393 return out
dict[str, Any]|None lookup_probe_record(str model_id, str|None probe_id=None)
torch.Tensor steer_vector_tensor(dict[str, Any] data, *, torch.device|str device, torch.dtype dtype)
dict[str, Any] extract_steer_vector(str model_id, int feature_idx, str|Path output, *, int|None layer=None, str|None feature_label=None, str|None probe_id=None)
tuple[torch.Tensor, int, dict[str, Any]] resolve_feature_direction(str model_id, int feature_idx, *, int|None layer=None, bool apply_norm_std=True)
dict[str, Any] build_steer_vector_payload(str model_id, int feature_idx, *, int|None layer=None, str|None feature_label=None, str|None probe_id=None)
dict[str, Any] run_steer_with_vector(*, str model_id, str|None prompt=None, float steer_strength, int|None layer=None, int|None feature_idx=None, str|Path|None vector_path=None, dict[str, Any]|None vector_data=None, str|None feature_label=None, int max_new_tokens=80, dict|None ctx=None, dict|None args=None)
Path write_steer_vector_file(dict[str, Any] payload, str|Path output)
Path _experiment_path(str model_id)
str _generate_steered_response(Any model, str prompt, *, torch.Tensor steer_vec, float steer_strength, int layer, int max_new_tokens)
dict[str, Any] load_steer_vector_file(str|Path path)