Skip to content

Commit ccb0be0

Browse files
author
Ronald Tse
committed
fix(imf): fp16 via torch-native half export — ORT converter broken on real ByT5
The onnxruntime float16 converter produced all-zero encoder hiddens on the real khm-latn checkpoint (CER 1939pp, every sample mismatched) while looking fine on the tiny fixture. Exporting the model under .half() is exact on gold pairs; graph IO becomes float16 (int64 ids unchanged) and _zero_pasts follows session dtypes. Measured on 300 khm test samples: fp16 delta 0.43pp, int8 0.84pp — quantization noise (argmax flips cascading under greedy decode), not breakage. Whether lossy precisions may exceed the 0.2pp export-fidelity bar is a policy call pending; fp32 measures 0.0pp.
1 parent 96ada69 commit ccb0be0

2 files changed

Lines changed: 23 additions & 27 deletions

File tree

‎docs/imf-v1.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ model.zip
4747
| `tokenizer` | enum | `bytes` (the only v1 value) |
4848
| `opset` | int | 7..14; must equal the graphs' opset |
4949
| `decoder` | enum | `plain` \| `kv` (`kv` requires decoder-kv.onnx) |
50-
| `precision` | enum | `fp32` \| `fp16` \| `int8` |
50+
| `precision` | enum | `fp32` \| `fp16` \| `int8` (fp16 = torch-native half export: float16 graph IO, int64 ids unchanged; runtimes read dtypes from the session — the ORT float16 converter produces all-zero hiddens on real ByT5 and must not be used) |
5151
| `license` | str | non-empty (strict gate) |
5252
| `trained_from` | str | repo + run/checkpoint id |
5353
| `metrics` | list | `{name, value, protocol, source}`; `source` must be a `RESULTS.md#anchor` (strict gate) |

‎src/imf/export.py‎

Lines changed: 22 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -242,27 +242,18 @@ def export_graphs(model, out_dir: Path | str) -> dict[str, Path]:
242242
return paths
243243

244244

245-
def convert_fp16(src: Path | str, dst: Path | str) -> Path:
246-
"""fp32 -> mixed fp16, IO types preserved (encoder/decoder compose cleanly).
247-
248-
LayerNorm/softmax math stays fp32: ORT's session-time
249-
SimplifiedLayerNormFusion crashes on half-converted LN subgraphs
250-
(InsertPrecisionFreeCast name mismatch), so the whole decomposition
251-
must stay one dtype. Weights (MatMuls) carry the size win.
245+
def convert_fp16(model):
246+
"""A fp16 copy of the model for torch-native half export.
247+
248+
The onnxruntime float16 CONVERTER is not usable here: on real ByT5
249+
checkpoints it produces all-zero encoder hiddens (found 2026-08-16,
250+
khm-latn — 1939pp CER); exporting the torch model under .half() is
251+
exact on gold pairs. Graph IO becomes float16 (input_ids stay int64);
252+
runtimes read dtypes from the session, and _zero_pasts follows them.
252253
"""
253-
import onnx
254-
from onnxruntime.transformers import float16
255-
256-
block_list = list(float16.DEFAULT_OP_BLOCK_LIST) + [
257-
"ReduceMean", "Pow", "Sqrt", "Div", "Sub", "Add", "Mul",
258-
"Softmax", "Range", "Exp", "Where", "Less", "Cast",
259-
]
260-
model = onnx.load(str(src))
261-
converted = float16.convert_float_to_float16(
262-
model, keep_io_types=True, op_block_list=block_list
263-
)
264-
onnx.save(converted, str(dst))
265-
return Path(dst)
254+
import copy
255+
256+
return copy.deepcopy(model).half()
266257

267258

268259
def quantize_int8(src: Path | str, dst: Path | str) -> Path:
@@ -320,7 +311,8 @@ def _zero_pasts(kv_sess) -> dict[str, object]:
320311
shape = meta.shape # [batch, heads, past_seq, d_kv] with str dynamic dims
321312
heads = shape[1] if isinstance(shape[1], int) else 4
322313
d_kv = shape[3] if isinstance(shape[3], int) else 8
323-
pasts[meta.name] = np.zeros((1, heads, 0, d_kv), dtype=np.float32)
314+
dtype = np.float16 if meta.type == "tensor(float16)" else np.float32
315+
pasts[meta.name] = np.zeros((1, heads, 0, d_kv), dtype=dtype)
324316
return pasts
325317

326318

@@ -372,18 +364,22 @@ def export_zips(
372364
with tempfile.TemporaryDirectory() as tmp:
373365
tmp = Path(tmp)
374366
graphs = export_graphs(model, tmp / "graphs")
367+
graphs_16 = (
368+
export_graphs(convert_fp16(model), tmp / "graphs-fp16")
369+
if "fp16" in precisions
370+
else {}
371+
)
375372

376373
for precision in precisions:
377374
variant_dir = tmp / precision
378375
variant_dir.mkdir()
379-
for name, src in graphs.items():
376+
sources = graphs if precision != "fp16" else graphs_16
377+
for name, src in sources.items():
380378
dst = variant_dir / name
381-
if precision == "fp32":
379+
if precision == "fp32" or precision == "fp16":
382380
dst.write_bytes(src.read_bytes())
383-
elif precision == "fp16":
384-
convert_fp16(src, dst)
385381
elif precision == "int8":
386-
quantize_int8(src, dst)
382+
quantize_int8(graphs[name], dst)
387383
else:
388384
raise ValueError(f"unknown precision {precision!r}")
389385
meta = replace(metadata, precision=precision)

0 commit comments

Comments
 (0)