"""Stage 6: Dr. GRPO RL — vllm-lens rollouts, no KL, no /std, global-token normalizer. Rollouts: ONE llm.generate() per step; per-request SteeringVector(norm_match=True) == our norm-matched inject@INJECT_LAYER at the marker. old_logp comes from vLLM's generation logprobs (valid behavior-policy logps at temperature 1.0 ONLY). new_logp is recomputed HF-side with the same inject hook; TIS (ratio capped at cfg.tis_cap, upper only) absorbs the residual vLLM/HF kernel mismatch; the LoRA-merged actor is pushed back into vLLM every --sync-every steps. Reward: each generation re-tokenized STANDALONE through the CLEAN base model (adapter disabled, no injection); reward = max over kept positions of x_t · unit(v) at READ_LAYER, position 0 skipped (attention-sink guard). No μ-centering: v is shared within a group, so μ·v is a constant that cancels exactly in the Dr. GRPO advantage (r − group_mean). python scripts/rl.py --tp 8 # full box (sbatch_rl.sh) python scripts/rl.py --groups-per-step 8 --group-size 4 --total-steps 3 --no-wandb # 1-GPU smoke """ import argparse import functools import json import math import os import time from collections import defaultdict os.environ.setdefault("VLLM_ALLOW_INSECURE_SERIALIZATION", "1") # pickle for apply_model(partial) import numpy as np import torch from peft import LoraConfig, PeftModel, get_peft_model from transformers import AutoModelForCausalLM, AutoTokenizer import wandb from mxf.config import D_MODEL, INJECT_LAYER, MODEL, READ_LAYER, STEER_COEFF, RLConfig, TrainConfig from mxf.inject import get_layer, hooked, make_inject_hook, read_resid from mxf.prompts import build_prompt_ids def _load_chunk(model, chunk): """Module-level (picklable) target for llm.apply_model — runs on every TP worker.""" model.load_weights(iter(chunk)) def sync_weights(actor, llm): """LoRA→vLLM colocate sync (TRL pattern): merge adapter, push HF-name/cpu-tensor pairs in per-layer chunks (msgspec caps one encode at 4GB), reset prefix cache, unmerge.""" t0 = time.time() actor.merge_adapter() try: buckets = defaultdict(list) for k, v in actor.state_dict().items(): if "lora_" in k or "modules_to_save" in k: continue k = k.removeprefix("base_model.model.") k = k.replace(".base_layer.weight", ".weight").replace(".base_layer.bias", ".bias") grp = f"layer_{int(k.split('.', 3)[2]):03d}" if k.startswith("model.layers.") else "_other" buckets[grp].append((k, v.detach().cpu())) for name in sorted(buckets): # "_other" (embed/norm/lm_head) first, then layers in order llm.apply_model(functools.partial(_load_chunk, chunk=buckets[name])) try: llm.llm_engine.reset_prefix_cache() # weights changed → cached prefixes are stale except AttributeError: pass # TODO(verify): vLLM 0.19 exposes reset_prefix_cache on llm_engine (0.19 should) finally: actor.unmerge_adapter() return time.time() - t0 @torch.no_grad() def rollout(llm, prompt_ids, marker, dirs, a): """B groups × G rollouts in ONE generate(). dirs: [B, d]. Returns flat group-major lists (texts, gen_ids, old_logps) — rollout i belongs to group i // group_size.""" from vllm import SamplingParams from vllm_lens import SteeringVector reqs, params = [], [] for v in dirs: # activations MUST be 3-D [1 layer, 1 pos, d]: a 2-D tensor hits vllm-lens's broadcast # branch and gets ADDed at EVERY token, silently ignoring position_indices. sv = SteeringVector(activations=v.view(1, 1, -1).cpu().float(), layer_indices=[INJECT_LAYER], scale=STEER_COEFF, norm_match=True, position_indices=[marker]) for _ in range(a.group_size): # TODO(verify): TokensPrompt dict form on vLLM 0.19 — reference passed text prompts; # we pass the exact chat-template ids so marker position is guaranteed. reqs.append({"prompt_token_ids": list(prompt_ids)}) params.append(SamplingParams(temperature=a.temperature, top_p=1.0, top_k=-1, logprobs=1, max_tokens=a.max_new_tokens, min_tokens=a.min_new_tokens, extra_args={"apply_steering_vectors": [sv]})) # TODO(verify): vLLM 0.19 reads Qwen3's generation_config for EOS (<|im_end|>) by default; # if smoke rollouts never stop early, pass stop_token_ids explicitly in SamplingParams. outs = llm.generate(reqs, params) assert len(outs) == len(reqs) texts, gen_ids, old_lps = [], [], [] for out in outs: o = out.outputs[0] ids = list(o.token_ids) # old_logp MUST come from vLLM (the behavior policy). Crash on absence — any substituted # value silently corrupts the importance ratio. assert o.logprobs is not None and len(o.logprobs) == len(ids), ( f"vLLM logprobs missing/short ({None if o.logprobs is None else len(o.logprobs)} vs " f"{len(ids)} tokens) — vLLM API drift?") lp = [] for t, tid in enumerate(ids): assert tid in o.logprobs[t], f"sampled token {tid} absent from logprobs at step {t}" lp.append(o.logprobs[t][tid].logprob) texts.append(o.text) gen_ids.append(ids) old_lps.append(torch.tensor(lp, dtype=torch.float32)) return texts, gen_ids, old_lps @torch.no_grad() def score(texts, dirs_rep, actor, tok, device, a): """reward[i] = max_t x_t·unit(v_i) at READ_LAYER — standalone re-tokenization, CLEAN base (adapter off, no injection), position 0 skipped. Rows with no kept token score 0.""" r = torch.zeros(len(texts)) valid = [i for i, t in enumerate(texts) if t.strip()] prev = tok.padding_side tok.padding_side = "right" # position 0 must be the first real token try: for s in range(0, len(valid), a.score_batch): idxs = valid[s : s + a.score_batch] enc = tok([texts[i] for i in idxs], return_tensors="pt", padding=True, truncation=True, max_length=a.max_new_tokens + 32, add_special_tokens=True).to(device) with actor.disable_adapter(): h, mask = read_resid(actor, READ_LAYER, dict(enc), pool="all") # [b,T,d] fp32, [b,T] keep = mask.clone() keep[:, 0] = False # attention-sink guard (old repo also norm-filtered; keep it simple) proj = torch.einsum("btd,bd->bt", h, dirs_rep[idxs]) best = proj.masked_fill(~keep, torch.finfo(proj.dtype).min).max(1).values has = keep.any(1) for row, i in enumerate(idxs): if has[row]: r[i] = best[row].item() finally: tok.padding_side = prev return r @torch.no_grad() def fluency(texts, actor, tok, device, a): """(mean clean-base logprob/token, distinct-token fraction) per standalone text — gate inputs. Adapter disabled so the policy can't inflate its own fluency score.""" logp, dis = torch.full((len(texts),), -20.0), torch.zeros(len(texts)) valid = [i for i, t in enumerate(texts) if t.strip()] prev = tok.padding_side tok.padding_side = "right" try: for s in range(0, len(valid), a.score_batch): idxs = valid[s : s + a.score_batch] enc = tok([texts[i] for i in idxs], return_tensors="pt", padding=True, truncation=True, max_length=a.max_new_tokens + 32, add_special_tokens=True).to(device) if enc["input_ids"].shape[1] < 2: continue with actor.disable_adapter(): logits = actor(**enc).logits[:, :-1].float() lp = torch.log_softmax(logits, -1).gather(-1, enc["input_ids"][:, 1:, None]).squeeze(-1) m = enc["attention_mask"][:, 1:].bool() for row, i in enumerate(idxs): n = int(m[row].sum()) if n: logp[i] = (lp[row][m[row]].sum() / n).item() ids = enc["input_ids"][row][enc["attention_mask"][row].bool()] dis[i] = len(set(ids.tolist())) / max(len(ids), 1) finally: tok.padding_side = prev return logp, dis def update(actor, opt, submodule, ids, attn, p_len, marker, old_lp, adv, dirs_rep, a, device): """ONE Dr. GRPO optimizer update. loss = Σ_tokens −min(ratio·A, clip(ratio)·A)·mask / TOTAL completion tokens in batch (GLOBAL constant normalizer — no per-sequence mean, no /std, no KL). ratio TIS-capped (upper only). new_logp forward runs with the SAME inject hook as rollout.""" n = ids.shape[0] gen_mask = attn[:, p_len:].bool() total_tok = max(int(gen_mask.sum()), 1) lo, hi = 1 - a.clip_eps, 1 + a.clip_eps loss_sum, clipped_tok, ent_sum = 0.0, 0, 0.0 opt.zero_grad(set_to_none=True) for s in range(0, n, a.micro_batch): e = min(s + a.micro_batch, n) b_ids, b_attn = ids[s:e].to(device), attn[s:e].to(device) hook = make_inject_hook([dirs_rep[i : i + 1] for i in range(s, e)], [[marker]] * (e - s), STEER_COEFF, device, torch.bfloat16) with hooked(submodule, hook): logits = actor(input_ids=b_ids, attention_mask=b_attn).logits[:, p_len - 1 : -1] logp_full = torch.log_softmax(logits.float(), -1) del logits new_lp = logp_full.gather(-1, b_ids[:, p_len:, None]).squeeze(-1) m = gen_mask[s:e].to(device) ratio = torch.exp(new_lp - old_lp[s:e].to(device)).clamp(max=a.tis_cap) # TIS, upper only A = adv[s:e, None].to(device) loss = (-torch.minimum(ratio * A, ratio.clamp(lo, hi) * A) * m).sum() / total_tok if a.entropy_coef > 0: # true per-token entropy (unbiased, logits are already here) — maximize r + β·H(π): # keeps the policy stochastic for Bo-N without KL's behavior-anchoring side effect ent = -(logp_full.exp() * logp_full).sum(-1) ent_sum += float((ent.detach() * m).sum()) loss = loss - a.entropy_coef * (ent * m).sum() / total_tok del logp_full loss.backward() # micro-losses share the global normalizer → grads sum correctly loss_sum += loss.item() clipped_tok += int((((ratio < lo) | (ratio > hi)) & m).sum()) gn = float(torch.nn.utils.clip_grad_norm_( [p for p in actor.parameters() if p.requires_grad], a.max_grad_norm)) if math.isfinite(gn): opt.step() else: # stepping Adam on nan/inf grads corrupts moments AND weights opt.zero_grad(set_to_none=True) print(f"[update] non-finite grad norm ({gn}) — skipping step", flush=True) return {"loss": loss_sum, "grad_norm": gn, "clipfrac": clipped_tok / total_tok, "entropy": ent_sum / total_tok} def main(): cfg, tr = RLConfig(), TrainConfig() ap = argparse.ArgumentParser() ap.add_argument("--data-dir", default="data/pretrain") ap.add_argument("--init-adapter", default=cfg.init_adapter) ap.add_argument("--save-dir", default=cfg.save_dir) ap.add_argument("--run-name", default=cfg.run_name) ap.add_argument("--direction-source", default=cfg.direction_source) ap.add_argument("--groups-per-step", type=int, default=cfg.groups_per_step) ap.add_argument("--group-size", type=int, default=cfg.group_size) ap.add_argument("--lr", type=float, default=cfg.lr) ap.add_argument("--clip-eps", type=float, default=cfg.clip_eps) ap.add_argument("--tis-cap", type=float, default=cfg.tis_cap) ap.add_argument("--max-new-tokens", type=int, default=cfg.max_new_tokens) ap.add_argument("--min-new-tokens", type=int, default=cfg.min_new_tokens) ap.add_argument("--temperature", type=float, default=cfg.temperature) ap.add_argument("--total-steps", type=int, default=cfg.total_steps) ap.add_argument("--sync-every", type=int, default=cfg.sync_every) ap.add_argument("--fluency-floor", type=float, default=cfg.fluency_floor) ap.add_argument("--distinct-floor", type=float, default=cfg.distinct_floor) ap.add_argument("--gate-penalty", type=float, default=cfg.gate_penalty) ap.add_argument("--len-penalty-start", type=int, default=cfg.len_penalty_start) ap.add_argument("--len-penalty-per-tok", type=float, default=cfg.len_penalty_per_tok) ap.add_argument("--no-gates", action="store_true", help="disable fluency/distinct/len shaping") ap.add_argument("--entropy-coef", type=float, default=cfg.entropy_coef, help="β for maximize r + β·H(π): direct diversity pressure (Bo-N depends on it)") ap.add_argument("--tp", type=int, default=int(os.environ.get("WORLD_SIZE", "1"))) ap.add_argument("--vllm-gpu-mem", type=float, default=0.35) ap.add_argument("--attn-backend", default="TRITON_ATTN", help="vLLM attention backend; TRITON_ATTN is the only one verified to expose " "the metadata vllm-lens needs (FLASHINFER silently breaks injection)") ap.add_argument("--vllm-max-len", type=int, default=1024) ap.add_argument("--micro-batch", type=int, default=8) ap.add_argument("--score-batch", type=int, default=64) ap.add_argument("--max-grad-norm", type=float, default=1.0) ap.add_argument("--save-every", type=int, default=500) ap.add_argument("--no-wandb", action="store_true") ap.add_argument("--seed", type=int, default=0) a = ap.parse_args() if a.no_gates: a.fluency_floor = a.distinct_floor = a.len_penalty_start = None # vLLM generation logprobs equal the sampling distribution's ONLY at T=1 (raw_logprobs). assert a.temperature == 1.0, "old_logp from vLLM is only valid at temperature 1.0" torch.manual_seed(a.seed) rng = np.random.default_rng(a.seed) device = "cuda:0" # HF actor lives here; vLLM TP shares all GPUs at gpu_memory_utilization tok = AutoTokenizer.from_pretrained(MODEL) if tok.pad_token is None: tok.pad_token = tok.eos_token prompt_ids, mpos = build_prompt_ids(tok) marker, p_len = mpos[0], len(prompt_ids) assert p_len + a.max_new_tokens <= a.vllm_max_len # ---- direction bank ---- if a.direction_source == "cluster": stats_p = f"{a.data_dir}/build_stats.json" n_vecs = (json.load(open(stats_p))["n_examples"] if os.path.exists(stats_p) else os.path.getsize(f"{a.data_dir}/vecs.f32") // (4 * D_MODEL)) bank = np.memmap(f"{a.data_dir}/vecs.f32", dtype=np.float32, mode="r", shape=(n_vecs, D_MODEL)) assert n_vecs >= a.groups_per_step else: # TODO: "sae" = unit encoder columns of the L27 SAE, "mix" = interleave cluster+sae. # The SAE loader isn't in this repo yet — port from max-activating-examples/src/maxact/sae.py. raise NotImplementedError(f"direction_source={a.direction_source!r}: only 'cluster' in the pilot") # ---- actor (HF + LoRA, cuda:0). NO gradient checkpointing EVER: recompute happens after the # inject-hook context exits → silently wrong grads. ---- actor = AutoModelForCausalLM.from_pretrained(MODEL, torch_dtype=torch.bfloat16, attn_implementation="sdpa", device_map={"": device}) if a.init_adapter: actor = PeftModel.from_pretrained(actor, a.init_adapter, is_trainable=True) else: actor = get_peft_model(actor, LoraConfig( r=tr.lora_r, lora_alpha=tr.lora_alpha, lora_dropout=0.0, use_rslora=True, target_modules="all-linear", bias="none", task_type="CAUSAL_LM")) actor.train() opt = torch.optim.AdamW([p for p in actor.parameters() if p.requires_grad], lr=a.lr, weight_decay=0.0) submodule = get_layer(actor, INJECT_LAYER) # ---- vLLM rollout engine (colocated, TP across all visible GPUs) ---- from vllm import LLM llm = LLM(model=MODEL, dtype="bfloat16", gpu_memory_utilization=a.vllm_gpu_mem, max_model_len=a.vllm_max_len, tensor_parallel_size=a.tp, enforce_eager=True, # MANDATORY — vllm-lens hooks don't fire under compiled graphs # MANDATORY — FLASHINFER (auto-picked on Blackwell) lacks query_start_loc metadata: # the injection hook SILENTLY skips every step. TRITON_ATTN provides it everywhere. attention_backend=a.attn_backend) print(f"[vllm] up tp={a.tp} | {n_vecs} directions | prompt {p_len} toks, marker @{marker}", flush=True) print(f"[sync] initial {sync_weights(actor, llm):.1f}s", flush=True) # vLLM == actor at step 0 if not a.no_wandb: wandb.init(project="maxact-fast", name=a.run_name, config=vars(a)) os.makedirs(a.save_dir, exist_ok=True) B, G = a.groups_per_step, a.group_size for step in range(a.total_steps): t0 = time.time() idx = np.sort(rng.choice(n_vecs, size=B, replace=False)) # B distinct vec_idx (sorted: memmap-friendly) dirs = torch.nn.functional.normalize( torch.from_numpy(np.asarray(bank[idx], dtype=np.float32)), dim=-1) texts, gen_ids, old_lps = rollout(llm, prompt_ids, marker, dirs, a) t_roll = time.time() - t0 dirs_rep = dirs.repeat_interleave(G, 0).to(device) # [B*G, d] rollout i's group direction r = score(texts, dirs_rep, actor, tok, device, a) raw_r, gate_frac = r.clone(), 1.0 if a.fluency_floor is not None or a.distinct_floor is not None: flu, dis = fluency(texts, actor, tok, device, a) gate = torch.ones(B * G, dtype=torch.bool) if a.fluency_floor is not None: gate &= flu >= a.fluency_floor if a.distinct_floor is not None: gate &= dis >= a.distinct_floor # sign-safe subtract, NOT zero: zeroing would rank gated garbage above coherent # negative-dot rollouts r = r - a.gate_penalty * (~gate).float() gate_frac = gate.float().mean().item() if a.len_penalty_start is not None: over = torch.tensor([max(0, len(g) - a.len_penalty_start) for g in gen_ids], dtype=torch.float32) r = r - a.len_penalty_per_tok * over adv = (r.view(B, G) - r.view(B, G).mean(1, keepdim=True)).flatten().detach() # NO /std # pad the batch — prompt is shared, so p_len is constant across rows L = p_len + max(len(g) for g in gen_ids) ids = torch.full((B * G, L), tok.pad_token_id, dtype=torch.long) attn = torch.zeros((B * G, L), dtype=torch.long) old_lp = torch.zeros((B * G, L - p_len)) pt = torch.tensor(prompt_ids, dtype=torch.long) for i, (g, lp) in enumerate(zip(gen_ids, old_lps)): ids[i, :p_len] = pt ids[i, p_len : p_len + len(g)] = torch.tensor(g) attn[i, : p_len + len(g)] = 1 old_lp[i, : len(g)] = lp stats = update(actor, opt, submodule, ids, attn, p_len, marker, old_lp, adv, dirs_rep, a, device) sync_s = sync_weights(actor, llm) if (step + 1) % a.sync_every == 0 else 0.0 secs = time.time() - t0 n_gen = float(sum(len(g) for g in gen_ids)) log = {"reward/mean": raw_r.mean().item(), "reward/std": raw_r.std().item(), "reward/max": raw_r.max().item(), "reward/shaped_mean": r.mean().item(), "reward/gate_frac": gate_frac, "ratio/clipfrac": stats["clipfrac"], "policy/entropy": stats["entropy"], "loss": stats["loss"], "grad_norm": stats["grad_norm"], "rollout/mean_logp": torch.cat(old_lps).mean().item(), "rollout/len_mean": n_gen / (B * G), "tokens_per_sec": n_gen / secs, "time/rollout_s": t_roll, "time/sync_s": sync_s, "time/step_s": secs} print(f"step {step:05d} | r {log['reward/mean']:.2f} (max {log['reward/max']:.1f}) | " f"gate {gate_frac:.0%} | clip {log['ratio/clipfrac']:.2%} | len {log['rollout/len_mean']:.0f} " f"| {log['tokens_per_sec']:.0f} tok/s | {secs:.0f}s", flush=True) if step % 10 == 0: print(f" sample r={raw_r[0]:.2f}: {texts[0][:110]!r}", flush=True) if not a.no_wandb: wandb.log(log, step=step) if a.save_every and step and step % a.save_every == 0: actor.save_pretrained(f"{a.save_dir}/step_{step}") actor.save_pretrained(f"{a.save_dir}/final") print("RL_DONE", flush=True) if __name__ == "__main__": main()