"""Elastic training of MM-Jev (run inside the Colab kernel after build_data.py; `jev`, `DATA` in globals). One run trains two inference configs that share every weight: full : 35 layers, FFN width 16384, media tokens kept fast : exit after layer 19 (head "20"), FFN width 8192 (MatFormer E2B slice), media tokens dropped after layer 8, 8 latent memory tokens per media segment Loss per step = proper score(full) + proper score(fast) + lambda * KL(stopgrad p_full || p_fast) (self-distillation, LayerSkip / MatFormer style). LoRA r=16 on attention q/k/v/o only, so FFN slicing stays exact. Trains in a background thread; progress in /content/train.log. """ import collections, math, random, re, threading, time, json import torch, torch.nn.functional as F from torch.utils.checkpoint import checkpoint from mmjev import FastConfig, decision_loss, set_ffn_width, QTYPES FULL = FastConfig(n_latents=8, exit_layer=35, media_exit=None, sibling_from=14) FAST = FastConfig(n_latents=8, exit_layer=20, media_exit=8, sibling_from=14) FAST_WIDTH = 8192 LOG = "/content/train.log" def log(*a): with open(LOG, "a") as f: print(time.strftime("%H:%M:%S"), *a, file=f, flush=True) class LoRALinear(torch.nn.Module): """y = W x + (B A x) * alpha / r ; W frozen fp16, A/B fp32 (B zero-initialised).""" def __init__(self, base, r=16, alpha=32, dropout=0.05): super().__init__() self.base_layer = base self.lora_A = torch.nn.Parameter(torch.randn(r, base.in_features, device=base.weight.device) / math.sqrt(base.in_features)) self.lora_B = torch.nn.Parameter(torch.zeros(base.out_features, r, device=base.weight.device)) self.scale, self.drop = alpha / r, torch.nn.Dropout(dropout) @property def weight(self): return self.base_layer.weight def forward(self, x): y = self.base_layer(x) out_f, in_f = self.base_layer.weight.shape # smaller than the LoRA shapes when the FFN is sliced A = self.lora_A[:, :in_f].to(x.dtype) Bm = self.lora_B[:out_f].to(x.dtype) lx = torch.nn.functional.linear(torch.nn.functional.linear(self.drop(x), A), Bm) return y + lx * self.scale @torch.no_grad() def reset_trainable(jev): """Fresh start: LoRA B = 0 (A re-drawn), heads = E[Yes] - E[No], latents re-initialised.""" for m in jev.lm.modules(): if isinstance(m, LoRALinear): m.lora_A.normal_().div_(math.sqrt(m.lora_A.shape[1])); m.lora_B.zero_() E = jev.lm.embed_tokens.weight w = (E[jev.yes_id].float() - E[jev.no_id].float())[None] for h in jev.heads.values(): h.weight.copy_(w); h.bias.zero_() init = jev._embed_ids(torch.tensor(jev._t(" summary"))).float().mean(0) jev.latents.copy_(init[None].repeat(jev.latents.shape[0], 1) + 0.02 * init.std() * torch.randn_like(jev.latents)) def add_lora(jev, r=16, alpha=32, mlp_from=10): """Attention q/k/v/o everywhere (layers 20-34 have no k/v: KV-shared) + MLP gate/up/down from layer `mlp_from`.""" for p in jev.base.parameters(): p.requires_grad = False n = 0 for li, layer in enumerate(jev.lm.layers): att = layer.self_attn for name in ("q_proj", "k_proj", "v_proj", "o_proj"): m = getattr(att, name, None) if isinstance(m, torch.nn.Linear): setattr(att, name, LoRALinear(m, r, alpha)); n += 1 if li >= mlp_from: for name in ("gate_proj", "up_proj", "down_proj"): m = getattr(layer.mlp, name) if isinstance(m, torch.nn.Linear): setattr(layer.mlp, name, LoRALinear(m, r, alpha)); n += 1 n_lora = sum(p.numel() for nme, p in jev.base.named_parameters() if "lora_" in nme) log(f"LoRA on {n} projections, {n_lora / 1e6:.1f}M params") EXCLUDE_TASKS = {"typed_decisions", "snake_syn"} # benchmark train split + demo task: both kept out (zero-shot) _SNAKE_Q = re.compile(r"\bsnake\b", re.I) def is_snake(r): """Snake-game records from any public corpus (Open-Jev snake-v1 boards, snake questions): kept out of training.""" if r["task"].startswith("snake"): return True if any(s.kind == "text" and '"game": "snake"' in str(s.data) for s in r["state"]): return True return any(_SNAKE_Q.search(q.get("instructions", "")) for q in r["qs"]) def train_records(DATA, calib_frac=0.05, seed=0, modalities=None, max_records=None): recs = [r for k, v in list(DATA.items()) for r in v if r["split"] == "train" and r["task"] not in EXCLUDE_TASKS and (modalities is None or r["modality"] in modalities)] rng = random.Random(seed) rng.shuffle(recs) if max_records: recs = recs[:max_records] n_cal = int(len(recs) * calib_frac) return recs[n_cal:], recs[:n_cal] def batch_logits(jev, batch, fc, width): set_ffn_width(jev.lm, width) plans = [jev.plan(jev.encode_state(r["state"], fc), r["qs"], fc) for r in batch] pk = jev.pack(plans) return jev.run(pk, fc) def step_loss(jev, batch, fc, width, teacher=None, kl_w=0.5): out = batch_logits(jev, batch, fc, width) loss, n, probs = 0.0, 0, [] for r, per_q in zip(batch, out): pq = [] for q, lg, tgt in zip(r["qs"], per_q, r["targets"]): t = torch.tensor(tgt, device=lg.device, dtype=torch.float32) loss = loss + decision_loss(lg, t, q["type"]) + 0.1 * ((torch.softmax(lg, -1) - t) ** 2).sum() pq.append(torch.softmax(lg.detach(), -1)) n += 1 probs.append(pq) if teacher is not None: kl = 0.0 for per_q, tq in zip(out, teacher): for lg, pt in zip(per_q, tq): kl = kl + F.kl_div(torch.log_softmax(lg, -1), pt, reduction="sum") loss = loss + kl_w * kl return loss / max(n, 1), probs def enable_ckpt(jev): """Gradient checkpointing per decoder layer (the shared-KV dict is refilled identically on recompute).""" for layer in jev.lm.layers: if not hasattr(layer, "_orig_fwd"): layer._orig_fwd = layer.forward def fwd(*a, _l=layer, **k): if _l.training and torch.is_grad_enabled(): return checkpoint(_l._orig_fwd, *a, use_reentrant=False, **k) return _l._orig_fwd(*a, **k) layer.forward = fwd TRAIN_MAX_STATE, TRAIN_MAX_OPTS = 1500, 32 def cap_record(r, rng): """Training-time cost cap: truncate long text states; for questions with > TRAIN_MAX_OPTS options keep the gold option plus a random subset (label-subsampling augmentation), renormalising soft targets.""" from mmjev import Seg r = dict(r) r["state"] = [Seg(s.kind, s.data[:TRAIN_MAX_STATE], s.audio, s.fps) if s.kind == "text" else s for s in r["state"]] qs, ts, ys = [], [], [] for q, t, y in zip(r["qs"], r["targets"], r["ys"]): if q["type"] == "choice" and len(t) > TRAIN_MAX_OPTS: keys = list(q["criteria"]) keep = sorted([y] + rng.sample([i for i in range(len(keys)) if i != y], TRAIN_MAX_OPTS - 1)) q = dict(q, criteria={keys[i]: q["criteria"][keys[i]] for i in keep}) tt = [t[i] for i in keep]; sm = sum(tt) or 1.0 t = [x / sm for x in tt]; y = keep.index(y) qs.append(q); ts.append(t); ys.append(y) r.update(qs=qs, targets=ts, ys=ys) return r def est_len(r): n = 0 for s in r["state"]: n += len(s.data) // 4 if s.kind == "text" else 64 if s.kind == "image" else 150 if s.kind == "video" else 60 for q in r["qs"]: n += 30 + 12 * len(q.get("criteria") or [0, 0]) return n def train(jev, DATA, epochs=1, bs=4, lr_lora=2e-4, lr_head=1e-3, warmup=50, modalities=None, max_records=None, save_path="/content/mmjev_adapter.pt", lora=True, ckpt_dir=None, ckpt_every=200, recs=None, on_ckpt=None, init=None): """Elastic training; with ckpt_dir it checkpoints (adapter + optimizer + scaler + step) and resumes.""" import os if recs is None: tr, cal = train_records(DATA, modalities=modalities, max_records=max_records) else: tr, cal = recs drop = collections.Counter(r["task"] for r in tr + cal if is_snake(r)) tr, cal = [r for r in tr if not is_snake(r)], [r for r in cal if not is_snake(r)] log(f"snake records removed: {sum(drop.values())} {dict(drop.most_common(8))}") globals()["CALIB"] = cal log(f"train records {len(tr)} calib {len(cal)}") if lora: add_lora(jev) if init: # warm start (e.g. the text-stage adapter); a checkpoint overrides it load_trainable(jev, torch.load(init, map_location="cpu", weights_only=False)) log(f"initialised trainable weights from {init}") enable_ckpt(jev) params = [ {"params": [p for n, p in jev.base.named_parameters() if "lora_" in n], "lr": lr_lora}, {"params": list(jev.heads.parameters()) + [jev.latents], "lr": lr_head}, ] opt = torch.optim.AdamW(params, weight_decay=0.0) scaler = torch.amp.GradScaler("cuda") start = 0 if ckpt_dir: os.makedirs(ckpt_dir, exist_ok=True) last = f"{ckpt_dir}/last.pt" if os.path.exists(last): ck = torch.load(last, map_location="cpu", weights_only=False) load_trainable(jev, ck) opt.load_state_dict(ck["opt"]); scaler.load_state_dict(ck["scaler"]) start = ck["step"] log(f"RESUMED from step {start}") rng = random.Random(0) tr = [cap_record(r, rng) for r in tr] # token-budget batches: sort inside chunks by estimated length, fill a batch until max_len * n exceeds the budget # (padded tokens ~ memory), cap at bs records; then shuffle the batch order. Deterministic -> resumable. budget = int(os.environ.get("MMJEV_TOKEN_BUDGET", 6000)) order = list(range(len(tr))); rng.shuffle(order) batches = [] for c in range(0, len(order), 64 * bs): chunk = sorted(order[c:c + 64 * bs], key=lambda i: est_len(tr[i])) cur, mx = [], 0 for i in chunk: l = est_len(tr[i]) if cur and (max(mx, l) * (len(cur) + 1) > budget or len(cur) >= bs): batches.append(cur); cur, mx = [], 0 cur.append(i); mx = max(mx, l) if cur: batches.append(cur) rng.shuffle(batches) steps = len(batches) * epochs sched = torch.optim.lr_scheduler.LambdaLR( opt, lambda s: min(1.0, (s + 1) / warmup) * 0.5 * (1 + math.cos(math.pi * min(1.0, s / steps)))) if start: sched.load_state_dict(ck["sched"]) jev.train() t0, ema = time.time(), None step = start base_mem = torch.cuda.memory_allocated() for bi in range(start, steps): batch = [tr[i] for i in batches[bi % len(batches)]] try: opt.zero_grad(set_to_none=True) lf, pf = step_loss(jev, batch, FULL, None) scaler.scale(lf).backward() ls, _ = step_loss(jev, batch, FAST, FAST_WIDTH, teacher=pf) scaler.scale(ls).backward() set_ffn_width(jev.lm, None) scaler.unscale_(opt) torch.nn.utils.clip_grad_norm_([p for g in params for p in g["params"]], 1.0) scaler.step(opt); scaler.update(); sched.step() except torch.OutOfMemoryError: oom = True else: oom = False if oom: # free the failed step OUTSIDE the except block: the live exception pins every activation of the step lf = ls = pf = None set_ffn_width(jev.lm, None) opt.zero_grad(set_to_none=True) import gc; gc.collect(); torch.cuda.empty_cache() leaked = torch.cuda.memory_allocated() - base_mem log(f"OOM at step {step}, skipped (tasks {[r['task'] for r in batch]}) " f"mem now {torch.cuda.memory_allocated() / 2**30:.1f}G, leaked {leaked / 2**30:.1f}G") if leaked > 2 * 2**30 and ckpt_dir: ck = trainable_state(jev) ck.update(opt=opt.state_dict(), scaler=scaler.state_dict(), sched=sched.state_dict(), step=step) torch.save(ck, f"{ckpt_dir}/last.tmp"); os.replace(f"{ckpt_dir}/last.tmp", f"{ckpt_dir}/last.pt") if on_ckpt: on_ckpt(f"{ckpt_dir}/last.pt", step) log(f"UNRECOVERABLE OOM leak -> checkpointed step {step}, exiting for a clean restart") time.sleep(60) # let the hub upload finish os._exit(3) step += 1; sched.step() continue step += 1 v = (float(lf), float(ls)) ema = v if ema is None else (0.98 * ema[0] + 0.02 * v[0], 0.98 * ema[1] + 0.02 * v[1]) if step % 50 == 0 or step == start + 1: el = time.time() - t0 done = step - start log(f"step {step}/{steps} loss full {ema[0]:.3f} fast {ema[1]:.3f} | {el / done:.2f}s/step " f"ETA {(steps - step) * el / done / 60:.0f}m | mem {torch.cuda.max_memory_allocated() / 2**30:.1f}G") if ckpt_dir and step % ckpt_every == 0: ck = trainable_state(jev) ck.update(opt=opt.state_dict(), scaler=scaler.state_dict(), sched=sched.state_dict(), step=step) torch.save(ck, f"{ckpt_dir}/last.tmp"); os.replace(f"{ckpt_dir}/last.tmp", f"{ckpt_dir}/last.pt") if on_ckpt: on_ckpt(f"{ckpt_dir}/last.pt", step) jev.eval() set_ffn_width(jev.lm, None) save(jev, save_path) log(f"TRAIN DONE {time.time() - t0:.0f}s") def trainable_state(jev): return {"lora": {n: p.detach().cpu().clone() for n, p in jev.base.named_parameters() if "lora_" in n}, "heads": {k: v.cpu().clone() for k, v in jev.heads.state_dict().items()}, "latents": jev.latents.detach().cpu().clone()} @torch.no_grad() def load_trainable(jev, ck): params = dict(jev.base.named_parameters()) for n, v in ck["lora"].items(): params[n].copy_(v.to(params[n].device)) jev.heads.load_state_dict(ck["heads"]) jev.latents.copy_(ck["latents"].to(jev.latents.device)) def save(jev, path): torch.save(trainable_state(jev), path) def k_bucket(k): return "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11-30" if k <= 30 else "31+" @torch.no_grad() def fit_temperature_k(jev, recs, fc, width, bs=8): """Temperature per (question type, option-count bucket) -- a 2-way noul and a 72-way choice need different scaling.""" jev.eval() per = {} for s in range(0, len(recs), bs): batch = recs[s:s + bs] out = batch_logits(jev, batch, fc, width) for r, per_q in zip(batch, out): for q, lg, tgt in zip(r["qs"], per_q, r["targets"]): per.setdefault(f"{q['type']}:{k_bucket(len(lg))}", []).append((lg.float().cpu(), torch.tensor(tgt))) set_ffn_width(jev.lm, None) grid = [0.5, 0.6, 0.7, 0.8, 0.9, 1.0, 1.15, 1.3, 1.5, 1.75, 2.0, 2.5, 3.0, 4.0] temps = {} for key, items in per.items(): if len(items) < 20: continue temps[key] = min(grid, key=lambda T: sum(float(-(tg * torch.log_softmax(lg / T, -1)).sum()) for lg, tg in items)) return temps @torch.no_grad() def fit_temperature(jev, recs, fc, width, bs=8): """Per question type temperature by NLL grid search on the calibration records.""" jev.eval() per = {t: [] for t in QTYPES} for s in range(0, len(recs), bs): batch = recs[s:s + bs] out = batch_logits(jev, batch, fc, width) for r, per_q in zip(batch, out): for q, lg, tgt in zip(r["qs"], per_q, r["targets"]): per[q["type"]].append((lg.float().cpu(), torch.tensor(tgt))) set_ffn_width(jev.lm, None) temps = {} for t, items in per.items(): if not items: temps[t] = 1.0; continue best = min((sum(float(-(tg * torch.log_softmax(lg / T, -1)).sum()) for lg, tg in items), T) for T in [0.5, 0.6, 0.7, 0.8, 0.9, 1.0, 1.15, 1.3, 1.5, 1.75, 2.0, 2.5, 3.0, 4.0]) temps[t] = best[1] return temps def start(jev, DATA, **kw): th = threading.Thread(target=lambda: _safe(train, jev, DATA, **kw), daemon=True) th.start() return th def _safe(fn, *a, **k): try: fn(*a, **k) except Exception as e: # noqa: BLE001 import traceback log("TRAIN FAILED", repr(e), traceback.format_exc()[-2000:])