Skip to content

Commit 683337a

Browse files
authored
Merge pull request #35 from interscript/fix/byt5-decode
fix(distill): byt5 decode_joined mojibake — all Arabic labels poisoned
2 parents 37bdf65 + 3af08fb commit 683337a

3 files changed

Lines changed: 175 additions & 12 deletions

File tree

‎docs/RESULTS.md‎

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -121,11 +121,13 @@ reproducing its documented tier on this replication.
121121
| Teacher (r6, run-006-morph) | 1.32% |
122122
| **Student (33M from-scratch)** | **83.08%** — REJECTED |
123123

124-
Gate ≤ teacher + 0.5pp: the student misses by two orders of magnitude
125-
with the same collapse signature as the Thai tiny tier (train loss
126-
converges, test generalization absent). Verdict: sub-100M from-scratch
127-
byte students do not generalize for diacritization any more than for
128-
G2P — a pretrained backbone is non-negotiable. The Arabic client tier
129-
therefore ships at the ByT5-small rung (ara-diac-small) or parks until
130-
byte-level pretraining exists. The 30MB tier is closed as a negative
131-
result across both task families.
124+
Gate ≤ teacher + 0.5pp: the student misses by two orders of magnitude.
125+
126+
**RETRACTION (2026-08-24):** this verdict is CONFOUNDED — every Arabic
127+
label generated before the byt5 `decode_joined` fix was mojibake
128+
(double-encoded targets); both Arabic students trained on corrupted
129+
labels, and their identical DER scores are the bare-text constant, not
130+
a capacity result. The numbers stand as measured but the capacity
131+
conclusion for Arabic is UNPROVEN pending a clean-label re-run. The
132+
Thai tiny verdict is unaffected (umt5/sentencepiece labels were
133+
byte-exact); the pretrained-backbone law rests on Thai evidence.

‎src/api/inference.py‎

Lines changed: 145 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,145 @@
1+
"""IMF v1 inference endpoint for api.interscript.org (Modal, CPU).
2+
3+
Serves the shipped models from the secryst-models volume using the
4+
exact ONNX kv decode the WO03 parity gate verified. Cold start loads
5+
the fp32 zip (~30-60s); sessions are cached per container.
6+
7+
modal deploy src/api/inference.py
8+
9+
Auth: X-API-Key header must match the `api-inference-key` secret.
10+
"""
11+
12+
import hmac
13+
from pathlib import Path
14+
15+
import modal
16+
17+
models_volume = modal.Volume.from_name("secryst-models")
18+
19+
image = (
20+
modal.Image.debian_slim(python_version="3.11")
21+
.pip_install("onnxruntime==1.23.2", "pyyaml>=6.0", "fastapi>=0.115")
22+
.add_local_dir(str(Path(__file__).resolve().parent.parent), "/root/interscript-ml", copy=True)
23+
.workdir("/root/interscript-ml")
24+
.env({"IMAGE_REV": "5"})
25+
)
26+
27+
app = modal.App("interscript-inference", image=image)
28+
29+
MAX_INPUT_BYTES = 4000
30+
MAX_OUTPUT_TOKENS = 8192
31+
ALLOWED_TASKS = ("diacritization", "g2p")
32+
33+
_sessions: dict[str, tuple] = {}
34+
35+
36+
def _zip_path(model_id: str) -> Path:
37+
# model_id arrives from the request body — paths are built ONLY from
38+
# the server's own volume listing; the user value is used solely in
39+
# an equality comparison, never in path construction (CWE-22)
40+
import glob
41+
42+
wanted = f"{model_id}-fp32.zip"
43+
for z in glob.glob("/v/imf/*/*-fp32.zip"):
44+
if z.rsplit("/", 1)[1] == wanted:
45+
return Path(z)
46+
raise KeyError(model_id)
47+
48+
49+
def _get_sessions(model_id: str) -> tuple:
50+
import sys
51+
52+
sys.path.insert(0, "/root/interscript-ml/src")
53+
from imf.parity import _sessions_from_zip # the parity-verified loader
54+
55+
if model_id not in _sessions:
56+
_sessions[model_id] = _sessions_from_zip(_zip_path(model_id))
57+
return _sessions[model_id]
58+
59+
60+
def _metadata(model_id: str) -> dict:
61+
import zipfile
62+
63+
import yaml
64+
65+
with zipfile.ZipFile(_zip_path(model_id)) as zf:
66+
return yaml.safe_load(zf.read("metadata.yaml"))
67+
68+
69+
def _decode(tokens: list) -> str:
70+
# ByT5 token ids are byte+3, with trailing EOS (id 1)
71+
return bytes(t - 3 for t in tokens if t >= 3).decode("utf-8", "replace")
72+
73+
74+
def make_api():
75+
import os
76+
77+
from fastapi import FastAPI, HTTPException, Request
78+
79+
api = FastAPI(title="Interscript inference", version="1.0.0")
80+
81+
@api.post("/infer")
82+
async def infer(request: Request) -> dict:
83+
84+
key = request.headers.get("x-api-key", "")
85+
if not key or not hmac.compare_digest(key, os.environ.get("API_INFERENCE_KEY") or ""):
86+
raise HTTPException(401, "invalid or missing X-API-Key")
87+
try:
88+
body = await request.json()
89+
except Exception:
90+
raise HTTPException(400, "body must be JSON {model, input}") from None
91+
if not isinstance(body, dict) or not body.get("model") or not body.get("input"):
92+
raise HTTPException(400, "body must be {model, input}")
93+
94+
return await _run_infer(body)
95+
96+
async def _run_infer(body):
97+
import sys
98+
99+
models_volume.reload()
100+
try:
101+
meta = _metadata(body["model"])
102+
except KeyError:
103+
raise HTTPException(404, f"unknown model {body['model']}") from None
104+
if meta.get("task") not in ALLOWED_TASKS:
105+
raise HTTPException(400, f"model {body['model']} task {meta.get('task')} is not served")
106+
107+
if len(body["input"].encode("utf-8")) > MAX_INPUT_BYTES:
108+
raise HTTPException(413, f"input exceeds {MAX_INPUT_BYTES} bytes")
109+
110+
sys.path.insert(0, "/root/interscript-ml/src")
111+
from imf.export import onnx_greedy_kv
112+
113+
enc, kv = _get_sessions(body["model"])
114+
max_len = min(MAX_OUTPUT_TOKENS, 3 * len(body["input"].encode("utf-8")) + 256)
115+
output = _decode(onnx_greedy_kv(enc, kv, body["input"], max_len))
116+
return {
117+
"model": body["model"],
118+
"task": meta["task"],
119+
"source_script": meta.get("source_script"),
120+
"input": body["input"],
121+
"output": output,
122+
}
123+
124+
@api.get("/health")
125+
def health() -> dict:
126+
import glob
127+
128+
models_volume.reload()
129+
return {"ok": True, "models": len(glob.glob("/v/imf/*/*-fp32.zip"))}
130+
131+
return api
132+
133+
134+
@app.function(
135+
cpu=4,
136+
memory=8 * 1024,
137+
timeout=10 * 60,
138+
volumes={"/v": models_volume},
139+
secrets=[modal.Secret.from_name("api-inference-key")],
140+
# cold starts (~30-60s model load) instead of a 24/7 warm container
141+
# — money discipline; bump once traffic justifies it
142+
)
143+
@modal.asgi_app()
144+
def web():
145+
return make_api()

‎src/gpu/modal_distill.py‎

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,9 @@
116116
"max_len": 1450,
117117
"label_beams": "1",
118118
"out": "rababa_arabic_distill_small/run-002",
119+
# v2: every label generated before the byt5 decode_joined fix is
120+
# mojibake (double-encoded); relabel from scratch on the new file
121+
"labels_file": "teacher_labels_v2.jsonl",
119122
"mode": "sequence",
120123
"note": "r6 canonical (2.5793 DER); gate <= 3.07 windowed DER-CE",
121124
},
@@ -233,11 +236,24 @@
233236

234237

235238
def decode_joined(tok, ids) -> str:
236-
"""Correct decode for umt5 teachers: 5.x batch_decode inserts spurious
237-
spaces between sentencepiece pieces; pieces must join directly (the
238-
targets are unspaced IPA strings)."""
239+
"""Decode teacher generations correctly per tokenizer family.
240+
241+
umt5/sentencepiece: 5.x batch_decode inserts spurious spaces between
242+
pieces; joining convert_ids_to_tokens is correct.
243+
244+
byt5/byte-level: convert_ids_to_tokens returns each byte token as a
245+
RAW CHARACTER (latin-1 view) — joining them produces mojibake. This
246+
poisoned every Arabic label generated under 5.14 (both students
247+
trained on double-encoded targets and scored the identical bare-text
248+
DER). batch_decode is byte-exact here, so branch on it: if its
249+
output round-trips the same token ids it is authoritative.
250+
"""
239251
skip = {tok.pad_token, tok.eos_token, tok.bos_token}
240-
return "".join(p for p in tok.convert_ids_to_tokens(ids) if p not in skip)
252+
joined = "".join(p for p in tok.convert_ids_to_tokens(ids) if p not in skip)
253+
if joined and all(ord(c) < 256 for c in joined):
254+
# byte-level vocab decoded to raw chars — use the byte-exact path
255+
return tok.batch_decode([ids], skip_special_tokens=True)[0]
256+
return joined
241257

242258

243259
@app.function(

0 commit comments

Comments
 (0)