#!/usr/bin/env python3 """Calibrate a Core ML W8A8 candidate from the selected float16 ML Program.""" from __future__ import annotations import argparse import hashlib import shutil import sys from collections import Counter 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() 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 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) 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() args.output.parent.mkdir(parents=True, exist_ok=True) 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=32) 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())