omnijev-work / code /train.py
fnruha0921's picture
stage3: drop snake records from training
5cede12 verified
Raw History Blame Contribute Delete
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)
@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:])