Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user