AQIT 0.1.0
Loading...
Searching...
No Matches
localize_collapse.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""Deception-representation collapse localization across layers."""
3from __future__ import annotations
4
5from pathlib import Path
6from typing import Any
7
8import torch
9
10from aquin.compute.find_feature import load_deception_probes
11
12SIGNAL_WEAK = 0.08
13SIGNAL_COLLAPSED = 0.03
14MAX_PROBES_PER_CLASS = 24
15
17def _unit(v: torch.Tensor) -> torch.Tensor:
18 return v / v.norm().clamp(min=1e-8)
19
20
21def _centroid_signal(honest: torch.Tensor, deceptive: torch.Tensor) -> dict[str, float]:
22 """honest/deceptive: (n, d). Relative centroid separation + cosine gap."""
23 if honest.shape[0] == 0 or deceptive.shape[0] == 0:
24 return {"signal": 0.0, "l2_sep": 0.0, "cos_sep": 0.0, "proj_delta": 0.0}
26 mu_h = honest.mean(dim=0)
27 mu_d = deceptive.mean(dim=0)
28 delta = mu_d - mu_h
29 l2_sep = float(delta.norm().item())
30 scale = float(mu_h.norm().item() + mu_d.norm().item()) + 1e-8
31 signal = l2_sep / scale
32
33 cos = float(torch.nn.functional.cosine_similarity(mu_h.unsqueeze(0), mu_d.unsqueeze(0)).item())
34 cos_sep = max(0.0, 1.0 - cos)
35
36 direction = _unit(delta)
37 proj_h = float((honest @ direction).mean().item())
38 proj_d = float((deceptive @ direction).mean().item())
39 proj_delta = abs(proj_d - proj_h)
40
41 return {
42 "signal": round(signal, 6),
43 "l2_sep": round(l2_sep, 6),
44 "cos_sep": round(cos_sep, 6),
45 "proj_delta": round(proj_delta, 6),
46 }
47
48
49def _status(signal: float) -> str:
50 if signal <= SIGNAL_COLLAPSED:
51 return "collapsed"
52 if signal <= SIGNAL_WEAK:
53 return "weak"
54 return "ok"
55
56
57def pick_collapse_layer(layers: list[dict[str, Any]]) -> dict[str, Any]:
58 """Peak = max signal; collapse = steepest drop after peak (else global min)."""
59 if not layers:
60 return {"peak_layer": None, "collapse_layer": None, "drop_from_peak": None}
62 peak = max(layers, key=lambda r: float(r.get("signal") or 0.0))
63 peak_layer = int(peak["layer"])
64 peak_signal = float(peak["signal"])
65
66 after = [r for r in layers if int(r["layer"]) > peak_layer]
67 if after:
68 worst = min(after, key=lambda r: float(r.get("signal") or 0.0))
69 collapse_layer = int(worst["layer"])
70 drop = round(peak_signal - float(worst["signal"]), 6)
71 else:
72 worst = min(layers, key=lambda r: float(r.get("signal") or 0.0))
73 collapse_layer = int(worst["layer"])
74 drop = round(peak_signal - float(worst["signal"]), 6)
75
76 return {
77 "peak_layer": peak_layer,
78 "collapse_layer": collapse_layer,
79 "drop_from_peak": drop,
80 "peak_signal": round(peak_signal, 6),
81 "collapse_signal": round(float(worst["signal"]), 6),
82 }
83
84
86 honest_by_layer: dict[int, torch.Tensor],
87 deceptive_by_layer: dict[int, torch.Tensor],
88) -> list[dict[str, Any]]:
89 layers_out: list[dict[str, Any]] = []
90 for layer in sorted(set(honest_by_layer) & set(deceptive_by_layer)):
91 metrics = _centroid_signal(honest_by_layer[layer], deceptive_by_layer[layer])
92 layers_out.append({
93 "layer": int(layer),
94 **metrics,
95 "status": _status(metrics["signal"]),
96 "n_honest": int(honest_by_layer[layer].shape[0]),
97 "n_deceptive": int(deceptive_by_layer[layer].shape[0]),
98 })
99 return layers_out
100
101
103 baseline_layers: list[dict[str, Any]],
104 stressor_layers: list[dict[str, Any]],
105) -> list[dict[str, Any]]:
106 by_base = {int(r["layer"]): r for r in baseline_layers}
107 by_stress = {int(r["layer"]): r for r in stressor_layers}
108 rows: list[dict[str, Any]] = []
109 for layer in sorted(set(by_base) & set(by_stress)):
110 b = float(by_base[layer]["signal"])
111 s = float(by_stress[layer]["signal"])
112 delta = round(b - s, 6)
113 rows.append({
114 "layer": layer,
115 "baseline_signal": round(b, 6),
116 "stressor_signal": round(s, 6),
117 "collapse_delta": delta,
118 "status": "collapsed" if delta >= 0.05 and s <= SIGNAL_WEAK else (
119 "weakened" if delta >= 0.02 else "stable"
120 ),
121 })
122 return rows
123
124
126 honest: torch.Tensor,
127 deceptive: torch.Tensor,
128 direction: torch.Tensor,
129) -> dict[str, float]:
130 d = _unit(direction.float().cpu())
131 if honest.shape[-1] != d.shape[-1]:
132 raise ValueError(
133 f"Direction d_model={d.shape[-1]} does not match activations d_model={honest.shape[-1]}"
134 )
135 proj_h = float((honest @ d).mean().item())
136 proj_d = float((deceptive @ d).mean().item())
137 return {
138 "honest_proj": round(proj_h, 6),
139 "deceptive_proj": round(proj_d, 6),
140 "proj_delta": round(abs(proj_d - proj_h), 6),
141 }
142
143
145 model: Any,
146 model_id: str,
147 layer: int,
148 honest_texts: list[str],
149 deceptive_texts: list[str],
150 *,
151 top_k: int = 8,
152) -> list[dict[str, Any]] | None:
153 """Rank SAE features by deceptive−honest mean activation at one layer."""
154 try:
155 from aquin.compute.feature_analysis import load_sae
156 from aquin.compute.find_feature import _rank_features
157 except Exception:
158 return None
159
160 try:
161 sae = load_sae(model_id, layer)
162 except Exception:
163 return None
164
165 def _mean_feats(texts: list[str]) -> torch.Tensor | None:
166 feats: list[torch.Tensor] = []
167 hook = f"blocks.{layer}.hook_resid_post"
168 for text in texts:
169 tokens = model.to_tokens(text)
170 with torch.no_grad():
171 _, cache = model.run_with_cache(tokens, names_filter=lambda n, h=hook: n == h)
172 if hook not in cache:
173 continue
174 act = cache[hook][0, -1].float()
175 with torch.no_grad():
176 f = sae.encode(act.unsqueeze(0))[0].float().cpu()
177 feats.append(f)
178 if not feats:
179 return None
180 return torch.stack(feats).mean(dim=0)
181
182 h = _mean_feats(honest_texts)
183 d = _mean_feats(deceptive_texts)
184 if h is None or d is None:
185 return None
186 return _rank_features(h, d, top_k=top_k, direction="both")
187
188
190 model: Any,
191 model_id: str,
192 *,
193 prompts: str | Path | None = None,
194 stressor_prompts: str | Path | None = None,
195 feature_idx: int | None = None,
196 vector_path: str | Path | None = None,
197 layer: int | None = None,
198 top_k_features: int = 8,
199 collect_layer_activations: Any | None = None,
200) -> dict[str, Any]:
201 """
202 Rank layers by honest vs deceptive representation strength.
203
204 Default: contrastive centroid signal per resid_post layer.
205 Optional stressor prompts → per-layer collapse_delta.
206 Optional feature_idx / vector → direction projection at the relevant layer.
207 """
208 from aquin.compute.layer_analysis import _collect_layer_activations
209
210 collect = collect_layer_activations or _collect_layer_activations
211
212 try:
213 honest, deceptive, meta = load_deception_probes(prompts)
214 except (FileNotFoundError, ValueError) as e:
215 return {"error": str(e)}
216
217 honest = honest[:MAX_PROBES_PER_CLASS]
218 deceptive = deceptive[:MAX_PROBES_PER_CLASS]
219
220 honest_acts = collect(model, honest)
221 deceptive_acts = collect(model, deceptive)
222 layers = score_layer_activations(honest_acts, deceptive_acts)
223 locus = pick_collapse_layer(layers)
224
225 out: dict[str, Any] = {
226 "mode": "contrastive",
227 "prompts_path": meta.get("prompts_path"),
228 "n_honest": len(honest),
229 "n_deceptive": len(deceptive),
230 "layers": layers,
231 **locus,
232 "n_collapsed": sum(1 for r in layers if r["status"] == "collapsed"),
233 "n_weak": sum(1 for r in layers if r["status"] == "weak"),
234 }
235
236 if stressor_prompts is not None:
237 try:
238 s_honest, s_deceptive, s_meta = load_deception_probes(stressor_prompts)
239 except (FileNotFoundError, ValueError) as e:
240 return {**out, "error": f"stressor probes: {e}"}
241 s_honest = s_honest[:MAX_PROBES_PER_CLASS]
242 s_deceptive = s_deceptive[:MAX_PROBES_PER_CLASS]
243 s_layers = score_layer_activations(
244 collect(model, s_honest),
245 collect(model, s_deceptive),
246 )
247 stress_rows = score_stressor_collapse(layers, s_layers)
248 out["stressor"] = {
249 "prompts_path": s_meta.get("prompts_path"),
250 "n_honest": len(s_honest),
251 "n_deceptive": len(s_deceptive),
252 "layers": stress_rows,
253 }
254 if stress_rows:
255 worst = max(stress_rows, key=lambda r: float(r["collapse_delta"]))
256 out["stressor"]["max_collapse_layer"] = int(worst["layer"])
257 out["stressor"]["max_collapse_delta"] = float(worst["collapse_delta"])
258 # Prefer stressor-driven collapse locus when available
259 out["collapse_layer"] = int(worst["layer"])
260 out["collapse_signal"] = float(worst["stressor_signal"])
261 out["drop_from_peak"] = float(worst["collapse_delta"])
262 out["mode"] = "contrastive+stressor"
263
264 direction_info: dict[str, Any] | None = None
265 direction_vec: torch.Tensor | None = None
266 direction_layer: int | None = layer
267
268 if vector_path:
269 from aquin.compute.steer_vector import load_steer_vector_file, steer_vector_tensor
270
271 try:
272 payload = load_steer_vector_file(vector_path)
273 direction_vec = steer_vector_tensor(
274 payload, device="cpu", dtype=torch.float32,
275 )
276 direction_layer = int(payload.get("layer", direction_layer or 0))
277 direction_info = {
278 "source": "vector",
279 "vector_path": str(vector_path),
280 "layer": direction_layer,
281 "feature_idx": payload.get("feature_idx"),
282 }
283 except (OSError, ValueError) as e:
284 return {**out, "error": f"vector: {e}"}
285 elif feature_idx is not None:
286 from aquin.compute.steer_vector import resolve_feature_direction
287
288 try:
289 direction_vec, direction_layer, _ = resolve_feature_direction(
290 model_id, int(feature_idx), layer=layer,
291 )
292 direction_vec = direction_vec.float().cpu()
293 direction_info = {
294 "source": "feature",
295 "feature_idx": int(feature_idx),
296 "layer": int(direction_layer),
297 }
298 except Exception as e:
299 return {**out, "error": f"feature direction: {e}"}
300
301 if direction_vec is not None and direction_layer is not None:
302 h = honest_acts.get(int(direction_layer))
303 d = deceptive_acts.get(int(direction_layer))
304 if h is not None and d is not None:
305 try:
306 proj = _project_on_direction(h, d, direction_vec)
307 direction_info = {**(direction_info or {}), **proj}
308 except ValueError as e:
309 return {**out, "error": str(e)}
310 out["direction"] = direction_info
311
312 collapse_l = out.get("collapse_layer")
313 if collapse_l is not None:
315 model,
316 model_id,
317 int(collapse_l),
318 honest,
319 deceptive,
320 top_k=int(top_k_features),
321 )
322 if feats is not None:
323 out["features_at_collapse"] = feats
324 out["features_layer"] = int(collapse_l)
325
326 return out
list[dict[str, Any]] score_layer_activations(dict[int, torch.Tensor] honest_by_layer, dict[int, torch.Tensor] deceptive_by_layer)
dict[str, float] _centroid_signal(torch.Tensor honest, torch.Tensor deceptive)
list[dict[str, Any]]|None rank_sae_features_at_layer(Any model, str model_id, int layer, list[str] honest_texts, list[str] deceptive_texts, *, int top_k=8)
dict[str, Any] pick_collapse_layer(list[dict[str, Any]] layers)
list[dict[str, Any]] score_stressor_collapse(list[dict[str, Any]] baseline_layers, list[dict[str, Any]] stressor_layers)
torch.Tensor _unit(torch.Tensor v)
dict[str, Any] run_localize_collapse(Any model, str model_id, *, str|Path|None prompts=None, str|Path|None stressor_prompts=None, int|None feature_idx=None, str|Path|None vector_path=None, int|None layer=None, int top_k_features=8, Any|None collect_layer_activations=None)
dict[str, float] _project_on_direction(torch.Tensor honest, torch.Tensor deceptive, torch.Tensor direction)