#!/usr/bin/env python3 """Export purpose-lite to fixed-shape fp16 and int8-QDQ ONNX artifacts.""" from __future__ import annotations import argparse import hashlib import shutil import sys import tempfile from pathlib import Path from typing import Any, Sequence import numpy as np from purpose_data import DataError, load_jsonl, normalize_prompt, write_json from train import ( HEAD_TOKENS, MAX_LENGTH, TAIL_TOKENS, encode_fixed_shape, ) SCRIPT_DIR = Path(__file__).resolve().parent DEFAULT_MODEL_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1" / "model" DEFAULT_CALIBRATION = SCRIPT_DIR / "outputs" / "purpose-lite-v1" / "calibration.json" DEFAULT_VALIDATION = SCRIPT_DIR / ".artifacts" / "dataset-v1" / "validation.jsonl" DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1" / "export" MODEL_VERSION = "purpose-lite-v1" SHIPPING_BUDGET_BYTES = 25 * 1024 * 1024 GOLDEN_PROMPTS = ( "Fix the typo in the README.", " Cafe\u0301\tdeploy\nnow ", ( "EXPLAIN ANALYZE shows a sequential scan before the incident notes. " + "context " * 180 + "Find the root cause and explain which query plan evidence proves it." ), "レビューだけして、コードは変更しないでください。", ) def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def _fixed_shape_inputs(model: Any) -> list[str]: names = ["input_ids", "attention_mask"] if getattr(model.config, "type_vocab_size", 0) > 1: names.append("token_type_ids") return names def _validate_graph(path: Path, input_names: Sequence[str]) -> dict[str, Any]: import onnx model = onnx.load(path) onnx.checker.check_model(model) opset = max( item.version for item in model.opset_import if item.domain in ("", "ai.onnx") ) if opset < 17: raise DataError(f"{path}: ONNX opset {opset} is below 17") shapes = {} for value in model.graph.input: dimensions = [item.dim_value for item in value.type.tensor_type.shape.dim] shapes[value.name] = dimensions expected = {name: [1, MAX_LENGTH] for name in input_names} if shapes != expected: raise DataError(f"{path}: expected fixed inputs {expected}, got {shapes}") custom_domains = sorted( { node.domain for node in model.graph.node if node.domain not in ("", "ai.onnx") } ) if custom_domains: raise DataError(f"{path}: custom ONNX op domains are not allowed: {custom_domains}") return { "opset": opset, "inputs": shapes, "nodes": len(model.graph.node), "operators": sorted({node.op_type for node in model.graph.node}), } def _write_tokenizer_contract( output_dir: Path, model_dir: Path, tokenizer: Any, torch: Any ) -> dict[str, Any]: vocab_destination = output_dir / "vocab.txt" vocabulary = tokenizer.get_vocab() ordered_vocabulary = sorted(vocabulary.items(), key=lambda item: item[1]) if [token_id for _, token_id in ordered_vocabulary] != list( range(len(ordered_vocabulary)) ): raise DataError("tokenizer vocabulary IDs are not contiguous") vocab_destination.write_text( "".join(f"{token}\n" for token, _ in ordered_vocabulary), encoding="utf-8", ) tokenizer_json_source = model_dir / "tokenizer.json" if not tokenizer_json_source.exists(): raise DataError(f"{tokenizer_json_source}: tokenizer JSON is missing") tokenizer_json_destination = output_dir / "tokenizer.json" shutil.copy2(tokenizer_json_source, tokenizer_json_destination) golden_values = [] for prompt in GOLDEN_PROMPTS: encoded = encode_fixed_shape(tokenizer, [prompt], torch) golden_values.append( { "prompt": prompt, "normalizedPrompt": normalize_prompt(prompt), "inputIds": encoded["input_ids"][0].tolist(), "attentionMask": encoded["attention_mask"][0].tolist(), "tokenTypeIds": encoded.get( "token_type_ids", torch.zeros_like(encoded["input_ids"]) )[0].tolist(), } ) write_json(output_dir / "tokenizer-goldens.json", golden_values) contract = { "schemaVersion": 1, "modelVersion": MODEL_VERSION, "tokenizer": { "family": "BERT WordPiece", "vocabFile": vocab_destination.name, "vocabSha256": _sha256(vocab_destination), "tokenizerJSON": tokenizer_json_destination.name, "tokenizerJSONSha256": _sha256(tokenizer_json_destination), "lowercase": bool(getattr(tokenizer, "do_lower_case", True)), }, "normalization": ["Unicode NFKC", "collapse Unicode whitespace", "trim"], "input": { "shape": [1, MAX_LENGTH], "padding": "right", "longPromptStrategy": "BERT sentence pair: head and tail", "headTokens": HEAD_TOKENS, "tailTokens": TAIL_TOKENS, "specialTokenLayout": "[CLS] head [SEP] tail [SEP]", }, "specialTokenIds": { "padding": tokenizer.pad_token_id, "unknown": tokenizer.unk_token_id, "classification": tokenizer.cls_token_id, "separator": tokenizer.sep_token_id, }, "goldens": "tokenizer-goldens.json", } write_json(output_dir / "tokenizer-spec.json", contract) return contract def export(args: argparse.Namespace) -> dict[str, Any]: try: import onnx import torch from onnxconverter_common import float16 from onnxruntime.quantization import ( CalibrationDataReader, QuantFormat, QuantType, quantize_static, ) from transformers import AutoModelForSequenceClassification, AutoTokenizer except ImportError as exc: raise DataError( "export dependencies are missing; install requirements.txt" ) from exc output_dir: Path = args.output_dir output_dir.mkdir(parents=True, exist_ok=True) tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True) model = AutoModelForSequenceClassification.from_pretrained( args.model_dir, local_files_only=True ) model.eval() input_names = _fixed_shape_inputs(model) example = encode_fixed_shape(tokenizer, ["Plan a safe cache migration."], torch) class LogitsModel(torch.nn.Module): def __init__(self, inner: Any) -> None: super().__init__() self.inner = inner def forward(self, *values: Any) -> Any: inputs = dict(zip(input_names, values)) return self.inner(**inputs).logits fp16_path = output_dir / f"{MODEL_VERSION}-fp16.onnx" int8_path = output_dir / f"{MODEL_VERSION}-int8-qdq.onnx" with tempfile.TemporaryDirectory() as temporary: fp32_path = Path(temporary) / f"{MODEL_VERSION}-fp32.onnx" torch.onnx.export( LogitsModel(model), tuple(example[name] for name in input_names), fp32_path, input_names=input_names, output_names=["logits"], opset_version=args.opset, do_constant_folding=True, dynamo=False, ) fp32_model = onnx.load(fp32_path) fp16_model = float16.convert_float_to_float16( fp32_model, keep_io_types=True, disable_shape_infer=False, ) onnx.save(fp16_model, fp16_path) validation = load_jsonl(args.validation) class Reader(CalibrationDataReader): def __init__(self) -> None: self.index = 0 self.samples = validation[: args.calibration_records] def get_next(self) -> dict[str, np.ndarray] | None: if self.index >= len(self.samples): return None record = self.samples[self.index] self.index += 1 encoded = encode_fixed_shape(tokenizer, [record["prompt"]], torch) return { name: encoded[name].numpy().astype(np.int64, copy=False) for name in input_names } quantize_static( fp32_path, int8_path, Reader(), quant_format=QuantFormat.QDQ, activation_type=QuantType.QUInt8, weight_type=QuantType.QInt8, per_channel=True, # Gather is essential: MiniLM's 30k x 384 embedding table is more than half # the checkpoint. Quantizing only MatMul/Gemm leaves a ~47 MB fp32 table and # cannot meet the 25 MB parity-floor artifact budget. op_types_to_quantize=["Gather", "MatMul", "Gemm"], extra_options={ "ActivationSymmetric": False, "WeightSymmetric": True, }, ) graph_reports = { "fp16": _validate_graph(fp16_path, input_names), "int8QDQ": _validate_graph(int8_path, input_names), } int8_size = int8_path.stat().st_size if int8_size > args.shipping_budget_bytes: raise DataError( f"{int8_path}: {int8_size} bytes exceeds the " f"{args.shipping_budget_bytes}-byte shipping budget" ) tokenizer_contract = _write_tokenizer_contract( output_dir, args.model_dir, tokenizer, torch ) calibration_destination = output_dir / "calibration.json" shutil.copy2(args.calibration, calibration_destination) artifacts = {} for name, path in (("fp16", fp16_path), ("int8QDQ", int8_path)): artifacts[name] = { "path": path.name, "bytes": path.stat().st_size, "sha256": _sha256(path), } report = { "schemaVersion": 1, "modelVersion": MODEL_VERSION, "sourceModel": str(args.model_dir), "opset": args.opset, "fixedInputShape": [1, MAX_LENGTH], "inputNames": input_names, "calibrationRecords": min(args.calibration_records, len(validation)), "shippingArtifact": "int8QDQ", "shippingBudgetBytes": args.shipping_budget_bytes, "shippingBudgetPassed": int8_size <= args.shipping_budget_bytes, "artifacts": artifacts, "graphs": graph_reports, "tokenizerSpecSha256": _sha256(output_dir / "tokenizer-spec.json"), "tokenizerContract": tokenizer_contract["input"], "calibrationSha256": _sha256(calibration_destination), } write_json(output_dir / "export-metrics.json", report) return report def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model-dir", type=Path, default=DEFAULT_MODEL_DIR) parser.add_argument("--calibration", type=Path, default=DEFAULT_CALIBRATION) parser.add_argument("--validation", type=Path, default=DEFAULT_VALIDATION) parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) parser.add_argument("--opset", type=int, default=17) parser.add_argument("--calibration-records", type=int, default=256) parser.add_argument( "--shipping-budget-bytes", type=int, default=SHIPPING_BUDGET_BYTES ) return parser def main(argv: Sequence[str] | None = None) -> int: parser = build_parser() args = parser.parse_args(argv) if ( args.opset < 17 or args.calibration_records <= 0 or args.shipping_budget_bytes <= 0 ): parser.error("opset must be >=17 and record/budget values must be positive") try: report = export(args) except (DataError, OSError, ValueError) as exc: print(f"error: {exc}", file=sys.stderr) return 1 print( f"Exported {MODEL_VERSION}: fp16={report['artifacts']['fp16']['bytes']} bytes, " f"int8-QDQ={report['artifacts']['int8QDQ']['bytes']} bytes." ) return 0 if __name__ == "__main__": raise SystemExit(main())