Jarbas's picture
Upload quantize_nar.py with huggingface_hub
095a3aa verified
Raw
History Blame Contribute Delete
1.46 kB
#!/usr/bin/env python3
"""Quantize the already exported granite-nar graphs to int8.
The encoder is quantized per channel with the in-graph front end kept in fp32.
The editor is quantized per tensor: per-channel needs more than 110 GB of RAM on
this graph, and the text side of this family is known to take per-tensor well.
"""
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent))
from export_granite_nar import consolidate, front_end_nodes # noqa: E402
OUT = Path(__file__).parent
TMP = OUT / "_staging"
def run(name: str, per_channel: bool, exclude: list[str]) -> None:
from onnxruntime.quantization import QuantType, quantize_dynamic
stage = TMP / f"{name}_int8"
stage.mkdir(parents=True, exist_ok=True)
print(f"quantizing {name} (per_channel={per_channel}, {len(exclude)} nodes excluded)...", flush=True)
quantize_dynamic(
OUT / f"{name}.onnx",
stage / f"{name}_int8.onnx",
weight_type=QuantType.QInt8,
per_channel=per_channel,
use_external_data_format=True,
nodes_to_exclude=exclude,
extra_options={"MatMulConstBOnly": True},
)
consolidate(stage / f"{name}_int8.onnx", OUT / f"{name}_int8.onnx")
print(name, (OUT / f"{name}_int8.onnx").stat().st_size, flush=True)
for arg in sys.argv[1:]:
if arg == "encoder":
run("encoder", True, front_end_nodes(OUT / "encoder.onnx"))
else:
run("editor", False, [])