File size: 28,268 Bytes
3d92ad2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
#!/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()