97 lines
3.1 KiB
Python
97 lines
3.1 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.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)
|
|
|
|
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()
|