Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 21:08:05 -07:00
parent 9e462a79fc
commit 40176d0c76
4 changed files with 373 additions and 3 deletions
+79
View File
@@ -0,0 +1,79 @@
import argparse
import sys
import tempfile
import unittest
from pathlib import Path
import numpy as np
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import quantize_coreml
from purpose_data import DataError
class FakeOpLinearQuantizerConfig:
def __init__(self, **values):
self.values = values
class FakeOptimizationConfig:
def __init__(self, *, global_config=None, op_type_configs=None):
self.global_config = global_config
self.op_type_configs = op_type_configs or {}
class FakeCoreML:
OpLinearQuantizerConfig = FakeOpLinearQuantizerConfig
OptimizationConfig = FakeOptimizationConfig
class FakeOptimize:
coreml = FakeCoreML
class CoreMLQuantizationConfigTests(unittest.TestCase):
def test_matches_accepted_qdq_policy(self):
activation, weights = quantize_coreml._optimization_configs(FakeOptimize)
self.assertEqual("linear", activation.global_config.values["mode"])
self.assertIs(np.uint8, activation.global_config.values["dtype"])
self.assertEqual(
"per_tensor",
activation.global_config.values["granularity"],
)
linear = weights.op_type_configs["linear"].values
self.assertEqual("linear_symmetric", linear["mode"])
self.assertIs(np.int8, linear["dtype"])
self.assertEqual("per_channel", linear["granularity"])
self.assertIs(
weights.op_type_configs["linear"],
weights.op_type_configs["matmul"],
)
embedding = weights.op_type_configs["gather"].values
self.assertEqual("linear", embedding["mode"])
self.assertIs(np.uint8, embedding["dtype"])
self.assertEqual("per_tensor", embedding["granularity"])
def test_rejects_overwriting_source_package(self):
with tempfile.TemporaryDirectory() as temp:
root = Path(temp)
package = root / "model.mlpackage"
package.mkdir()
model_dir = root / "model"
model_dir.mkdir()
args = argparse.Namespace(
model=package,
model_dir=model_dir,
output=package,
overwrite_output=True,
)
with self.assertRaisesRegex(DataError, "must not overwrite"):
quantize_coreml._validate_args(args)
if __name__ == "__main__":
unittest.main()