Files
nucleic-purpose-classifier/export.py
T

335 lines
12 KiB
Python

#!/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())