capabvector-research / scripts /16_learned /train_steer_vec.py
AlexWortega's picture
Upload folder using huggingface_hub
aaf1c39 verified
Raw
History Blame Contribute Delete
8.27 kB
"""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()