Files
nucleic-purpose-classifier/tests/test_quantize_coreml.py
T

100 lines
3.3 KiB
Python

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.assertIsNone(activation.global_config)
activation_linear = activation.op_type_configs["linear"].values
self.assertEqual("linear", activation_linear["mode"])
self.assertIs(np.uint8, activation_linear["dtype"])
self.assertEqual(
"per_tensor",
activation_linear["granularity"],
)
self.assertEqual({"linear"}, set(activation.op_type_configs))
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)
def test_calibration_packages_are_removed_after_each_prediction(self):
class FakeDebugger:
def predict_intermediate_outputs(self):
package = Path(tempfile.mkdtemp(suffix=".mlpackage"))
(package / "weight.bin").write_bytes(b"weights")
return {"output": np.array([1.0])}
with tempfile.TemporaryDirectory() as temp:
root = Path(temp) / "calibration"
with quantize_coreml._bounded_calibration_packages(
FakeDebugger,
root,
):
output = FakeDebugger().predict_intermediate_outputs()
self.assertEqual([1.0], output["output"].tolist())
self.assertEqual([], list(root.glob("*.mlpackage")))
if __name__ == "__main__":
unittest.main()