"""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 # tail tokens scored for the KL guard (keeps vocab-sized tensors small) 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): # orthonormalize (QR) so rank>1 learns genuinely distinct directions 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): # sum of the (orthonormal) basis directions, scaled 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) # hook: add the current (differentiable) delta to the residual stream at `layer` 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()