Download code/train.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 16.5 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/train.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/train.py
-
curl -L -o train.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/train.py
16.5 kB
| """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) | |
| 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 | |
| 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()} | |
| 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+" | |
| 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 | |
| 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:]) | |