Text Classification
Transformers
Safetensors
English
qwen3_5_text
text-generation
system-one
typed-decisions
decision-model
calibrated-probabilities
knowledge-distillation
jev
noul
choice
score
lora
qwen3_5
dual-head
vllm
Eval Results (legacy)
Instructions to use autotrust/JEV with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use autotrust/JEV with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="autotrust/JEV")# Load model directly from transformers import AutoTokenizer, AutoModelForCausalLM tokenizer = AutoTokenizer.from_pretrained("autotrust/JEV") model = AutoModelForCausalLM.from_pretrained("autotrust/JEV", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """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 | |
| 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) | |
| 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 | |
| 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 | |