JEV / code /src /jev_judge /train_loop.py
cloudyu's picture
vLLM inference: adapter_vllm/ (decision head as lm_head LoRA), model card section "Inference with vLLM", OpenAI client, measured speed
b16c3a6 verified
Raw
History Blame Contribute Delete
16.4 kB
"""Single-GPU training loop (B200 edition) for S1 (head probe) / S2 (head + LoRA) / S3 (full).
Key B200 adaptations vs DESIGN v0.6 (§7):
* optimizer batch = `batch_size` samples, internally split into micro-batches capped by
`max_padded_tokens` (B*T) so memory stays bounded regardless of the length mix; the loss is
weighted so the step equals the weighted mean over the whole batch;
* backbone runs under bf16 autocast, LoRA adapters keep fp32 master weights (peft default dtype is
bf16 -> we upcast) and AdamW is fused; the head is always fp32 outside autocast;
* optional non-reentrant gradient checkpointing (needed for the 27B / very large micro-batches);
* S1 runs the frozen backbone under no_grad (no activation storage at all).
"""
from __future__ import annotations
import json
import math
import os
import time
from dataclasses import asdict, dataclass, field
import numpy as np
import pandas as pd
import torch
from torch.utils.data import DataLoader
from .checkpointing import save_checkpoint
from .data import JevDataset, KindBatchSampler, apply_d1_policy, collate, load_split, subset
from .losses import judge_loss, per_sample_kl
from .model import JevJudge, LoraSpec, slot_mask
from .template import KINDS
@dataclass
class TrainConfig:
model_path: str = "/root/models/Qwen3.5-9B"
data_dir: str = "data"
out_dir: str = "checkpoints/run"
stage: str = "s2" # s1 | s2 | s3
seed: int = 42
subset_frac: float = 1.0
max_seq_len: int = 1024
epochs: int = 2
batch_size: int = 128 # samples per optimizer step
max_padded_tokens: int = 10000 # B*T cap per micro-batch (memory bound)
gradient_checkpointing: bool = False
lr_head: float = 2e-4
lr_lora: float = 1e-4
lr_backbone: float = 1e-5
weight_decay: float = 0.0
warmup_frac: float = 0.03
min_lr_frac: float = 0.10
grad_clip: float = 1.0
lambda_rps: float = 0.5
d1_policy: str = "downweight" # keep | downweight | drop
d1_downweight: float = 0.05
d1_scope: str = "yuri_v1"
choice_permute_prob: float = 0.30
kind_floor: float = 1 / 6
lora: dict = field(default_factory=lambda: {"r": 16, "alpha": 32, "dropout": 0.05,
"target_modules": ["in_proj_qkv", "in_proj_z", "q_proj", "k_proj", "v_proj",
"o_proj", "gate_proj", "up_proj", "down_proj", "out_proj"]})
eval_every: int = 500 # optimizer steps
eval_rows: int = 4000 # validation rows for periodic eval (full validation at epoch end)
patience: int = 3 # evals without val-KL improvement before stopping
log_every: int = 20
num_workers: int = 8
max_steps: int = 0 # >0: stop after this many optimizer steps (smoke)
init_from: str = "" # checkpoint dir to initialise head (+adapter) from, e.g. S1 -> S2
save_optimizer: bool = False
skip_batches: int = 0 # skip the first N batches of epoch 0 (continue an interrupted epoch on unseen data)
@classmethod
def from_yaml(cls, path: str, overrides: dict | None = None) -> "TrainConfig":
import yaml
with open(path) as f:
d = yaml.safe_load(f) or {}
d.update({k: v for k, v in (overrides or {}).items() if v is not None})
return cls(**d)
def cosine_lr(step: int, total: int, warmup: int, min_frac: float) -> float:
if step < warmup:
return (step + 1) / max(1, warmup)
prog = min(1.0, (step - warmup) / max(1, total - warmup))
return min_frac + (1 - min_frac) * 0.5 * (1 + math.cos(math.pi * prog))
def split_micro(items: list[dict], max_padded_tokens: int) -> list[list[dict]]:
"""Sort by length and greedily pack into micro-batches with B*T <= max_padded_tokens."""
items = sorted(items, key=lambda d: len(d["input_ids"]))
out: list[list[dict]] = []
cur: list[dict] = []
for it in items:
L = len(it["input_ids"])
if cur and (len(cur) + 1) * L > max_padded_tokens: # L is the max so far (sorted)
out.append(cur)
cur = []
cur.append(it)
if cur:
out.append(cur)
return out
class Trainer:
def __init__(self, cfg: TrainConfig):
self.cfg = cfg
torch.manual_seed(cfg.seed)
np.random.seed(cfg.seed)
os.makedirs(cfg.out_dir, exist_ok=True)
self.log_f = open(os.path.join(cfg.out_dir, "log.jsonl"), "a")
self.judge = JevJudge.from_base(cfg.model_path)
self.tok = self.judge.tokenizer
self.dev = self.judge.device
self.lora_spec: LoraSpec | None = None
if cfg.init_from:
self._init_from(cfg.init_from)
if cfg.stage in ("s2",) and not (cfg.init_from and self._has_adapter):
self.lora_spec = LoraSpec(**cfg.lora)
targets = self.judge.attach_lora(self.lora_spec)
self.log({"event": "lora", "targets": targets,
"adapter_params": sum(p.numel() for n, p in self.judge.lm.named_parameters() if "lora_" in n)})
self.judge.set_stage_grads(cfg.stage)
if cfg.gradient_checkpointing and cfg.stage != "s1":
self.judge.backbone.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
self.judge.backbone.train(cfg.stage != "s1")
self.judge.head.train()
# data
train_df = load_split(cfg.data_dir, "train")
if cfg.subset_frac < 1.0:
train_df = subset(train_df, cfg.subset_frac, cfg.seed)
train_df, weights = apply_d1_policy(train_df, cfg.d1_policy, cfg.d1_downweight, cfg.d1_scope)
self.train_ds = JevDataset(train_df, self.tok, cfg.max_seq_len, weights, cfg.choice_permute_prob, cfg.seed)
val_df = load_split(cfg.data_dir, "validation")
self.val_full = JevDataset(val_df, self.tok, cfg.max_seq_len)
sub = val_df.sample(min(cfg.eval_rows, len(val_df)), random_state=cfg.seed).reset_index(drop=True)
self.val_sub = JevDataset(sub, self.tok, cfg.max_seq_len)
self.sampler = KindBatchSampler(self.train_ds.kind_ids, self.train_ds.lengths, cfg.batch_size, cfg.kind_floor, seed=cfg.seed)
self.steps_per_epoch = len(self.sampler)
planned = self.steps_per_epoch * cfg.epochs - cfg.skip_batches
self.total_steps = planned if cfg.max_steps <= 0 else min(cfg.max_steps, planned)
# optimizer
groups = self.judge.trainable_parameters(cfg.stage)
lrs = {"head": cfg.lr_head, "lora": cfg.lr_lora, "backbone": cfg.lr_backbone}
self.param_groups = [{"params": ps, "lr": lrs[name], "base_lr": lrs[name], "name": name, "weight_decay": cfg.weight_decay}
for name, ps in groups.items() if ps]
self.opt = torch.optim.AdamW(self.param_groups, betas=(0.9, 0.98), eps=1e-8, fused=True)
self.n_trainable = sum(p.numel() for g in self.param_groups for p in g["params"])
self.log({"event": "setup", "config": asdict(cfg), "train_rows": len(self.train_ds), "steps_per_epoch": self.steps_per_epoch,
"total_steps": self.total_steps, "trainable_params": self.n_trainable,
"d1_rows": int((weights < 1).sum()), "gpu": torch.cuda.get_device_name()})
# ------------------------------------------------------------------------------------------
def _init_from(self, ckpt_dir: str) -> None:
from safetensors.torch import load_file
head_sd = load_file(os.path.join(ckpt_dir, "head.safetensors"))
self.judge.head.load_state_dict({k: v.to(self.dev) for k, v in head_sd.items()})
adapter = os.path.join(ckpt_dir, "adapter")
self._has_adapter = os.path.isdir(adapter)
if self._has_adapter:
from peft import PeftModel
self.judge.lm = PeftModel.from_pretrained(self.judge.lm, adapter, is_trainable=True)
for n, p in self.judge.lm.named_parameters():
if "lora_" in n:
p.data = p.data.to(torch.float32)
with open(os.path.join(ckpt_dir, "judge_config.json")) as f:
lc = json.load(f).get("lora")
self.lora_spec = LoraSpec(**lc) if lc else None
self.log({"event": "init_from", "path": ckpt_dir, "adapter": self._has_adapter})
_has_adapter = False
def log(self, rec: dict) -> None:
rec = {"t": time.time(), **rec}
self.log_f.write(json.dumps(rec) + "\n")
self.log_f.flush()
if rec.get("event") not in ("train",):
print(json.dumps({k: v for k, v in rec.items() if k not in ("config",)})[:400], flush=True)
# ------------------------------------------------------------------------------------------
def _forward_logits(self, mb: dict, grad_backbone: bool) -> tuple[torch.Tensor, torch.Tensor]:
ids = mb["input_ids"].to(self.dev, non_blocking=True)
am = mb["attention_mask"].to(self.dev, non_blocking=True)
L = mb["lengths"].to(self.dev)
if grad_backbone:
with torch.autocast("cuda", dtype=torch.bfloat16):
h = self.judge.hidden_last(ids, am, L)
else:
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
h = self.judge.hidden_last(ids, am, L)
z = self.judge.head(h.float())
return z, slot_mask(mb["kind_ids"].to(self.dev), mb["n_options"].to(self.dev))
def train_step(self, items: list[dict]) -> dict:
cfg = self.cfg
micro = split_micro(items, cfg.max_padded_tokens)
w_total = float(sum(it["weight"] for it in items))
stats = {"loss": 0.0, "kl": 0.0, "n_micro": len(micro), "tokens": 0, "padded": 0, "shapes": []}
for m_items in micro:
mb = collate(m_items, self.tok.pad_token_id)
stats["shapes"].append([int(mb["input_ids"].shape[0]), int(mb["input_ids"].shape[1])])
z, mask = self._forward_logits(mb, grad_backbone=cfg.stage != "s1")
w = mb["weight"].to(self.dev)
loss, st = judge_loss(z, mb["target"].to(self.dev), mask, mb["kind_ids"].to(self.dev), w, cfg.lambda_rps)
scale = float(w.sum()) / w_total # so the sum over micro-batches = weighted mean over the batch
(loss * scale).backward()
stats["loss"] += st["loss"] * scale
stats["kl"] += st["kl"] * scale
stats["tokens"] += int(mb["lengths"].sum())
stats["padded"] += int(mb["input_ids"].numel())
if cfg.grad_clip > 0:
gn = torch.nn.utils.clip_grad_norm_([p for g in self.param_groups for p in g["params"]], cfg.grad_clip)
stats["grad_norm"] = float(gn)
self.opt.step()
self.opt.zero_grad(set_to_none=True)
stats["peak_mem_gb"] = torch.cuda.max_memory_allocated() / 1024**3
return stats
@torch.no_grad()
def evaluate(self, ds: JevDataset, desc: str) -> dict:
self.judge.eval()
order = np.argsort(ds.lengths, kind="stable")
batches, cur = [], []
for i in order:
L = int(ds.lengths[i])
if cur and (len(cur) + 1) * L > 32768:
batches.append(cur); cur = []
cur.append(int(i))
if cur:
batches.append(cur)
dl = DataLoader(ds, batch_sampler=batches, num_workers=self.cfg.num_workers, collate_fn=lambda b: collate(b, self.tok.pad_token_id))
kls, kinds, uni = [], [], []
for mb in dl:
z, mask = self._forward_logits(mb, grad_backbone=False)
kl = per_sample_kl(z, mb["target"].to(self.dev), mask)
kls.append(kl.cpu().numpy()); kinds.append(mb["kind_ids"].numpy())
uni.append(ds.df["is_uniform"].to_numpy()[mb["index"].numpy()])
kl = np.concatenate(kls); kind = np.concatenate(kinds); uni = np.concatenate(uni)
out = {"kl": float(kl.mean()), "kl_nonuniform": float(kl[~uni].mean()) if (~uni).any() else float("nan")}
for i, k in enumerate(KINDS):
sel = kind == i
if sel.any():
out[f"kl_{k}"] = float(kl[sel].mean())
self.judge.backbone.train(self.cfg.stage != "s1")
self.judge.head.train()
return out
# ------------------------------------------------------------------------------------------
def fit(self) -> dict:
cfg = self.cfg
warmup = int(cfg.warmup_frac * self.total_steps)
best = float("inf"); bad = 0; step = 0; stop = False
t_start = time.time(); tok_seen = 0
dl_kwargs = dict(num_workers=cfg.num_workers, collate_fn=lambda b: b, pin_memory=False, persistent_workers=False, prefetch_factor=4)
for epoch in range(cfg.epochs):
self.sampler.set_epoch(epoch)
self.train_ds.set_epoch(epoch)
batch_sampler = self.sampler
if epoch == 0 and cfg.skip_batches > 0:
order = list(iter(self.sampler)) # deterministic for (seed, epoch)
batch_sampler = order[cfg.skip_batches:]
self.log({"event": "skip_batches", "skipped": cfg.skip_batches, "remaining": len(batch_sampler)})
dl = DataLoader(self.train_ds, batch_sampler=batch_sampler, **dl_kwargs)
t_wait = time.time()
data_s_acc = 0.0
for items in dl:
data_s = time.time() - t_wait
data_s_acc += data_s
lr_mult = cosine_lr(step, self.total_steps, warmup, cfg.min_lr_frac)
for g in self.opt.param_groups:
g["lr"] = g["base_lr"] * lr_mult
t0 = time.time()
st = self.train_step(items)
dt = time.time() - t0
step += 1; tok_seen += st["tokens"]
if step % cfg.log_every == 0 or step == 1:
el = time.time() - t_start
rec = {"event": "train", "step": step, "epoch": epoch, "lr_mult": lr_mult, **st, "step_s": dt, "data_s": data_s,
"data_s_total": data_s_acc, "tok_per_s": st["tokens"] / dt, "avg_tok_per_s": tok_seen / el,
"eta_h": (self.total_steps - step) * (el / step) / 3600}
self.log(rec)
print(f"step {step}/{self.total_steps} ep{epoch} loss {st['loss']:.4f} kl {st['kl']:.4f} gn {st.get('grad_norm', 0):.2f} "
f"{st['tokens']/dt:.0f} tok/s (avg {tok_seen/el:.0f}) micro {st['n_micro']} {st['shapes']} peak {st['peak_mem_gb']:.0f}GB step {dt:.2f}s data {data_s:.2f}s "
f"(Σdata {data_s_acc:.0f}s) eta {rec['eta_h']:.2f}h", flush=True)
if step % cfg.eval_every == 0 or step == self.total_steps:
ev = self.evaluate(self.val_sub, "val_sub")
improved = ev["kl"] < best - 1e-4
self.log({"event": "eval", "step": step, "epoch": epoch, "split": "val_sub", **ev, "best": min(best, ev["kl"]), "improved": improved})
state = {"step": step, "epoch": epoch, "val_kl": ev["kl"], "best_val_kl": min(best, ev["kl"])}
save_checkpoint(self.judge, os.path.join(cfg.out_dir, "last"), cfg.stage, self.lora_spec, state,
optimizer=self.opt if cfg.save_optimizer else None)
if improved:
best = ev["kl"]; bad = 0
save_checkpoint(self.judge, os.path.join(cfg.out_dir, "best"), cfg.stage, self.lora_spec, state)
else:
bad += 1
if bad >= cfg.patience:
self.log({"event": "early_stop", "step": step, "best": best}); stop = True
if stop or step >= self.total_steps:
break
t_wait = time.time()
if not stop:
ev = self.evaluate(self.val_full, "validation")
self.log({"event": "eval", "step": step, "epoch": epoch, "split": "validation_full", **ev})
if stop or step >= self.total_steps:
break
summary = {"event": "done", "steps": step, "best_val_kl": best, "hours": (time.time() - t_start) / 3600,
"avg_tok_per_s": tok_seen / max(1e-9, time.time() - t_start)}
self.log(summary)
return summary