muscgab's picture
JMangaTranslator-Fast v1
3d92ad2 verified
Raw History Blame Contribute Delete
18.3 kB
#!/usr/bin/env python3
"""Train the ModernBERT-ja + AR decoder translator (model.py).
Stage A (steps < --stage-a): encoder frozen and run without grad; bridge, fusion, decoder train.
Stage B: encoder unfrozen at lr * --enc-lr-mult, warmed up over --enc-warmup steps.
Schedule: linear warmup over --warmup steps, cosine decay to --final-lr-frac * lr at --steps.
Batches: token budget (padded src + padded tgt per row times rows <= --max-tokens), built from pools of
mixed-source examples sorted by length. Sources are memory-mapped packed files from build_data.py and are
mixed by --mix weights; each source walks a fresh seeded permutation per epoch. The sampler state is saved
with every checkpoint, so --resume continues on exactly the next batch.
Checkpoints: last.pt (model + optimizer + sampler, every --save-minutes and at the end), best.pt (model only,
best dev chrF). Log: train_log.jsonl.
"""
from __future__ import annotations
import argparse
import json
import math
import os
import random
import sys
import time
from collections import Counter
from pathlib import Path
import numpy as np
import sentencepiece as spm
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent))
from model import ARMT # noqa: E402
ROOT = Path(__file__).resolve().parents[1]
DAT_TOK = ROOT / "artifacts/preview200m_training_20260923/tokenizer/joint.model"
DAT_SHARED = ROOT / "artifacts/ar_mt_20261005/dat_shared_fp16.pt" # {"shared.weight": [40960,768]} from base_fp16.pt
DAT_WEIGHTS = ROOT / "artifacts/dat_context_20260926/pack/base_fp16.pt"
UNK, BOS, EOS, PAD = 0, 1, 2, 3
MB_PAD = 3
def chrf(hyps: list[str], refs: list[str]) -> float:
try:
import sacrebleu
return float(sacrebleu.metrics.CHRF().corpus_score(hyps, [refs]).score)
except ImportError: # fallback: char 6-gram, beta 2, whitespace removed
stats = np.zeros((6, 3))
for h, r in zip(hyps, refs):
h, r = h.replace(" ", ""), r.replace(" ", "")
for n in range(1, 7):
hc = Counter(h[i:i + n] for i in range(len(h) - n + 1))
rc = Counter(r[i:i + n] for i in range(len(r) - n + 1))
stats[n - 1] += (sum((hc & rc).values()), sum(hc.values()), sum(rc.values()))
p = np.mean([m / hh if hh else 0 for m, hh, _ in stats])
rr = np.mean([m / rf if rf else 0 for m, _, rf in stats])
return float(100 * 5 * p * rr / (4 * p + rr)) if p + rr else 0.0
class Source:
def __init__(self, data: Path, name: str, lut: np.ndarray):
self.name = name
self.src = np.memmap(data / f"{name}.src.bin", dtype=np.uint32, mode="r")
self.tgt = np.memmap(data / f"{name}.tgt.bin", dtype=np.uint16, mode="r")
self.so = np.load(data / f"{name}.src.idx.npy")
self.to = np.load(data / f"{name}.tgt.idx.npy")
self.n = len(self.so) - 1
self.slen, self.tlen = np.diff(self.so), np.diff(self.to)
self.lut = lut
def get(self, i: int):
return np.asarray(self.src[self.so[i]:self.so[i + 1]], dtype=np.int64), \
self.lut[np.asarray(self.tgt[self.to[i]:self.to[i + 1]], dtype=np.int64)]
class Sampler:
def __init__(self, sources: list[Source], weights: list[float], seed: int, max_tokens: int, pool: int):
self.sources, self.seed, self.max_tokens, self.pool = sources, seed, max_tokens, pool
self.w = np.asarray(weights, dtype=np.float64) / sum(weights)
self.state = {"rng": np.random.default_rng(seed).bit_generator.state, "epoch": [0] * len(sources),
"pos": [0] * len(sources), "consumed": 0}
self._perm = {}
self._pool_state = None
self._batches = []
def perm(self, k: int, epoch: int):
"""Permutation of source k for this epoch; one cached per source (sources alternate within a pool)."""
cached = self._perm.get(k)
if cached is None or cached[0] != epoch:
cached = (epoch, np.random.default_rng([self.seed, k, epoch]).permutation(self.sources[k].n))
self._perm[k] = cached
return cached[1]
def _build_pool(self):
self._pool_state = json.loads(json.dumps(self.state))
rng = np.random.default_rng()
rng.bit_generator.state = self.state["rng"]
picks = rng.choice(len(self.sources), size=self.pool, p=self.w)
items = []
for k in picks:
s = self.sources[k]
if self.state["pos"][k] >= s.n:
self.state["epoch"][k] += 1
self.state["pos"][k] = 0
items.append((k, int(self.perm(k, self.state["epoch"][k])[self.state["pos"][k]])))
self.state["pos"][k] += 1
items.sort(key=lambda it: (int(self.sources[it[0]].slen[it[1]]), int(self.sources[it[0]].tlen[it[1]])))
batches, cur, ms, mt = [], [], 0, 0
for k, i in items:
s, t = int(self.sources[k].slen[i]), int(self.sources[k].tlen[i])
if cur and (len(cur) + 1) * (max(ms, s) + max(mt, t)) > self.max_tokens:
batches.append(cur)
cur, ms, mt = [], 0, 0
cur.append((k, i))
ms, mt = max(ms, s), max(mt, t)
if cur:
batches.append(cur)
order = rng.permutation(len(batches))
self._batches = [batches[j] for j in order]
self.state["rng"] = rng.bit_generator.state
self.state["consumed"] = 0
def resume(self, saved: dict):
"""saved = pool-start state + batches consumed from that pool."""
self.state = json.loads(json.dumps(saved["pool_state"]))
self._build_pool()
self._batches = self._batches[saved["consumed"]:]
self.state["consumed"] = saved["consumed"]
def checkpoint(self) -> dict:
return {"pool_state": self._pool_state, "consumed": self.state["consumed"]}
def next(self):
if not self._batches:
self._build_pool()
b = self._batches.pop(0)
self.state["consumed"] += 1
return [self.sources[k].get(i) for k, i in b], Counter(self.sources[k].name for k, _ in b)
def collate(rows, device):
b = len(rows)
ls, lt = max(len(s) for s, _ in rows), max(len(t) for _, t in rows)
src = np.full((b, ls), MB_PAD, dtype=np.int64)
tgt = np.full((b, lt), PAD, dtype=np.int64)
for j, (s, t) in enumerate(rows):
src[j, :len(s)] = s
tgt[j, :len(t)] = t
src = torch.from_numpy(src).to(device)
mask = torch.from_numpy(np.arange(ls)[None] < np.array([len(s) for s, _ in rows])[:, None]).to(device)
return src, mask, torch.from_numpy(tgt).to(device)
def load_eval(path: Path, sp, lut, n: int | None):
rows = [json.loads(x) for x in path.read_text(encoding="utf-8").splitlines()]
rows = rows[:n] if n else rows
for r in rows:
r["tgt"] = lut[np.asarray([BOS] + sp.encode(r["reference"]) + [EOS], dtype=np.int64)]
r["src"] = np.asarray(r["src"], dtype=np.int64)
return rows
def decode_ids(ids, dat_ids, sp):
return sp.decode([int(dat_ids[i]) for i in ids if dat_ids[i] >= 4])
@torch.no_grad()
def evaluate(model, rows, sp, dat_ids, device, amp, bsz=64, byte_compact=None):
model.eval()
tot, cnt, hyps = 0.0, 0, []
order = sorted(range(len(rows)), key=lambda i: len(rows[i]["src"]))
for j in range(0, len(order), bsz):
part = [rows[i] for i in order[j:j + bsz]]
src, mask, tgt = collate([(r["src"], r["tgt"]) for r in part], device)
with amp():
_, nll, n = model.loss(src, mask, tgt, PAD, 0.0)
out = model.generate(src, mask, BOS, EOS, PAD, max_len=3 * src.shape[1] + 10, byte_compact=byte_compact)
tot += float(nll) * int(n)
cnt += int(n)
hyps += list(zip([order[k] for k in range(j, j + len(part))], out))
hyps = [decode_ids(o, dat_ids, sp) for _, o in sorted(hyps)]
model.train()
return tot / cnt, chrf(hyps, [r["reference"] for r in rows]), hyps
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--data", type=Path, default=ROOT / "artifacts/ar_mt_20261005/data")
ap.add_argument("--mix", default="luna:1")
ap.add_argument("--encoder", default=str(ROOT / "models/modernbert_ja/modernbert-ja-310m"))
ap.add_argument("--out", type=Path, required=True)
ap.add_argument("--dec-layers", type=int, default=2)
ap.add_argument("--ffn", type=int, default=4096)
ap.add_argument("--bridge", default="swiglu", choices=["swiglu", "mlp", "linear"])
ap.add_argument("--bridge-hidden", type=int, default=2048)
ap.add_argument("--fusion", type=int, default=1)
ap.add_argument("--null-tokens", type=int, default=2)
ap.add_argument("--max-tgt", type=int, default=512)
ap.add_argument("--dropout", type=float, default=0.1)
ap.add_argument("--smoothing", type=float, default=0.1)
ap.add_argument("--init-emb", type=int, default=1)
ap.add_argument("--steps", type=int, required=True)
ap.add_argument("--stage-a", type=int, default=2000)
ap.add_argument("--lr", type=float, default=3e-4)
ap.add_argument("--enc-lr-mult", type=float, default=0.1)
ap.add_argument("--enc-warmup", type=int, default=1000)
ap.add_argument("--warmup", type=int, default=1000)
ap.add_argument("--final-lr-frac", type=float, default=0.1)
ap.add_argument("--wd", type=float, default=0.01)
ap.add_argument("--clip", type=float, default=1.0)
ap.add_argument("--max-tokens", type=int, default=16384)
ap.add_argument("--accum", type=int, default=1)
ap.add_argument("--pool", type=int, default=65536)
ap.add_argument("--seed", type=int, default=20261005)
ap.add_argument("--bf16", type=int, default=1)
ap.add_argument("--log-every", type=int, default=50)
ap.add_argument("--eval-every", type=int, default=2000)
ap.add_argument("--eval-n", type=int, default=500)
ap.add_argument("--save-minutes", type=float, default=20)
ap.add_argument("--resume", action="store_true")
ap.add_argument("--max-steps-this-run", type=int, default=0, help="stop early (smoke tests)")
ap.add_argument("--stop-at", type=int, default=0, help="stop at this absolute step; the LR schedule still uses --steps")
args = ap.parse_args()
args.out.mkdir(parents=True, exist_ok=True)
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
if device == "cuda":
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
use_bf16 = bool(args.bf16) and device == "cuda"
amp = (lambda: torch.autocast("cuda", dtype=torch.bfloat16)) if use_bf16 else (lambda: torch.autocast(device, enabled=False))
torch.manual_seed(args.seed)
random.seed(args.seed)
vocab = json.loads((args.data / "vocab.json").read_text())
dat_ids = vocab["dat_ids"]
lut = np.zeros(vocab["dat_vocab"], dtype=np.int64) # DAT id -> compact id (unknown -> UNK)
for i, d in enumerate(dat_ids):
if d >= 0:
lut[d] = i
sp = spm.SentencePieceProcessor(model_file=str(DAT_TOK))
byte_compact = [int(lut[x]) for x in vocab["byte_piece_ids"]] if "byte_piece_ids" in vocab else None
cfg = {"encoder": args.encoder, "dec_layers": args.dec_layers, "ffn": args.ffn, "bridge": args.bridge,
"bridge_hidden": args.bridge_hidden, "fusion": bool(args.fusion), "null_tokens": args.null_tokens,
"vocab": len(dat_ids), "max_tgt": args.max_tgt, "dropout": args.dropout}
model = ARMT(cfg)
if args.init_emb and not args.resume and model.d == 768:
table = (torch.load(DAT_SHARED, map_location="cpu")["shared.weight"] if DAT_SHARED.exists() else
torch.load(DAT_WEIGHTS, map_location="cpu", weights_only=False)["model"]["shared.weight"]).float()
print(json.dumps({"emb_rows_from_dat": model.init_embeddings(table, dat_ids)}), flush=True)
del table
model.to(device).train()
enc_params = list(model.encoder.parameters())
dec_named = [(n, p) for n, p in model.named_parameters() if not n.startswith("encoder.")]
# no weight decay on embeddings, positions, null memory, norms, and the fusion scales/logits (gamma, logits)
nd = lambda n, p: p.ndim < 2 or n.startswith(("emb.", "pos.", "null", "fusion.gamma", "fusion.logits")) # noqa: E731
decay = [p for n, p in dec_named if not nd(n, p)]
no_decay = [p for n, p in dec_named if nd(n, p)]
enc_decay = [p for p in enc_params if p.ndim >= 2]
enc_no = [p for p in enc_params if p.ndim < 2]
opt = torch.optim.AdamW([{"params": decay, "weight_decay": args.wd, "kind": "dec"},
{"params": no_decay, "weight_decay": 0.0, "kind": "dec"},
{"params": enc_decay, "weight_decay": args.wd, "kind": "enc"},
{"params": enc_no, "weight_decay": 0.0, "kind": "enc"}],
lr=args.lr, betas=(0.9, 0.98), eps=1e-8, fused=device == "cuda")
mix = [(s.split(":")[0], float(s.split(":")[1])) for s in args.mix.split(",")]
sources = [Source(args.data, name, lut) for name, _ in mix]
sampler = Sampler(sources, [w for _, w in mix], args.seed, args.max_tokens, args.pool)
dev = load_eval(args.data / "eval_dev.jsonl", sp, lut, args.eval_n)
step, best = 0, -1.0
if args.resume and (args.out / "last.pt").exists():
ck = torch.load(args.out / "last.pt", map_location="cpu", weights_only=False)
model.load_state_dict(ck["model"])
opt.load_state_dict(ck["opt"])
sampler.resume(ck["sampler"])
step, best = ck["step"], ck["best"]
torch.set_rng_state(ck["torch_rng"])
print(json.dumps({"resumed_step": step}), flush=True)
del ck
(args.out / "config.json").write_text(json.dumps({"cfg": cfg, "args": {k: str(v) for k, v in vars(args).items()},
"params_total": sum(p.numel() for p in model.parameters()),
"params_encoder": sum(p.numel() for p in enc_params),
"sources": {s.name: s.n for s in sources}}, indent=1) + "\n")
def save(name: str, full: bool):
tmp = args.out / f"{name}.tmp"
obj = {"model": model.state_dict(), "cfg": cfg, "vocab": str(args.data / "vocab.json"), "step": step, "best": best}
if full:
obj.update(opt=opt.state_dict(), sampler=sampler.checkpoint(), torch_rng=torch.get_rng_state())
torch.save(obj, tmp)
os.replace(tmp, args.out / f"{name}.pt")
log = open(args.out / "train_log.jsonl", "a")
t0 = last_save = time.time()
win = {"loss": 0.0, "nll": 0.0, "tok": 0, "rows": 0, "n": 0, "src": Counter()}
started = step
end = min(args.steps, args.stop_at) if args.stop_at else args.steps
while step < end:
stage_b = step >= args.stage_a
warm = min(1.0, (step + 1) / args.warmup)
prog = min(1.0, step / max(1, args.steps))
lr = args.lr * warm * (args.final_lr_frac + (1 - args.final_lr_frac) * 0.5 * (1 + math.cos(math.pi * prog)))
enc_lr = lr * args.enc_lr_mult * min(1.0, (step - args.stage_a + 1) / args.enc_warmup) if stage_b else 0.0
for g in opt.param_groups:
g["lr"] = lr if g["kind"] == "dec" else enc_lr
for _ in range(args.accum):
rows, names = sampler.next()
src, mask, tgt = collate(rows, device)
with amp():
loss, nll, n = model.loss(src, mask, tgt, PAD, args.smoothing, encoder_grad=stage_b)
(loss / args.accum).backward()
win["loss"] += float(loss.detach())
win["nll"] += float(nll.detach())
win["tok"] += int(n)
win["rows"] += len(rows)
win["n"] += 1
win["src"].update(names)
gn = torch.nn.utils.clip_grad_norm_([p for g in opt.param_groups for p in g["params"]], args.clip)
opt.step()
opt.zero_grad(set_to_none=True)
step += 1
if step % args.log_every == 0 or step == args.steps:
dt = time.time() - t0
rec = {"step": step, "stage": "B" if stage_b else "A", "loss": round(win["loss"] / win["n"], 4),
"nll": round(win["nll"] / win["n"], 4), "lr": lr, "enc_lr": enc_lr, "grad_norm": round(float(gn), 3),
"tgt_tok_s": round(win["tok"] / dt, 1), "rows_s": round(win["rows"] / dt, 1),
"rows_per_step": round(win["rows"] / max(1, win["n"] / args.accum), 1),
"mix": dict(win["src"]), "epochs": dict(zip([s.name for s in sources], sampler.state["epoch"]))}
if model.fusion is not None:
rec["gamma_abs"] = round(float(model.fusion.gamma.abs().mean()), 5)
a = model.fusion.logits.softmax(-1)
rec["depth_entropy"] = [round(float(-(r * r.log()).sum()), 3) for r in a]
if device == "cuda":
rec["peak_gib"] = round(torch.cuda.max_memory_allocated() / 2 ** 30, 2)
log.write(json.dumps(rec) + "\n")
log.flush()
print(json.dumps(rec), flush=True)
t0, win = time.time(), {"loss": 0.0, "nll": 0.0, "tok": 0, "rows": 0, "n": 0, "src": Counter()}
if step % args.eval_every == 0 or step == end:
dev_nll, dev_chrf, hyps = evaluate(model, dev, sp, dat_ids, device, amp, byte_compact=byte_compact)
rec = {"step": step, "dev_nll": round(dev_nll, 4), "dev_chrf": round(dev_chrf, 3), "examples": hyps[:3]}
log.write(json.dumps(rec, ensure_ascii=False) + "\n")
log.flush()
print(json.dumps(rec, ensure_ascii=False), flush=True)
if dev_chrf > best:
best = dev_chrf
save("best", full=False)
t0 = time.time()
if time.time() - last_save > 60 * args.save_minutes or step == end:
save("last", full=True)
last_save = time.time()
if args.max_steps_this_run and step - started >= args.max_steps_this_run:
save("last", full=True)
break
(args.out / "done.json").write_text(json.dumps({"step": step, "best_dev_chrf": best}) + "\n")
if __name__ == "__main__":
main()