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."""
52 cfg = get_config(model_id)
53 resolved_layer = int(layer
if layer
is not None else cfg[
"sae_layer"])
55 sae = load_sae(model_id, resolved_layer)
56 except FileNotFoundError:
59 available = get_available_sae_layers(model_id)
60 fallback = int(get_sae_layer(model_id))
62 resolved_layer = fallback
if fallback
in available
else int(available[0])
63 sae = load_sae(model_id, resolved_layer)
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})")
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:
80 std = norm[
"std"].detach().float().to(device=feat_dir.device)
81 feat_dir = feat_dir * std
82 mean = norm.get(
"mean")
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,
90 "d_model": int(feat_dir.shape[-1]),
91 "vector_l2_norm": round(float(feat_dir.norm().item()), 6),
94 return feat_dir, resolved_layer, meta
101 layer: int |
None =
None,
102 feature_label: str |
None =
None,
103 probe_id: str |
None =
None,
106 model_id, feature_idx, layer=layer,
109 label = feature_label
or f
"F{feature_idx}"
112 "schema_version": SCHEMA_VERSION,
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(),
125 "source":
"steer --save",
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")
203 steer_vec: torch.Tensor,
204 steer_strength: float,
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)
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] = []
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:
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)
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,
244 """Steer generation using a live feature index or a saved vector file."""
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
255 prompt_text = (prompt
or "").strip()
if prompt
else ""
256 if not prompt_text
and explicit_prompt:
257 prompt_text = explicit_prompt
260 run_single = bool(prompt_text)
and (
not do_eval
or explicit_prompt
is not None)
261 if do_eval
and not explicit_prompt:
265 model_id = resolve_model_id(model_id)
266 model = get_loaded_model()
268 model = load_model(model_id)
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
275 if vector_path
or vector_data:
279 device=model.W_E.device,
280 dtype=model.W_E.dtype,
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:
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."
295 elif feature_idx
is not None:
297 model_id, int(feature_idx), layer=layer,
299 steer_vec = direction.to(device=model.W_E.device, dtype=model.W_E.dtype)
300 if not feature_label:
304 feature_label = resolve_feature_label(
307 args={**args,
"model_id": model_id},
308 layer=resolved_layer,
311 feature_label = f
"F{feature_idx}"
313 return {
"error":
"Provide feature_idx or vector (path to saved LAT file)."}
315 resolved_layer = int(resolved_layer
if resolved_layer
is not None else cfg[
"sae_layer"])
316 assert steer_vec
is not None
319 steered_response =
""
323 prompt_text, model_id=model_id, max_new_tokens=max_new_tokens, temperature=0.0,
325 except Exception
as e:
326 return {
"error": f
"Baseline generation failed: {e}"}
332 steer_strength=steer_strength,
333 layer=resolved_layer,
334 max_new_tokens=max_new_tokens,
336 except Exception
as e:
337 steered_response = f
"[steer error: {e}]"
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,
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")
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)
362 def _baseline(p: str) -> str:
363 return run_chat(p, model_id=model_id, max_new_tokens=eval_tokens, temperature=0.0)
365 def _steered(p: str) -> str:
370 steer_strength=steer_strength,
371 layer=resolved_layer,
372 max_new_tokens=eval_tokens,
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"),
381 max_probes=int(args.get(
"max_probes")
or 50),
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
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 ""
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)