"""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