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