| """Learn a steering vector (or k-dim subspace) by gradient descent instead of difference-of-means. |
| |
| Difference-of-means is a closed form that DESCRIBES how activations differ; it never optimizes |
| behavior. That is why CV-AUC can be 1.000 on six layers at once while their behavioral effects |
| differ. Here `d` is a trainable parameter optimized for the thing we actually want: |
| |
| loss = -log P(pos_response | prompt, +alpha*d) # +d should produce positive-pole behavior |
| -log P(neg_response | prompt, -alpha*d) # -d should produce negative-pole behavior |
| + kl_coef * KL(steered || base) on neutral text # keep the model coherent |
| |
| The bidirectional term is what makes the direction an axis rather than a generic "be better" push, |
| and the KL term is what stops the alpha>=4 incoherence we measured with diff-of-means vectors. |
| |
| With --rank k>1 it learns an orthonormal subspace (k directions, orthogonality enforced by QR), |
| which difference-of-means cannot express at all. |
| |
| Saves the same blob format the vLLM patch and the HF server read: |
| {trait, layer, dir, rank, cv_auc(None), norm, learned: True} |
| Rank>1 additionally saves `basis` [k, D]. |
| |
| Usage: |
| .venv/bin/python scripts/16_learned/train_steer_vec.py --trait verification-before-claiming \ |
| --layer 16 --alpha 2.0 --rank 1 --epochs 3 |
| """ |
| from __future__ import annotations |
| import argparse, json, math, sys, time |
| from pathlib import Path |
|
|
| import torch |
| import torch.nn.functional as F |
| from transformers import AutoTokenizer, AutoModelForCausalLM |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) |
| from common import DATA, VECTORS, MODEL_PATH |
|
|
| MAX_PROMPT = 1024 |
| MAX_RESP = 384 |
| KL_WINDOW = 128 |
|
|
|
|
| def build_examples(trait, tok): |
| """Return [(input_ids, resp_start, pole)] from the judged pool.""" |
| rows = [json.loads(l) for l in open(DATA / "traits" / trait / "responses" / "scored.jsonl")] |
| rows = [r for r in rows if r.get("coherence") and r["coherence"] >= 60 and r.get("response")] |
| out = [] |
| for r in rows: |
| msgs = [{"role": "system", "content": r.get("system", "")}, |
| {"role": "user", "content": r["question"]}] |
| prompt = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True) |
| p_ids = tok(prompt, add_special_tokens=False).input_ids[-MAX_PROMPT:] |
| r_ids = tok(r["response"], add_special_tokens=False).input_ids[:MAX_RESP] |
| if len(r_ids) < 8: |
| continue |
| out.append((p_ids + r_ids, len(p_ids), r["pole"])) |
| return out |
|
|
|
|
| class Steer(torch.nn.Module): |
| """Trainable direction / subspace applied additively at one layer.""" |
|
|
| def __init__(self, d_model, rank, init=None): |
| super().__init__() |
| w = torch.randn(rank, d_model) * 0.02 if init is None else init.clone().reshape(rank, -1) |
| self.w = torch.nn.Parameter(w) |
|
|
| def basis(self): |
| |
| if self.w.shape[0] == 1: |
| return self.w / (self.w.norm() + 1e-9) |
| q, _ = torch.linalg.qr(self.w.float().T) |
| return q.T.to(self.w.dtype) |
|
|
| def delta(self, sign, alpha): |
| |
| return sign * alpha * self.basis().sum(0) |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--trait", required=True) |
| ap.add_argument("--layer", type=int, required=True) |
| ap.add_argument("--alpha", type=float, default=2.0) |
| ap.add_argument("--rank", type=int, default=1) |
| ap.add_argument("--epochs", type=int, default=3) |
| ap.add_argument("--lr", type=float, default=5e-3) |
| ap.add_argument("--kl-coef", type=float, default=1.0) |
| ap.add_argument("--init-from-meandiff", action="store_true", |
| help="initialize from the diff-of-means vector if present") |
| ap.add_argument("--model-path", default=MODEL_PATH) |
| ap.add_argument("--out", default=None) |
| args = ap.parse_args() |
|
|
| tok = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True) |
| model = AutoModelForCausalLM.from_pretrained( |
| args.model_path, dtype=torch.bfloat16, device_map={"": 0}, |
| trust_remote_code=True, low_cpu_mem_usage=True) |
| model.eval() |
| for p in model.parameters(): |
| p.requires_grad_(False) |
| model.gradient_checkpointing_enable() |
|
|
| d_model = model.config.hidden_size |
| init = None |
| if args.init_from_meandiff: |
| cand = VECTORS / f"{args.trait}_L{args.layer}.pt" |
| if cand.exists(): |
| init = torch.load(cand, weights_only=False)["dir"].float().to("cuda") |
| print(f"[init] from {cand}", flush=True) |
| steer = Steer(d_model, args.rank, init).to("cuda", torch.float32) |
| opt = torch.optim.Adam(steer.parameters(), lr=args.lr) |
|
|
| |
| ctl = {"sign": 0.0} |
|
|
| def pre_hook(_m, a, kw): |
| if ctl["sign"] == 0.0: |
| return None |
| h = a[0] if a else kw["hidden_states"] |
| h = h + steer.delta(ctl["sign"], args.alpha).to(h.dtype) |
| if a: |
| return ((h,) + a[1:], kw) |
| kw["hidden_states"] = h |
| return (a, kw) |
|
|
| handle = model.model.layers[args.layer].register_forward_pre_hook(pre_hook, with_kwargs=True) |
|
|
| examples = build_examples(args.trait, tok) |
| pos = [e for e in examples if e[2] == "pos"] |
| neg = [e for e in examples if e[2] == "neg"] |
| print(f"[data] pos={len(pos)} neg={len(neg)} layer={args.layer} alpha={args.alpha} rank={args.rank}", flush=True) |
| if not pos or not neg: |
| print("need both poles; abort", flush=True); return |
|
|
| def seq_loss(ids, resp_start, sign): |
| ctl["sign"] = sign |
| x = torch.tensor([ids], device="cuda") |
| out = model(input_ids=x, use_cache=False).logits[:, :-1] |
| tgt = x[:, 1:] |
| lp = F.cross_entropy(out.float().reshape(-1, out.shape[-1]), tgt.reshape(-1), reduction="none") |
| mask = torch.zeros_like(tgt, dtype=torch.float32) |
| mask[:, resp_start - 1:] = 1.0 |
| return (lp.reshape(tgt.shape) * mask).sum() / mask.sum().clamp(min=1) |
|
|
| def kl_term(ids): |
| """KL(steered || base) — coherence guard. |
| |
| Scored on a short tail window only: full-length log_softmax over a 151k vocab |
| allocates hundreds of MB per tensor and OOMs next to a running vLLM. A tail |
| window is a sufficient regularizer and keeps memory flat. |
| """ |
| x = torch.tensor([ids[: min(len(ids), MAX_PROMPT)][-KL_WINDOW:]], device="cuda") |
| ctl["sign"] = 0.0 |
| with torch.no_grad(): |
| base = model(input_ids=x, use_cache=False).logits[:, -KL_WINDOW:].float().log_softmax(-1) |
| ctl["sign"] = 1.0 |
| stee = model(input_ids=x, use_cache=False).logits[:, -KL_WINDOW:].float().log_softmax(-1) |
| kl = F.kl_div(stee, base, log_target=True, reduction="batchmean") |
| del base, stee |
| return kl |
|
|
| n_steps = min(len(pos), len(neg)) |
| for ep in range(args.epochs): |
| t0, tot = time.time(), 0.0 |
| for i in range(n_steps): |
| p_ids, p_start, _ = pos[i % len(pos)] |
| n_ids, n_start, _ = neg[i % len(neg)] |
| loss = seq_loss(p_ids, p_start, +1.0) + seq_loss(n_ids, n_start, -1.0) |
| if args.kl_coef > 0: |
| loss = loss + args.kl_coef * kl_term(p_ids) |
| opt.zero_grad(); loss.backward(); opt.step() |
| tot += float(loss.item()) |
| if (i + 1) % 10 == 0: |
| print(f" ep{ep} {i+1}/{n_steps} loss={tot/(i+1):.4f} ({time.time()-t0:.0f}s)", flush=True) |
| print(f"[epoch {ep}] mean loss={tot/max(n_steps,1):.4f}", flush=True) |
|
|
| handle.remove() |
| basis = steer.basis().detach().float().cpu() |
| d = basis.sum(0) |
| d = d / (d.norm() + 1e-9) |
| out = Path(args.out) if args.out else VECTORS / f"{args.trait}_L{args.layer}_learned_r{args.rank}.pt" |
| blob = {"trait": args.trait, "layer": args.layer, "dir": d, "rank": args.rank, |
| "learned": True, "alpha_trained": args.alpha, "norm": 1.0} |
| if args.rank > 1: |
| blob["basis"] = basis |
| torch.save(blob, out) |
| print(f"saved -> {out}", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|