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/export_coreml.py from muscgab/JMangaTranslator-Fast: direct link, hf CLI and curl.
- Browser
- Download file 28.3 kB
-
https://huggingface.co/muscgab/JMangaTranslator-Fast/resolve/main/src/export_coreml.py
- Command line
-
hf download hf://muscgab/JMangaTranslator-Fast/src/export_coreml.py
-
curl -L -o export_coreml.py https://huggingface.co/muscgab/JMangaTranslator-Fast/resolve/main/src/export_coreml.py
28.3 kB
| #!/usr/bin/env python3 | |
| """Core ML / ANE export of the single-block ARMT (R2 54k): an encoder model and a one-step decoder model, plus a | |
| host-side greedy loop identical to ARMT.generate_ctx with an empty prefix (kana byte rules included). | |
| encoder ids [1, L] int32 (L in --buckets, padded with pad id 3; mask derived inside as ids != 3) | |
| -> ck0, cv0, ck1, cv1 [1, H, 2 + L, 64] cross-attention K/V of both decoder layers | |
| (ModernBERT with eager attention, bidirectional sliding window |i - j| <= 64 on sliding layers, RoPE tables | |
| as constants; bridge, gamma depth fusion, null tokens, kv2 of each decoder layer). | |
| decoder x [1, 1, d] (host: emb[tok] * sqrt(d) + pos[t]), self K/V caches [1, H, T, 64] x 2 layers with additive mask | |
| smask [1, 1, 1, T] (positions < t), cross K/V padded to M = 2 + max bucket with additive cmask [1, 1, 1, M] | |
| -> logits [1, 1, V], k0, v0, k1, v1 [1, H, 1, 64] (host writes them into the caches at t). | |
| fp16 every LayerNorm / RMSNorm divides its input by a calibrated per-norm power of two s and uses eps / s^2 (the | |
| same function): x^2 stays below the fp16 limit where |h| reaches ~2,865 (encoder layers 14-24), while small | |
| inputs keep s = 1 (a global s = 64 underflowed x^2 at the embedding norm: 10.6 % error, 2026-10-07). | |
| Modes: | |
| check torch fp32: export modules vs ARMT (memories / logits) and host loop vs Ours.translate on --n boxes | |
| convert write <out>/{encoder,decoder}_<prec>.mlpackage | |
| eval Core ML host loop on Manga109 clean (--n boxes): agreement with ARMT outputs, chrF, latency per box | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import math | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from torch import nn | |
| # coremltools imports tensorflow when it is installed; TF's native library deadlocks / aborts on an absl mutex next to | |
| # sentencepiece / tokenizers (2026-10-07). The export does not need TF, so hide it before coremltools is imported. | |
| sys.modules.setdefault("tensorflow", None) | |
| try: | |
| import coremltools # noqa: F401 | |
| except ImportError: | |
| pass | |
| from torch.nn import functional as F | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT / "ar_mt")) | |
| sys.path.insert(0, str(ROOT / "benchmarks/sakura_cmp")) | |
| from model import ARMT # noqa: E402 | |
| from train import BOS, EOS, PAD, decode_ids # noqa: E402 | |
| NEG = -1e4 | |
| CALIB: dict | None = None # module -> max |input| while calibrating (torch fp32 only, never while tracing) | |
| def norm_scaled(x, mod, eps: float, center: bool): | |
| """LayerNorm (center=True, no bias) / RMSNorm of x with weight mod.weight, computed on x / s with eps / s^2 (the | |
| same function). s = mod._s is a per-norm power of two from calibrate(): 1 where inputs are small (dividing them | |
| would underflow x^2 in fp16), larger where the residual stream carries massive activations.""" | |
| if CALIB is not None: | |
| CALIB[mod] = max(CALIB.get(mod, 0.0), float(x.detach().abs().max())) | |
| s = getattr(mod, "_s", 1.0) | |
| w = mod.weight | |
| x = x * (1.0 / s) | |
| if center: | |
| x = x - x.mean(-1, keepdim=True) | |
| return x * torch.rsqrt((x * x).mean(-1, keepdim=True) + eps / (s * s)) * w | |
| def rope(x, cos, sin): | |
| x1, x2 = x.chunk(2, -1) # rotate_half without shape-derived ints (coremltools aten::Int) | |
| return x * cos + torch.cat((-x2, x1), -1) * sin | |
| class EncoderExport(nn.Module): | |
| def __init__(self, m: ARMT, lmax: int, s: float): | |
| super().__init__() | |
| enc, c = m.encoder, m.encoder.config | |
| self.c, self.s, self.lmax = c, s, lmax | |
| self.tok = enc.embeddings.tok_embeddings | |
| self.emb_norm = enc.embeddings.norm | |
| self.layers = enc.layers | |
| self.final_norm = enc.final_norm | |
| self.H, self.hd = c.num_attention_heads, c.hidden_size // c.num_attention_heads | |
| for kind in ("full_attention", "sliding_attention"): # no length-dependent slicing: RoPE angles | |
| theta = c.rope_parameters[kind]["rope_theta"] # and the band mask come from positions | |
| inv = 1.0 / theta ** (torch.arange(0, self.hd, 2, dtype=torch.float32) / self.hd) | |
| ang = torch.arange(lmax, dtype=torch.float32)[:, None] * torch.cat((inv, inv))[None] # fp32 table: | |
| self.register_buffer(f"cos_{kind}", ang.cos(), persistent=False) # angles up to ~lmax rad would | |
| self.register_buffer(f"sin_{kind}", ang.sin(), persistent=False) # lose ~0.06 rad in fp16 | |
| self.bridge, self.fusion, self.null = m.bridge, m.fusion, m.null | |
| self.register_buffer("fw", m.fusion.logits.detach().softmax(-1), persistent=False) # [J, D] | |
| self.kv2 = nn.ModuleList(layer.kv2 for layer in m.layers) | |
| self.eps = c.norm_eps | |
| def ln(self, x, mod): | |
| return norm_scaled(x, mod, self.eps, True) | |
| def rms(self, x, mod): | |
| return norm_scaled(x, mod, mod.eps, False) | |
| def forward(self, ids): | |
| h = self.ln(self.tok(ids.long()), self.emb_norm) | |
| valid = (ids != 3).to(h.dtype) # [1, L] | |
| pos = torch.cumsum(torch.ones_like(valid), 1)[0] - 1.0 # [L] = 0 .. L-1 | |
| key = (1.0 - valid)[:, None, None, :] * NEG | |
| far = ((pos[:, None] - pos[None, :]).abs() > self.c.sliding_window).to(h.dtype) * NEG | |
| masks = {"full_attention": key, "sliding_attention": key + far[None, None]} | |
| trig = {} | |
| for kind in ("full_attention", "sliding_attention"): | |
| pi = pos.long() # table lookup by position | |
| trig[kind] = (F.embedding(pi, getattr(self, f"cos_{kind}")).to(h.dtype), | |
| F.embedding(pi, getattr(self, f"sin_{kind}")).to(h.dtype)) | |
| return self.body(h, masks, trig) | |
| def body(self, h, masks, trig): | |
| states = [h] | |
| for i, layer in enumerate(self.layers): | |
| kind = layer.attention_type | |
| a = h if i == 0 else self.ln(h, layer.attn_norm) | |
| qkv = layer.attn.Wqkv(a).view(1, -1, 3, self.H, self.hd) | |
| q, k, v = (qkv[:, :, j].transpose(1, 2) for j in range(3)) | |
| cos, sin = trig[kind] | |
| q, k = rope(q, cos, sin), rope(k, cos, sin) | |
| p = torch.softmax(q @ k.transpose(2, 3) * self.hd ** -0.5 + masks[kind], -1) | |
| h = h + layer.attn.Wo((p @ v).transpose(1, 2).reshape(1, -1, self.H * self.hd)) | |
| x1, x2 = layer.mlp.Wi(self.ln(h, layer.mlp_norm)).chunk(2, -1) | |
| h = h + layer.mlp.Wo(F.gelu(x1) * x2) | |
| states.append(h) | |
| last = self.ln(h, self.final_norm) | |
| b = self.bridge # RMSNorm, SwiGLU, RMSNorm | |
| base = self.rms(b[1](self.rms(last, b[0])), b[2]) # [1, L, d] | |
| normed = torch.stack([self.rms(st, n) for st, n in zip(states[:len(self.fusion.norms)], self.fusion.norms)]) | |
| out = [] | |
| for j, kv2 in enumerate(self.kv2): | |
| f = (self.fw[j][:, None, None, None] * normed).sum(0) # weighted depth sum | |
| mem = base + self.fusion.gamma[j] * self.fusion.wo(f) | |
| mem = torch.cat((self.null[None], mem), 1) # [1, 2 + L, d] | |
| k, v = kv2(mem).chunk(2, -1) | |
| out += [k.view(1, -1, self.H, self.hd).transpose(1, 2), v.view(1, -1, self.H, self.hd).transpose(1, 2)] | |
| return tuple(out) | |
| class EncoderStatic(EncoderExport): | |
| """All-ANE encoder for one fixed length L: the token-embedding lookup moves to the host (input x = raw token | |
| embeddings [1, L, E] before the embedding LayerNorm), the padding mask is an input (kmask [1, 1, 1, L], 0 / NEG), | |
| RoPE tables and the sliding-window band are constants of this L. No gather / cast / comparison left in the graph.""" | |
| def __init__(self, m: ARMT, L: int): | |
| super().__init__(m, L, 1.0) | |
| pos = torch.arange(L, dtype=torch.float32) | |
| self.register_buffer("far", torch.where((pos[:, None] - pos[None, :]).abs() > self.c.sliding_window, NEG, 0.0) | |
| [None, None], persistent=False) # [1, 1, L, L] | |
| def forward(self, x, kmask): | |
| h = self.ln(x, self.emb_norm) | |
| masks = {"full_attention": kmask, "sliding_attention": kmask + self.far} | |
| trig = {k: (getattr(self, f"cos_{k}"), getattr(self, f"sin_{k}")) for k in masks} | |
| return self.body(h, masks, trig) | |
| class DecoderExport(nn.Module): | |
| def __init__(self, m: ARMT, s: float): | |
| super().__init__() | |
| self.layers, self.norm, self.s = m.layers, m.norm, s | |
| self.register_buffer("emb_t", m.emb.weight.detach().T.contiguous(), persistent=False) # [d, V] | |
| self.H, self.hd = m.layers[0].h, m.layers[0].hd | |
| def rms(self, x, mod): | |
| return norm_scaled(x, mod, mod.eps, False) | |
| def heads(self, x): | |
| return x.view(1, 1, self.H, self.hd).transpose(1, 2) | |
| def forward(self, x, kc0, vc0, kc1, vc1, smask, ck0, cv0, ck1, cv1, cmask): | |
| news = [] | |
| zero = torch.zeros_like(smask[..., :1]) | |
| for layer, kc, vc, ck, cv in zip(self.layers, (kc0, kc1), (vc0, vc1), (ck0, ck1), (cv0, cv1)): | |
| q, k, v = layer.qkv(self.rms(x, layer.n1)).chunk(3, -1) | |
| q, k, v = self.heads(q), self.heads(k), self.heads(v) | |
| K, V = torch.cat((kc, k), 2), torch.cat((vc, v), 2) | |
| p = torch.softmax(q @ K.transpose(2, 3) * self.hd ** -0.5 + torch.cat((smask, zero), -1), -1) | |
| x = x + layer.o1((p @ V).transpose(1, 2).reshape(1, 1, -1)) | |
| q2 = self.heads(layer.q2(self.rms(x, layer.n2))) | |
| p = torch.softmax(q2 @ ck.transpose(2, 3) * self.hd ** -0.5 + cmask, -1) | |
| x = x + layer.o2((p @ cv).transpose(1, 2).reshape(1, 1, -1)) | |
| x = x + layer.ffn(self.rms(x, layer.n3)) | |
| news += [k, v] | |
| return (self.rms(x, self.norm) @ self.emb_t, *news) | |
| class Host: | |
| """Greedy loop of ARMT.generate_ctx (empty prefix) around an encoder / decoder step backend.""" | |
| def __init__(self, m: ARMT, vocab: dict, tok, sp, buckets, T: int, enc_fn, dec_fn): | |
| self.tok, self.sp, self.buckets, self.T = tok, sp, buckets, T | |
| self.enc_fn, self.dec_fn = enc_fn, dec_fn | |
| self.d = m.d | |
| self.emb = m.emb.weight.detach().float().numpy() * math.sqrt(m.d) | |
| self.pos = m.pos.weight.detach().float().numpy() | |
| self.dat_ids = vocab["dat_ids"] | |
| lut = np.zeros(vocab["dat_vocab"], dtype=np.int64) | |
| for i, dd in enumerate(self.dat_ids): | |
| if dd >= 0: | |
| lut[dd] = i | |
| bc = [int(lut[x]) for x in vocab["byte_piece_ids"]] | |
| self.rules = [(p2, p1, mk.numpy()) for p2, p1, mk in ARMT.kana_byte_rules(bc, self.emb.shape[0], "cpu")] | |
| self.M = 2 + max(buckets) | |
| self.H, self.hd = m.layers[0].h, m.layers[0].hd | |
| def __call__(self, text: str): | |
| ids = self.tok(text, add_special_tokens=True, truncation=True, max_length=256)["input_ids"] | |
| n = len(ids) | |
| Lb = next((b for b in self.buckets if b >= n), None) | |
| if Lb is None: # longer than the largest bucket: keep the head | |
| ids, n, Lb = ids[:self.buckets[-1]], self.buckets[-1], self.buckets[-1] | |
| arr = np.full((1, Lb), 3, dtype=np.int32) | |
| arr[0, :n] = ids | |
| t0 = time.perf_counter() | |
| cross = self.enc_fn(arr) # 4 x [1, H, 2 + Lb, hd] | |
| t_enc = time.perf_counter() - t0 | |
| cr = [] | |
| for c in cross: | |
| z = np.zeros((1, self.H, self.M, self.hd), dtype=np.float32) | |
| z[:, :, :c.shape[2]] = c | |
| cr.append(z) | |
| cmask = np.full((1, 1, 1, self.M), NEG, dtype=np.float32) | |
| cmask[..., :2 + n] = 0.0 | |
| caches = [np.zeros((1, self.H, self.T, self.hd), dtype=np.float32) for _ in range(4)] | |
| smask = np.full((1, 1, 1, self.T), NEG, dtype=np.float32) | |
| steps = min(3 * n + 10, 512, self.T + 1) | |
| prev1 = prev2 = -1 | |
| tok, out = BOS, [] | |
| for t in range(steps): | |
| x = (self.emb[tok] + self.pos[t])[None, None].astype(np.float32) | |
| logits, *news = self.dec_fn(x, caches, smask, cr, cmask) | |
| lg = logits.reshape(-1).astype(np.float32) | |
| for p2, p1, mk in self.rules: | |
| if prev1 == p1 and (p2 is None or prev2 == p2): | |
| lg[mk] = -np.inf | |
| nxt = int(lg.argmax()) | |
| if nxt == EOS or nxt == PAD: | |
| break | |
| out.append(nxt) | |
| if t < self.T: | |
| for c, nw in zip(caches, news): | |
| c[:, :, t] = nw[:, :, 0] | |
| smask[..., t] = 0.0 | |
| prev2, prev1, tok = prev1, nxt, nxt | |
| return {"hyp": decode_ids(out, self.dat_ids, self.sp), "ntok": n, "nout": len(out), | |
| "enc_ms": round(1000 * t_enc, 2), "ms": round(1000 * (time.perf_counter() - t0), 2)} | |
| def load(args): | |
| from run_ours import Ours | |
| o = Ours(args.ckpt, args.encoder, args.data, "cpu") | |
| o.model.float().eval() | |
| return o | |
| def torch_backends(o, enc, dec): | |
| def enc_fn(arr): | |
| with torch.no_grad(): | |
| return [t.numpy() for t in enc(torch.from_numpy(arr))] | |
| def dec_fn(x, caches, smask, cr, cmask): | |
| with torch.no_grad(): | |
| r = dec(torch.from_numpy(x), *map(torch.from_numpy, caches), torch.from_numpy(smask), | |
| *map(torch.from_numpy, cr), torch.from_numpy(cmask)) | |
| return [t.numpy() for t in r] | |
| return enc_fn, dec_fn | |
| def calibrate(m, o, enc, dec, vocab, buckets, T, texts, cap: float, path: Path): | |
| """Per-norm scales from the max |input| seen on texts (torch fp32, encoder + host decoding loop): | |
| s = 2^ceil(log2(max(1, absmax / cap))). Saved to path (module name -> s, absmax) and set as mod._s.""" | |
| global CALIB | |
| names = {mod: n for n, mod in m.named_modules()} | |
| if path.exists(): | |
| saved = json.loads(path.read_text()) | |
| for mod, n in names.items(): | |
| if n in saved: | |
| mod._s = saved[n]["s"] | |
| return saved | |
| CALIB = {} | |
| host = Host(m, vocab, o.tok, o.sp, buckets, T, *torch_backends(o, enc, dec)) | |
| for t in texts: | |
| host(t) | |
| seen, CALIB = CALIB, None | |
| out = {} | |
| for mod, a in seen.items(): | |
| mod._s = float(2 ** math.ceil(math.log2(max(1.0, a / cap)))) | |
| out[names[mod]] = {"s": mod._s, "absmax": round(a, 2)} | |
| path.write_text(json.dumps(out, indent=1) + "\n") | |
| return out | |
| def coreml_backends(out: Path, prec: str, units: str, dprec: str | None = None, dunits: str | None = None, | |
| static: bool = False, table=None, buckets=()): | |
| import coremltools as ct | |
| if static: # all-ANE encoder: one function per bucket, embedding lookup + padding mask on host | |
| fm = {b: ct.models.MLModel(str(out / f"encoder_static_{prec}.mlpackage"), function_name=f"L{b}", | |
| compute_units=getattr(ct.ComputeUnit, units)) for b in buckets} | |
| else: | |
| em = ct.models.MLModel(str(out / f"encoder_{prec}.mlpackage"), compute_units=getattr(ct.ComputeUnit, units)) | |
| dm = ct.models.MLModel(str(out / f"decoder_{dprec or prec}.mlpackage"), | |
| compute_units=getattr(ct.ComputeUnit, dunits or units)) | |
| names = ("ck0", "cv0", "ck1", "cv1") | |
| def enc_fn(arr): | |
| if static: | |
| L = arr.shape[1] | |
| km = np.where(arr == 3, NEG, 0.0).astype(np.float32)[:, None, None, :] | |
| r = fm[L].predict({"x": table[arr], "kmask": km}) | |
| else: | |
| r = em.predict({"ids": arr}) | |
| return [r[k] for k in names] | |
| def dec_fn(x, caches, smask, cr, cmask): | |
| feed = {"x": x, "kc0": caches[0], "vc0": caches[1], "kc1": caches[2], "vc1": caches[3], "smask": smask, | |
| "ck0": cr[0], "cv0": cr[1], "ck1": cr[2], "cv1": cr[3], "cmask": cmask} | |
| r = dm.predict(feed) | |
| return [r["logits"], r["k0"], r["v0"], r["k1"], r["v1"]] | |
| return enc_fn, dec_fn | |
| def main() -> None: | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("mode", choices=("check", "convert", "eval", "encerr")) | |
| ap.add_argument("--ckpt", type=Path, default=ROOT / "artifacts/r2_20261006/weights/r2_step54000.pt") | |
| ap.add_argument("--encoder", default=str(ROOT / "artifacts/ar_mt_20261005/weights/dapt_v1_model")) | |
| ap.add_argument("--data", type=Path, default=ROOT / "artifacts/ar_mt_20261005/data_v2_ext") | |
| ap.add_argument("--out", type=Path, default=ROOT / "artifacts/coreml_20261007") | |
| ap.add_argument("--buckets", default="32,64,128") | |
| ap.add_argument("--T", type=int, default=128, help="self-attention cache length (max output tokens ~ T + 1)") | |
| ap.add_argument("--norm-cap", type=float, default=32.0, help="calibrated norm scale keeps max|x|/s <= cap") | |
| ap.add_argument("--prec", choices=("fp32", "fp16"), default="fp16") | |
| ap.add_argument("--units", default="CPU_AND_NE", help="ALL / CPU_ONLY / CPU_AND_GPU / CPU_AND_NE") | |
| ap.add_argument("--dec-prec", choices=("fp32", "fp16"), help="decoder precision (default --prec)") | |
| ap.add_argument("--dec-units", help="decoder compute units (default --units)") | |
| ap.add_argument("--backend", choices=("torch", "coreml"), default="coreml") | |
| ap.add_argument("--static", type=int, default=0, help="1 = all-ANE encoder (encoder_static_<prec>.mlpackage)") | |
| ap.add_argument("--n", type=int, default=200) | |
| ap.add_argument("--tag", default="") | |
| args = ap.parse_args() | |
| torch.set_num_threads(4) | |
| buckets = [int(b) for b in args.buckets.split(",")] | |
| o = load(args) | |
| m = o.model | |
| enc = EncoderExport(m, max(buckets), 1.0).eval() | |
| dec = DecoderExport(m, 1.0).eval() | |
| vocab = json.loads((args.data / "vocab.json").read_text()) | |
| args.out.mkdir(parents=True, exist_ok=True) | |
| from common import load_m109, load_murasaki | |
| recs = load_m109() | |
| # calibration texts: Manga109 boxes 200-999 (evaluation uses the first 200) + 300 Murasaki segments (longer) | |
| ctexts = [r["clean"] for r in recs[200:1000]] + [s_ for r in load_murasaki() for s_ in r["segs"]][:300] | |
| scales = calibrate(m, o, enc, dec, vocab, buckets, args.T, ctexts, args.norm_cap, args.out / "norm_scales.json") | |
| print(json.dumps({"norm_scales": {str(k): v for k, v in sorted( | |
| __import__("collections").Counter(x["s"] for x in scales.values()).items())}}), flush=True) | |
| if args.mode == "check": | |
| # 1) memories: export encoder vs ARMT._encode (+ null, kv2), on real boxes padded to a bucket | |
| worst = {} | |
| for r in recs[:args.n]: | |
| ids = o.tok(r["clean"], add_special_tokens=True, truncation=True, max_length=256)["input_ids"] | |
| n = len(ids) | |
| Lb = next(b for b in buckets if b >= n) | |
| arr = torch.full((1, Lb), 3, dtype=torch.int32) | |
| arr[0, :n] = torch.tensor(ids) | |
| with torch.no_grad(): | |
| got = enc(arr) | |
| s = torch.tensor([ids]) | |
| mems, _ = m.memories(s, torch.ones_like(s, dtype=torch.bool), False) | |
| ref = [t for layer, mem in zip(m.layers, mems) for t in layer.cross_kv(mem)] | |
| for name, a, b in zip(("ck0", "cv0", "ck1", "cv1"), got, ref): | |
| e = float((a[:, :, :n + 2] - b).abs().max() / b.abs().max()) | |
| worst[name] = max(worst.get(name, 0.0), e) | |
| print(json.dumps({"check": "encoder rel max err over valid positions", **worst}), flush=True) | |
| # 2) full host loop (torch backends) vs Ours.translate | |
| host = Host(m, vocab, o.tok, o.sp, buckets, args.T, *torch_backends(o, enc, dec)) | |
| same = 0 | |
| diffs = [] | |
| for r in recs[:args.n]: | |
| a = host(r["clean"])["hyp"] | |
| b = o.translate([r["clean"]])[0] | |
| same += a == b | |
| if a != b and len(diffs) < 5: | |
| diffs.append([r["clean"], a, b]) | |
| print(json.dumps({"check": "host loop vs ARMT", "n": args.n, "identical": same, "diffs": diffs}, | |
| ensure_ascii=False), flush=True) | |
| # 3) fp16 simulation in torch (CPU): any inf / nan in the encoder outputs? | |
| import copy | |
| enc16 = copy.deepcopy(enc).half().eval() | |
| bad = 0 | |
| for r in recs[:min(args.n, 50)]: | |
| ids = o.tok(r["clean"], add_special_tokens=True, truncation=True, max_length=256)["input_ids"] | |
| arr = torch.full((1, next(b for b in buckets if b >= len(ids))), 3, dtype=torch.int32) | |
| arr[0, :len(ids)] = torch.tensor(ids) | |
| with torch.no_grad(): | |
| bad += any(not torch.isfinite(t).all() for t in enc16(arr)) | |
| print(json.dumps({"check": "torch fp16 encoder non-finite outputs", "boxes": min(args.n, 50), "bad": bad}), | |
| flush=True) | |
| return | |
| if args.mode == "encerr": | |
| # encoder outputs of the Core ML model (--prec / --units) and of torch fp16 (CPU) vs torch fp32, valid positions | |
| enc_fn, _ = coreml_backends(args.out, args.prec, args.units) | |
| import copy | |
| enc16 = copy.deepcopy(enc).half().eval() | |
| err = {"coreml": [], "torch_fp16": []} | |
| for r in recs[:args.n]: | |
| ids = o.tok(r["clean"], add_special_tokens=True, truncation=True, max_length=256)["input_ids"] | |
| n = len(ids) | |
| arr = np.full((1, next(b for b in buckets if b >= n)), 3, dtype=np.int32) | |
| arr[0, :n] = ids | |
| with torch.no_grad(): | |
| ref = [t[:, :, :n + 2].numpy() for t in enc(torch.from_numpy(arr))] | |
| t16 = [t[:, :, :n + 2].float().numpy() for t in enc16(torch.from_numpy(arr))] | |
| cm = [c[:, :, :n + 2].astype(np.float32) for c in enc_fn(arr)] | |
| for name, got in (("coreml", cm), ("torch_fp16", t16)): | |
| err[name].append([float(np.linalg.norm(g - r_) / np.linalg.norm(r_)) for g, r_ in zip(got, ref)]) | |
| rep = {k: {"rel_l2_mean": np.round(np.mean(v, 0), 5).tolist(), "rel_l2_max": np.round(np.max(v, 0), 5).tolist()} | |
| for k, v in err.items()} | |
| print(json.dumps({"encerr": f"{args.prec}_{args.units}", "n": args.n, "outputs": "ck0 cv0 ck1 cv1", **rep}), | |
| flush=True) | |
| return | |
| if args.mode == "convert" and args.static: | |
| import coremltools as ct | |
| prec = ct.precision.FLOAT16 if args.prec == "fp16" else ct.precision.FLOAT32 | |
| E_ = m.encoder.config.hidden_size | |
| desc = ct.utils.MultiFunctionDescriptor() | |
| parts = [] | |
| for b in buckets: | |
| es = EncoderStatic(m, b).eval() | |
| ex = (torch.randn(1, b, E_) * 0.05, torch.zeros(1, 1, 1, b)) | |
| with torch.no_grad(): | |
| ts = torch.jit.trace(es, ex, check_trace=False) | |
| mb = ct.convert(ts, inputs=[ct.TensorType(name="x", shape=(1, b, E_)), ct.TensorType(name="kmask", shape=(1, 1, 1, b))], | |
| outputs=[ct.TensorType(name=k) for k in ("ck0", "cv0", "ck1", "cv1")], | |
| convert_to="mlprogram", compute_precision=prec, minimum_deployment_target=ct.target.macOS15) | |
| pth = args.out / f"_enc_static_L{b}_{args.prec}.mlpackage" | |
| mb.save(str(pth)) | |
| parts.append(pth) | |
| desc.add_function(str(pth), src_function_name="main", target_function_name=f"L{b}") | |
| desc.default_function_name = f"L{buckets[0]}" | |
| ct.utils.save_multifunction(desc, str(args.out / f"encoder_static_{args.prec}.mlpackage")) | |
| import shutil | |
| for pth in parts: | |
| shutil.rmtree(pth) | |
| print(json.dumps({"converted": "encoder_static", "functions": [f"L{b}" for b in buckets]}), flush=True) | |
| return | |
| if args.mode == "convert": | |
| import coremltools as ct | |
| prec = ct.precision.FLOAT16 if args.prec == "fp16" else ct.precision.FLOAT32 | |
| ex = torch.full((1, buckets[0]), 3, dtype=torch.int32) | |
| ex[0, :5] = torch.tensor([6, 100, 200, 300, 4]) | |
| with torch.no_grad(): | |
| te = torch.jit.trace(enc, ex, check_trace=False) | |
| t0 = time.time() | |
| me = ct.convert(te, inputs=[ct.TensorType(name="ids", shape=ct.EnumeratedShapes( | |
| shapes=[[1, b] for b in buckets], default=[1, buckets[0]]), dtype=np.int32)], | |
| outputs=[ct.TensorType(name=k) for k in ("ck0", "cv0", "ck1", "cv1")], | |
| convert_to="mlprogram", compute_precision=prec, minimum_deployment_target=ct.target.macOS15) | |
| me.save(str(args.out / f"encoder_{args.prec}.mlpackage")) | |
| print(json.dumps({"converted": "encoder", "secs": round(time.time() - t0, 1)}), flush=True) | |
| H, hd, d, M, T = dec.H, dec.hd, m.d, 2 + max(buckets), args.T | |
| shapes = {"x": (1, 1, d), "kc0": (1, H, T, hd), "vc0": (1, H, T, hd), "kc1": (1, H, T, hd), "vc1": (1, H, T, hd), | |
| "smask": (1, 1, 1, T), "ck0": (1, H, M, hd), "cv0": (1, H, M, hd), "ck1": (1, H, M, hd), | |
| "cv1": (1, H, M, hd), "cmask": (1, 1, 1, M)} | |
| exs = tuple(torch.zeros(s) for s in shapes.values()) | |
| with torch.no_grad(): | |
| td = torch.jit.trace(dec, exs, check_trace=False) | |
| t0 = time.time() | |
| md = ct.convert(td, inputs=[ct.TensorType(name=k, shape=s) for k, s in shapes.items()], | |
| outputs=[ct.TensorType(name=k) for k in ("logits", "k0", "v0", "k1", "v1")], | |
| convert_to="mlprogram", compute_precision=prec, minimum_deployment_target=ct.target.macOS15) | |
| md.save(str(args.out / f"decoder_{args.prec}.mlpackage")) | |
| print(json.dumps({"converted": "decoder", "secs": round(time.time() - t0, 1)}), flush=True) | |
| return | |
| # eval: Manga109 clean, first --n boxes; reference = ARMT (Ours.translate, CPU fp32) | |
| try: | |
| import sacrebleu | |
| except ImportError: # base env has none: chrF is added later (comet env) | |
| sacrebleu = None | |
| from common import M109 | |
| refs = {json.loads(x)["id"]: json.loads(x)["reference"] for x in open(M109 / "refs.jsonl", encoding="utf-8")} | |
| table = m.encoder.embeddings.tok_embeddings.weight.detach().float().numpy() if args.static else None | |
| fns = torch_backends(o, enc, dec) if args.backend == "torch" else \ | |
| coreml_backends(args.out, args.prec, args.units, args.dec_prec, args.dec_units, bool(args.static), table, buckets) | |
| host = Host(m, vocab, o.tok, o.sp, buckets, args.T, *fns) | |
| for r in recs[:10]: # warm-up (model load / ANE compile) | |
| host(r["clean"]) | |
| rows = [] | |
| for r in recs[:args.n]: | |
| x = host(r["clean"]) | |
| x.update(id=r["id"], ref_armt=o.translate([r["clean"]])[0]) | |
| rows.append(x) | |
| tag = args.tag or f"{args.backend}_{args.prec}_{args.units}" + \ | |
| (f"__dec_{args.dec_prec or args.prec}_{args.dec_units or args.units}" if args.dec_prec or args.dec_units else "") | |
| with open(args.out / f"eval_{tag}.jsonl", "w", encoding="utf-8") as f: | |
| for x in rows: | |
| f.write(json.dumps(x, ensure_ascii=False) + "\n") | |
| ms = np.array([x["ms"] for x in rows]) | |
| em = np.array([x["enc_ms"] for x in rows]) | |
| nout = np.array([x["nout"] for x in rows]) | |
| hyp, ref_armt = [x["hyp"] for x in rows], [x["ref_armt"] for x in rows] | |
| gold = [refs[x["id"]] for x in rows] | |
| rep = {"tag": tag, "n": len(rows), "identical_to_armt": round(float(np.mean([a == b for a, b in zip(hyp, ref_armt)])), 4), | |
| "chrf": sacrebleu and round(sacrebleu.corpus_chrf(hyp, [gold]).score, 2), | |
| "chrf_armt": sacrebleu and round(sacrebleu.corpus_chrf(ref_armt, [gold]).score, 2), | |
| "ms_p50": round(float(np.median(ms)), 1), "ms_p90": round(float(np.percentile(ms, 90)), 1), | |
| "enc_ms_p50": round(float(np.median(em)), 1), | |
| "dec_ms_per_step": round(float(((ms - em) / (nout + 1)).mean()), 2), "mean_out_tokens": round(float(nout.mean()), 1), | |
| "bad_outputs": sum(not h.strip() for h in hyp)} | |
| (args.out / f"eval_{tag}.json").write_text(json.dumps(rep, ensure_ascii=False) + "\n") | |
| print(json.dumps(rep, ensure_ascii=False), flush=True) | |
| if __name__ == "__main__": | |
| main() | |