Translation
MLX
Core ML
ONNX
Safetensors
Japanese
Chinese
jmangatranslator-fast
manga
japanese
chinese
Instructions to use muscgab/JMangaTranslator-Fast with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use muscgab/JMangaTranslator-Fast with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] hf download muscgab/JMangaTranslator-Fast --local-dir JMangaTranslator-Fast
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download src/train.py from muscgab/JMangaTranslator-Fast: direct link, hf CLI and curl.
- Browser
- Download file 18.3 kB
-
https://huggingface.co/muscgab/JMangaTranslator-Fast/resolve/main/src/train.py
- Command line
-
hf download hf://muscgab/JMangaTranslator-Fast/src/train.py
-
curl -L -o train.py https://huggingface.co/muscgab/JMangaTranslator-Fast/resolve/main/src/train.py
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]) | |
| 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() | |