From f18399d7d1e4b2d71511213c55e9a9d5f3a6f4e4 Mon Sep 17 00:00:00 2001 From: Ronald Tse Date: Tue, 29 Sep 2026 23:09:43 +0800 Subject: [PATCH] feat: int8-static export builds and gates the browser-size composition MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit build_int8_static_zip in imf.export: dynamic-int8 encoder + static-int8 decoder in one zip, calibrated on the model's own fp32 decode across framings (prefill, single steps, 8-token windows); metadata precision rewritten to int8 (PRECISIONS is closed — the recipe is the zip name, the gate the value). Fixture specs cover the release sequence end to end: build, parity gate, write_parity, strict validation, both quant recipes present, head MatMul kept fp32. modal_export::export_int8_static becomes thin orchestration over the helper plus the rebuild_int8_head32 gate stack (reference decode, parity, margin analysis, confident-flip budget), landing {mid}-int8static.zip on the volume. --- src/gpu/modal_export.py | 117 ++++++++++++++++++++------------------- src/imf/export.py | 83 +++++++++++++++++++++++++++ tests/test_imf_export.py | 64 +++++++++++++++++++++ 3 files changed, 208 insertions(+), 56 deletions(-) diff --git a/src/gpu/modal_export.py b/src/gpu/modal_export.py index 17b818c..a1037b0 100644 --- a/src/gpu/modal_export.py +++ b/src/gpu/modal_export.py @@ -590,33 +590,40 @@ def rebuild_int8_head32(model_id: str, limit: int = 0) -> dict: timeout=5 * 3600, volumes={**CHECKPOINT_VOLUMES, **DATASET_VOLUMES, "/outputs": MODELS_VOLUME}, ) -def export_int8_static(model_id: str, limit: int = 0) -> dict: - """Build the static-activation int8 artifact (TODO.impl/11's - positive branch): fp32 encoder + static-int8 decoder, calibrated on - the model's own eval pairs across BOTH decode framings. - - The MEASURED composition (scored 4.6241 full-set vs dynamic 4.5701, - +8% CPU decode): the fp32 encoder keeps the artifact server-sized; - a browser-sized int8-encoder + static-decoder composition needs its - own gate run before it ships. Lands as {mid}-int8static.zip; the - release swap is a version decision.""" - import re +def export_int8_static(model_id: str, limit: int = 0, calibration: int = 5) -> dict: + """Build and gate the static-activation int8 artifact (TODO.impl/11's + positive branch). Composition: dynamic-int8 encoder + static-int8 + decoder — measured 4.6045 full-set (vs dynamic 4.5701, not + separated) with +8% CPU decode speed, at 491 MiB. Calibration walks + the model's own fp32 decode over a small slice of eval inputs + (scripts/static_int8_experiment.py's recipe: 5 rows, framing- + diverse feeds). + + Gates, identical to rebuild_int8_head32: CER parity written into + the zip, margin analysis, confident-flip budget. Lands as + {mid}-int8static.zip; publication is a version decision.""" import sys - import tempfile - import zipfile from pathlib import Path sys.path.insert(0, "/root/interscript-ml/src") - from imf.export import ( - collect_decode_calibration, - head_matmul_names, - quantize_int8_static, - refresh_member_shas, - ) spec = MODELS[model_id] + checkpoint = Path(spec["volume"]) / spec["checkpoint"] test_path = Path(spec["test_volume"]) / spec["test_data"] - pairs = _load_pairs(test_path)[: limit or None] + + from imf.export import build_int8_static_zip, load_byte_seq2seq + from imf.parity import ( + reference_decode, + run_margin_analysis, + run_parity, + write_margin_report, + write_parity, + ) + + model = load_byte_seq2seq(checkpoint) + pairs = _load_pairs(test_path) + if limit: + pairs = pairs[:limit] out_dir = Path("/outputs/imf") / model_id meta_path = Path("/root/interscript-ml", spec["metadata"]) @@ -625,41 +632,42 @@ def export_int8_static(model_id: str, limit: int = 0) -> dict: if not fp32_zip.exists(): raise RuntimeError(f"{fp32_zip.name} missing on the volume") - import onnxruntime as ort + new_zip = build_int8_static_zip( + model, fp32_zip, [src for src, _ in pairs[:calibration]], + out_dir / f"{mid}-int8static.zip", + ) - with tempfile.TemporaryDirectory() as tmp: - tmp = Path(tmp) - with zipfile.ZipFile(fp32_zip) as zf: - zf.extract("encoder.onnx", tmp) - dec = "decoder-kv.onnx" if "decoder-kv.onnx" in zf.namelist() else "decoder.onnx" - zf.extract(dec, tmp) - opts = ort.SessionOptions() - opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_DISABLE_ALL - enc_sess = ort.InferenceSession(str(tmp / "encoder.onnx"), opts) - dec_sess = ort.InferenceSession(str(tmp / dec), opts) - calibration = collect_decode_calibration(enc_sess, dec_sess, [s for s, _ in pairs]) - print(f"[{model_id}] calibration: {len(calibration)} feeds", flush=True) - dec_static = tmp / dec.replace(".onnx", "-static.onnx") - quantize_int8_static( - tmp / dec, dec_static, calibration, - nodes_to_exclude=head_matmul_names(tmp / dec), + reference = reference_decode( + model, + [src for src, _ in pairs], + max_len=128, + resume_path=out_dir / f"{mid}-int8static-reference.jsonl", + ) + report = run_parity(model, new_zip, pairs, max_len=128, reference=reference) + if not report.passed: + raise RuntimeError(f"parity gate FAILED for {new_zip.name}: {report}") + write_parity(new_zip, report) + margins = run_margin_analysis(model, new_zip, pairs, max_len=128) + write_margin_report(margins, out_dir / f"{mid}-int8static-margins.json") + confident = margins.flip_rate * (1 - margins.flip_low_margin_share) + if confident > 0.01: + raise RuntimeError( + f"margin gate FAILED for {new_zip.name}: {confident:.2%} confident flips" ) - new_zip = out_dir / f"{mid}-int8static.zip" - with zipfile.ZipFile(fp32_zip) as src, zipfile.ZipFile( - new_zip, "w", zipfile.ZIP_DEFLATED - ) as dst: - for member in src.namelist(): - if member == "metadata.yaml": - meta = src.read(member).decode("utf-8") - meta = re.sub(r"^precision:\s*\S+", "precision: int8-static", meta, flags=re.M) - meta = re.sub(r"^id:\s*\S+", f"id: {mid}-int8static", meta, flags=re.M) - dst.writestr(member, meta) - elif member == dec: - dst.writestr(member, dec_static.read_bytes()) - else: - dst.writestr(member, src.read(member)) - refresh_member_shas(new_zip) - return {"artifact": str(new_zip), "calibration": len(calibration)} + MODELS_VOLUME.commit() + return { + "model": model_id, "zip": new_zip.name, + "parity": {"samples": report.samples, "cer_delta": report.cer_delta}, + "margins": {"flip_rate": margins.flip_rate, "kld": margins.kld_mean, + "low_share": margins.flip_low_margin_share, + "confident_flip_rate": round(confident, 6)}, + "size_bytes": new_zip.stat().st_size, + } + + +@app.local_entrypoint() +def static(model_id: str, limit: int = 0, calibration: int = 5) -> None: + print(export_int8_static.remote(model_id, limit, calibration)) @app.local_entrypoint() @@ -696,8 +704,5 @@ def zip_meta(model_id: str, precision: str) -> dict: def zmeta(model: str, precisions: str = "fp32,fp16,int8") -> None: for precision in precisions.split(","): print(precision, zip_meta.remote(model, precision)) -@app.local_entrypoint() -def static(model_id: str, limit: int = 0) -> None: - print(export_int8_static.remote(model_id, limit)) diff --git a/src/imf/export.py b/src/imf/export.py index 151ecd1..5c194dc 100644 --- a/src/imf/export.py +++ b/src/imf/export.py @@ -525,6 +525,89 @@ def export_zips( +def build_int8_static_zip( + model, + fp32_zip: Path | str, + calibration_texts: list[str], + out_zip: Path | str, +) -> Path: + """Dynamic-int8 encoder + static-int8 decoder in one zip — the + browser-size composition measured in TODO.impl/11 (4.6045 full-set + vs dynamic 4.5701, not separated; +8% CPU decode). Calibration walks + the model's own fp32 decode over ``calibration_texts`` across + framings (prefill, single steps, 8-token windows). Metadata + precision becomes ``int8``: PRECISIONS is closed, so the static + recipe is carried by the zip name and the gate by the value. + + Graph-only work; gates (parity, margins) are the caller's.""" + import tempfile + import zipfile + from dataclasses import replace as dc_replace + + import onnxruntime as ort + import yaml + + from imf.pack import _to_dict + from imf.schema import ModelMetadata + + fp32_zip, out_zip = Path(fp32_zip), Path(out_zip) + out_zip.parent.mkdir(parents=True, exist_ok=True) + + with tempfile.TemporaryDirectory() as tmp: + tmp = Path(tmp) + with zipfile.ZipFile(fp32_zip) as zf: + zf.extract("encoder.onnx", tmp) + dec = "decoder-kv.onnx" if "decoder-kv.onnx" in zf.namelist() else "decoder.onnx" + zf.extract(dec, tmp) + meta_text = zf.read("metadata.yaml").decode("utf-8") + + opts = ort.SessionOptions() + opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_DISABLE_ALL + enc_sess = ort.InferenceSession(str(tmp / "encoder.onnx"), opts) + dec_sess = ort.InferenceSession(str(tmp / dec), opts) + + # scales must hold on the shapes the runtime actually feeds; + # the fp32 decode supplies real hidden states and cache lengths + calibration = collect_decode_calibration( + enc_sess, dec_sess, calibration_texts + ) + print(f" calibration: {len(calibration)} feeds", flush=True) + + enc_q = tmp / "encoder-dyn.onnx" + quantize_int8(tmp / "encoder.onnx", enc_q) + dec_q = tmp / dec.replace(".onnx", "-static.onnx") + quantize_int8_static( + tmp / dec, dec_q, calibration, + nodes_to_exclude=head_matmul_names(tmp / dec), + ) + + metadata = dc_replace( + ModelMetadata.from_yaml(meta_text), precision="int8" + ) + with zipfile.ZipFile(fp32_zip) as src, zipfile.ZipFile( + out_zip, "w", zipfile.ZIP_DEFLATED + ) as dst: + for name in src.namelist(): + if name == "metadata.yaml": + dst.writestr( + name, + yaml.safe_dump( + _to_dict(metadata), sort_keys=False, allow_unicode=True + ), + ) + elif name == "encoder.onnx": + dst.writestr(name, enc_q.read_bytes()) + elif name == dec: + dst.writestr(name, dec_q.read_bytes()) + else: + dst.writestr(name, src.read(name)) + + # re-quantized graphs replaced members; the internal sha table must + # refresh or strict validation (and write_parity) rejects the zip + refresh_member_shas(out_zip) + return out_zip + + def collect_decode_calibration( enc_sess, dec_sess, texts: list[str], steps: int = 64, ) -> list[dict]: diff --git a/tests/test_imf_export.py b/tests/test_imf_export.py index bc2bfc2..181561e 100644 --- a/tests/test_imf_export.py +++ b/tests/test_imf_export.py @@ -189,3 +189,67 @@ def test_refresh_member_shas_after_member_replacement(zips: dict[str, Path]) -> with zipfile.ZipFile(path) as after_zf: assert after_zf.read("encoder.onnx") == tampered assert after_zf.read("decoder.onnx") == members["decoder.onnx"] + + +# --------------------------------------------------------------------------- +# int8-static: calibrated-activation decoder (TODO.impl/11's positive +# branch) — dynamic-int8 encoder + static-int8 decoder composition. + + +@pytest.fixture(scope="module") +def static_zip( + zips: dict[str, Path], reference_model, tmp_path_factory: pytest.TempPathFactory +) -> Path: + """The release sequence end-to-end: build, gate, write parity.""" + from imf.export import build_int8_static_zip + from imf.parity import run_parity, write_parity + + path = build_int8_static_zip( + reference_model, + zips["fixture-1.0-fp32.zip"], + TEXTS, + tmp_path_factory.mktemp("static") / "fixture-1.0-int8static.zip", + ) + report = run_parity( + reference_model, path, [(t, t) for t in TEXTS] * 170, max_len=MAX_LEN + ) + assert report.passed, report + write_parity(path, report) + return path + + +def test_static_zip_declares_int8_and_validates_strict(static_zip: Path) -> None: + """The closed PRECISIONS set has no 'int8-static': the recipe lives + in the zip name, the gate (int8's 2pp cer_delta limit) in the value.""" + strict = validate_zip(static_zip, strict=True) + assert strict.ok, strict.errors + with zipfile.ZipFile(static_zip) as zf: + meta = yaml.safe_load(zf.read("metadata.yaml")) + assert meta["precision"] == "int8" + + +def test_static_zip_graphs_use_both_quant_recipes(static_zip: Path) -> None: + """Encoder: dynamic int8 (MatMulInteger). Decoder: static + QOperator (QLinearMatMul) with calibrated activation scales.""" + import onnx + + with zipfile.ZipFile(static_zip) as zf: + enc_ops = {n.op_type for n in onnx.load(zf.open("encoder.onnx")).graph.node} + dec_ops = {n.op_type for n in onnx.load(zf.open("decoder-kv.onnx")).graph.node} + assert "MatMulInteger" in enc_ops + assert "QLinearMatMul" in dec_ops + + +def test_static_zip_keeps_head_matmul_fp32(static_zip: Path, zips: dict[str, Path]) -> None: + """The head-fp32 rule applies to the static recipe unchanged.""" + import onnx + + from imf.export import _head_matmul_names + + with zipfile.ZipFile(zips["fixture-1.0-fp32.zip"]) as zf: + heads = _head_matmul_names(onnx.load(zf.open("decoder-kv.onnx")).graph) + assert heads + with zipfile.ZipFile(static_zip) as zf: + nodes = {n.name: n for n in onnx.load(zf.open("decoder-kv.onnx")).graph.node} + for head in heads: + assert head in nodes and nodes[head].op_type == "MatMul", head