Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
117 changes: 61 additions & 56 deletions src/gpu/modal_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Expand All @@ -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()
Expand Down Expand Up @@ -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))


83 changes: 83 additions & 0 deletions src/imf/export.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
64 changes: 64 additions & 0 deletions tests/test_imf_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading