diff --git a/README.md b/README.md index a8a8e14..4be24c2 100644 --- a/README.md +++ b/README.md @@ -237,7 +237,10 @@ ml/purpose-classifier/venv/bin/python ml/purpose-classifier/quantize_coreml.py \ Activation calibration writes its temporary packages under the candidate output directory and removes each package immediately after prediction; this avoids Core ML Tools retaining one full weight copy per calibration step until process exit. It prints progress while it -runs. The candidate uses per-tensor asymmetric uint8 activations, +runs. The successfully rewritten A8 package is cached beside the W8A8 output and reused +only when its source hash, Core ML Tools version, activation policy, calibration seed, and +prompt hashes match exactly. This prevents a later weight-stage failure from forcing +another calibration. The candidate uses per-tensor asymmetric uint8 activations, per-channel symmetric int8 linear weights, and per-tensor asymmetric uint8 embedding weights. Activation quantization is limited to floating-point linear operations; applying Core ML Tools' global policy also selects integer embedding-index additions and produces diff --git a/quantize_coreml.py b/quantize_coreml.py index 0505b2a..00905e5 100644 --- a/quantize_coreml.py +++ b/quantize_coreml.py @@ -6,6 +6,7 @@ from __future__ import annotations import argparse import gc import hashlib +import json import shutil import sys import tempfile @@ -73,8 +74,39 @@ def _bounded_calibration_packages(debugger_type: Any, temporary_root: Path): finally: debugger_type.predict_intermediate_outputs = original_predict tempfile.tempdir = previous_tempdir - for package in temporary_root.glob("*.mlpackage"): - shutil.rmtree(package) + + +def _activation_cache_contract( + source_sha256: str, + calibration: Sequence[dict[str, Any]], + *, + calibration_seed: int, +) -> dict[str, Any]: + return { + "schemaVersion": 1, + "sourcePackageSha256": source_sha256, + "coremltoolsVersion": COREMLTOOLS_VERSION, + "activationQuantization": "linear:per-tensor-asymmetric-uint8", + "activationOpTypes": ["linear"], + "calibrationRecords": len(calibration), + "calibrationSeed": calibration_seed, + "calibrationPromptHashes": sorted( + prompt_hash(item["prompt"]) for item in calibration + ), + } + + +def _activation_cache_is_valid( + package: Path, + manifest: Path, + expected: dict[str, Any], +) -> bool: + if not package.is_dir() or not manifest.is_file(): + return False + try: + return json.loads(manifest.read_text(encoding="utf-8")) == expected + except (OSError, UnicodeError, json.JSONDecodeError): + return False def _optimization_configs(optimize: Any) -> tuple[Any, Any]: @@ -160,24 +192,38 @@ def quantize(args: argparse.Namespace) -> dict[str, Any]: args.calibration_records, seed=args.calibration_seed, ) - tokenizer = AutoTokenizer.from_pretrained( - args.model_dir, - local_files_only=True, + source_sha256 = _tree_sha256(args.model) + activation_cache = args.output.with_name( + f"{args.output.stem}-a8-cache.mlpackage" + ) + activation_cache_manifest = activation_cache.with_name( + f"{activation_cache.stem}-manifest.json" + ) + cache_contract = _activation_cache_contract( + source_sha256, + calibration, + calibration_seed=args.calibration_seed, + ) + reuse_activation_cache = _activation_cache_is_valid( + activation_cache, + activation_cache_manifest, + cache_contract, ) 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() - } + if not reuse_activation_cache: + tokenizer = AutoTokenizer.from_pretrained( + args.model_dir, + local_files_only=True, ) + 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, @@ -196,34 +242,58 @@ def quantize(args: argparse.Namespace) -> dict[str, Any]: 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, + if reuse_activation_cache: + print( + f"Reusing Core ML A8 cache: {activation_cache}", + flush=True, + ) + activation_quantized = ct.models.MLModel( + str(activation_cache), + compute_units=ct.ComputeUnit.CPU_ONLY, + ) + else: + print( + f"Core ML activation calibration: {len(sample_data)} records", + flush=True, + ) + activation_quantized = cto.coreml.linear_quantize_activations( + source_model, + activation_config, + sample_data, + calibration_op_group_size=args.calibration_op_group_size, + ) + activation_quantized.save(str(activation_cache)) + write_json(activation_cache_manifest, cache_contract) + print( + f"Core ML A8 cache: {activation_cache}", + flush=True, + ) + print("Core ML weight quantization: W8", flush=True) + quantized = cto.coreml.linear_quantize_weights( + activation_quantized, + weight_config, ) - 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}" - ) + quantized.user_defined_metadata["com.nucleic.model.quantization"] = ( + "W8A8" + ) + quantized.user_defined_metadata[ + "com.nucleic.model.quantizationCalibration" + ] = f"stratified:{len(calibration)}: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)) + if args.output.exists(): + if args.output.is_dir(): + shutil.rmtree(args.output) + else: + args.output.unlink() + # Save while the Core ML Tools result package and its source weights are + # still alive inside the dedicated temporary directory. + quantized.save(str(args.output)) manifest = _package_manifest(args.output) manifest.update( { "sourcePackage": str(args.model), - "sourcePackageSha256": _tree_sha256(args.model), + "sourcePackageSha256": source_sha256, "coremltoolsVersion": ct.__version__, "quantization": { "name": "W8A8", diff --git a/tests/test_quantize_coreml.py b/tests/test_quantize_coreml.py index 8d4ed1f..68b8063 100644 --- a/tests/test_quantize_coreml.py +++ b/tests/test_quantize_coreml.py @@ -1,4 +1,5 @@ import argparse +import json import sys import tempfile import unittest @@ -94,6 +95,36 @@ class CoreMLQuantizationConfigTests(unittest.TestCase): self.assertEqual([1.0], output["output"].tolist()) self.assertEqual([], list(root.glob("*.mlpackage"))) + result_package = Path(tempfile.mkdtemp(suffix=".mlpackage")) + self.assertTrue(result_package.exists()) + + def test_activation_cache_requires_exact_contract(self): + expected = { + "schemaVersion": 1, + "sourcePackageSha256": "abc", + } + with tempfile.TemporaryDirectory() as temp: + root = Path(temp) + package = root / "a8.mlpackage" + package.mkdir() + manifest = root / "a8-manifest.json" + manifest.write_text(json.dumps(expected), encoding="utf-8") + self.assertTrue( + quantize_coreml._activation_cache_is_valid( + package, + manifest, + expected, + ) + ) + manifest.write_text("{}", encoding="utf-8") + self.assertFalse( + quantize_coreml._activation_cache_is_valid( + package, + manifest, + expected, + ) + ) + if __name__ == "__main__": unittest.main()