Files
nucleic-purpose-classifier/quantize_coreml.py
T

305 lines
11 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Calibrate a Core ML W8A8 candidate from the selected float16 ML Program."""
from __future__ import annotations
import argparse
import gc
import hashlib
import shutil
import sys
import tempfile
from collections import Counter
from contextlib import contextmanager
from pathlib import Path
from typing import Any, Sequence
import numpy as np
from convert_coreml import COREMLTOOLS_VERSION, _package_manifest
from export import stratified_calibration_sample
from purpose_data import DataError, load_jsonl, prompt_hash, write_json
from train import encode_fixed_shape
SCRIPT_DIR = Path(__file__).resolve().parent
CANDIDATE_DIR = (
SCRIPT_DIR / "outputs" / "purpose-lite-v1-distilled-qat-mlx-4e"
)
DEFAULT_MODEL = CANDIDATE_DIR / "coreml" / "purpose-lite-v1-fp16.mlpackage"
DEFAULT_MODEL_DIR = CANDIDATE_DIR / "model"
DEFAULT_VALIDATION = SCRIPT_DIR / ".artifacts" / "dataset-v1" / "validation.jsonl"
DEFAULT_OUTPUT = CANDIDATE_DIR / "coreml" / "purpose-lite-v1-w8a8.mlpackage"
SHIPPING_BUDGET_BYTES = 25 * 1024 * 1024
def _tree_sha256(package: Path) -> str:
digest = hashlib.sha256()
for path in sorted(item for item in package.rglob("*") if item.is_file()):
relative = str(path.relative_to(package)).encode("utf-8")
digest.update(relative)
digest.update(b"\0")
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
@contextmanager
def _bounded_calibration_packages(debugger_type: Any, temporary_root: Path):
"""Eagerly remove Core ML Tools' per-prediction temporary packages.
Core ML Tools registers these packages for process-exit cleanup. Activation
calibration creates one package per intermediate-output group per record, so retaining
all of them can consume tens of gigabytes before the process exits.
"""
original_predict = debugger_type.predict_intermediate_outputs
previous_tempdir = tempfile.tempdir
def predict_and_cleanup(*args: Any, **kwargs: Any) -> Any:
try:
return original_predict(*args, **kwargs)
finally:
gc.collect()
for package in temporary_root.glob("*.mlpackage"):
shutil.rmtree(package)
temporary_root.mkdir(parents=True, exist_ok=True)
tempfile.tempdir = str(temporary_root)
debugger_type.predict_intermediate_outputs = predict_and_cleanup
try:
yield
finally:
debugger_type.predict_intermediate_outputs = original_predict
tempfile.tempdir = previous_tempdir
for package in temporary_root.glob("*.mlpackage"):
shutil.rmtree(package)
def _optimization_configs(optimize: Any) -> tuple[Any, Any]:
"""Return the Core ML analogue of the accepted ONNX QDQ policy."""
activation = optimize.coreml.OpLinearQuantizerConfig(
mode="linear",
dtype=np.uint8,
granularity="per_tensor",
)
activation_config = optimize.coreml.OptimizationConfig(
global_config=activation,
)
linear_weight = optimize.coreml.OpLinearQuantizerConfig(
mode="linear_symmetric",
dtype=np.int8,
granularity="per_channel",
weight_threshold=2048,
)
embedding_weight = optimize.coreml.OpLinearQuantizerConfig(
mode="linear",
dtype=np.uint8,
granularity="per_tensor",
weight_threshold=2048,
)
weight_config = optimize.coreml.OptimizationConfig(
op_type_configs={
"gather": embedding_weight,
"linear": linear_weight,
"matmul": linear_weight,
}
)
return activation_config, weight_config
def _validate_args(args: argparse.Namespace) -> None:
if not args.model.is_dir():
raise DataError(f"{args.model}: source Core ML package is missing")
if not args.model_dir.is_dir():
raise DataError(f"{args.model_dir}: tokenizer directory is missing")
if args.output.suffix != ".mlpackage":
raise DataError("Core ML output must end in .mlpackage")
try:
same_output = args.model.resolve() == args.output.resolve()
except OSError as exc:
raise DataError(f"cannot resolve Core ML package paths: {exc}") from exc
if same_output:
raise DataError("W8A8 output must not overwrite its float16 source package")
if args.output.exists() and not args.overwrite_output:
raise DataError(
f"{args.output}: output exists; pass --overwrite-output intentionally"
)
def quantize(args: argparse.Namespace) -> dict[str, Any]:
_validate_args(args)
try:
import coremltools as ct
import coremltools.optimize as cto
import torch
from coremltools.optimize.coreml.experimental._model_debugger import (
ModelDebugger,
)
from transformers import AutoTokenizer
except ImportError as exc:
raise DataError(
"Core ML quantization requires requirements-coreml.txt on macOS"
) from exc
if ct.__version__ != COREMLTOOLS_VERSION:
raise DataError(
f"expected coremltools {COREMLTOOLS_VERSION}, found {ct.__version__}"
)
validation = load_jsonl(args.validation)
calibration = stratified_calibration_sample(
validation,
args.calibration_records,
seed=args.calibration_seed,
)
tokenizer = AutoTokenizer.from_pretrained(
args.model_dir,
local_files_only=True,
)
sample_data = []
for record in calibration:
encoded = encode_fixed_shape(tokenizer, [record["prompt"]], torch)
sample_data.append(
{
name: value.numpy().astype(np.int32, copy=False)
for name, value in encoded.items()
}
)
print(
f"Core ML activation calibration: {len(sample_data)} records",
flush=True,
)
source_model = ct.models.MLModel(
str(args.model),
compute_units=ct.ComputeUnit.CPU_ONLY,
)
input_names = {item.name for item in source_model.get_spec().description.input}
expected_inputs = {"input_ids", "attention_mask", "token_type_ids"}
if input_names != expected_inputs:
raise DataError(
f"Core ML inputs changed: expected {sorted(expected_inputs)}, "
f"got {sorted(input_names)}"
)
activation_config, weight_config = _optimization_configs(cto)
args.output.parent.mkdir(parents=True, exist_ok=True)
with tempfile.TemporaryDirectory(
prefix=".purpose-coreml-calibration-",
dir=args.output.parent,
) as temporary:
with _bounded_calibration_packages(ModelDebugger, Path(temporary)):
activation_quantized = cto.coreml.linear_quantize_activations(
source_model,
activation_config,
sample_data,
calibration_op_group_size=args.calibration_op_group_size,
)
print("Core ML weight quantization: W8", flush=True)
quantized = cto.coreml.linear_quantize_weights(
activation_quantized,
weight_config,
)
quantized.user_defined_metadata["com.nucleic.model.quantization"] = "W8A8"
quantized.user_defined_metadata["com.nucleic.model.quantizationCalibration"] = (
f"stratified:{args.calibration_records}:seed={args.calibration_seed}"
)
if args.output.exists():
if args.output.is_dir():
shutil.rmtree(args.output)
else:
args.output.unlink()
quantized.save(str(args.output))
manifest = _package_manifest(args.output)
manifest.update(
{
"sourcePackage": str(args.model),
"sourcePackageSha256": _tree_sha256(args.model),
"coremltoolsVersion": ct.__version__,
"quantization": {
"name": "W8A8",
"activations": "per-tensor asymmetric uint8",
"linearWeights": "per-channel symmetric int8",
"embeddingWeights": "per-tensor asymmetric uint8",
},
"calibrationRecords": len(calibration),
"calibrationSeed": args.calibration_seed,
"calibrationOpGroupSize": args.calibration_op_group_size,
"calibrationSample": {
"strategy": "stratified by purpose, slice, and primary language",
"purposeCounts": dict(
sorted(Counter(item["purpose"] for item in calibration).items())
),
"sliceCounts": dict(
sorted(Counter(item["slice"] for item in calibration).items())
),
"promptHashes": sorted(
prompt_hash(item["prompt"]) for item in calibration
),
},
"shippingBudgetBytes": args.shipping_budget_bytes,
"shippingBudgetPassed": (
manifest["bytes"] <= args.shipping_budget_bytes
),
}
)
manifest_path = args.output.with_name(f"{args.output.stem}-manifest.json")
write_json(manifest_path, manifest)
print(
f"Core ML W8A8 package: {args.output} ({manifest['bytes']} bytes)",
flush=True,
)
if not manifest["shippingBudgetPassed"]:
raise DataError(
f"{args.output}: {manifest['bytes']} bytes exceeds the "
f"{args.shipping_budget_bytes}-byte shipping budget"
)
return manifest
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", type=Path, default=DEFAULT_MODEL)
parser.add_argument("--model-dir", type=Path, default=DEFAULT_MODEL_DIR)
parser.add_argument("--validation", type=Path, default=DEFAULT_VALIDATION)
parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument("--calibration-records", type=int, default=256)
parser.add_argument("--calibration-seed", type=int, default=20260730)
parser.add_argument("--calibration-op-group-size", type=int, default=-1)
parser.add_argument(
"--shipping-budget-bytes",
type=int,
default=SHIPPING_BUDGET_BYTES,
)
parser.add_argument("--overwrite-output", action="store_true")
return parser
def main(argv: Sequence[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
if (
args.calibration_records <= 0
or args.calibration_op_group_size == 0
or args.calibration_op_group_size < -1
or args.shipping_budget_bytes <= 0
):
parser.error(
"calibration records/budget must be positive and op group size must be "
"-1 or positive"
)
try:
quantize(args)
except (DataError, OSError, RuntimeError, ValueError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main())