AQIT 0.1.0
Loading...
Searching...
No Matches
steer.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6"""
7Copied from inspection-backend/server.py steer logic.
8Synchronous generator — yields SSE event dicts.
9"""
10from __future__ import annotations
11
12from typing import Generator
13
14import torch
15
16from aquin.compute.device import resolve_compute_device
17
18DEVICE = resolve_compute_device()
19
20
22 prompt: str,
23 feature_idx: int,
24 strength: float = 20.0,
25 model_id: str = "llama-3.2-1b",
26 max_new_tokens: int = 200,
27 temperature: float = 0.7,
28) -> Generator[dict, None, None]:
29 from aquin.compute.model_loader import get_config, load_model, resolve_model_id
30 from aquin.compute.feature_analysis import _load_sae, _load_norm
31 from aquin.compute.causal_trace import _format_prompt
32
33 model_id = resolve_model_id(model_id)
34 m = load_model(model_id)
35 cfg = get_config(model_id)
36 layer = cfg.get("sae_layer", 8)
37
38 sae = _load_sae(model_id, layer)
39 norm = _load_norm(model_id, layer)
40
41 feat_dir = sae.W_dec[feature_idx]
42 if norm is not None and norm.get("std") is not None:
43 feat_dir = feat_dir * norm["std"].to(device=feat_dir.device, dtype=feat_dir.dtype)
44 steer_vec = feat_dir.to(device=m.W_E.device, dtype=m.W_E.dtype)
45 hook_name = f"blocks.{layer}.hook_resid_post"
46
47 formatted = _format_prompt(m, prompt)
48 input_ids = m.tokenizer(formatted, return_tensors="pt").input_ids.to(m.W_E.device)
49
50 # Baseline pass
51 with torch.no_grad():
52 cur = input_ids
53 for _ in range(max_new_tokens):
54 logits = m(cur)
55 next_logits = logits[0, -1] / max(temperature, 1e-6)
56 probs = torch.softmax(next_logits, dim=-1)
57 next_id = int(torch.multinomial(probs, 1).item())
58 if next_id == m.tokenizer.eos_token_id:
59 break
60 tok_str = m.tokenizer.decode([next_id], skip_special_tokens=True)
61 yield {"type": "baseline", "token": tok_str}
62 cur = torch.cat([cur, torch.tensor([[next_id]], device=cur.device)], dim=1)
63 yield {"type": "baseline_done"}
64
65 def steer_hook(value, hook):
66 value = value.clone()
67 vec = steer_vec.to(device=value.device, dtype=value.dtype)
68 value[:, -1, :] = value[:, -1, :] + strength * vec
69 return value
70
71 # Steered pass
72 with torch.no_grad():
73 cur = input_ids
74 for _ in range(max_new_tokens):
75 logits = m.run_with_hooks(cur, fwd_hooks=[(hook_name, steer_hook)])
76 next_logits = logits[0, -1] / max(temperature, 1e-6)
77 probs = torch.softmax(next_logits, dim=-1)
78 next_id = int(torch.multinomial(probs, 1).item())
79 if next_id == m.tokenizer.eos_token_id:
80 break
81 tok_str = m.tokenizer.decode([next_id], skip_special_tokens=True)
82 yield {"type": "steered", "token": tok_str}
83 cur = torch.cat([cur, torch.tensor([[next_id]], device=cur.device)], dim=1)
84
85 yield {"type": "done"}
Generator[dict, None, None] run_steer_stream(str prompt, int feature_idx, float strength=20.0, str model_id="llama-3.2-1b", int max_new_tokens=200, float temperature=0.7)
Definition steer.py:32