| """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") |
|
|
| 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): |
| llm.apply_model(functools.partial(_load_chunk, chunk=buckets[name])) |
| try: |
| llm.llm_engine.reset_prefix_cache() |
| except AttributeError: |
| pass |
| 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: |
| |
| |
| 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): |
| |
| |
| 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]})) |
| |
| |
| 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) |
| |
| |
| 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" |
| 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") |
| keep = mask.clone() |
| keep[:, 0] = False |
| 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) |
| 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: |
| |
| |
| 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() |
| 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: |
| 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 |
| |
| 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" |
|
|
| 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 |
|
|
| |
| 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: |
| |
| |
| raise NotImplementedError(f"direction_source={a.direction_source!r}: only 'cluster' in the pilot") |
|
|
| |
| |
| 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) |
|
|
| |
| 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, |
| |
| |
| 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) |
|
|
| 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)) |
| 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) |
|
|
| 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 |
| |
| |
| 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() |
|
|
| |
| 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() |
|
|