diff --git a/tensorflow/lite/micro/compression/lut.py b/tensorflow/lite/micro/compression/lut.py index d50baf3e413..06e9915d9c2 100644 --- a/tensorflow/lite/micro/compression/lut.py +++ b/tensorflow/lite/micro/compression/lut.py @@ -15,7 +15,7 @@ import sys from dataclasses import dataclass, field -from typing import Optional +from typing import ClassVar, Optional import bitarray import bitarray.util @@ -58,7 +58,8 @@ class LutAncillaryData: The LUT ancillary data uses the DCM user_data bytes (4-15) plus value tables: - Byte 4: LUT version (currently 1) - - Byte 5: Params (lower 3 bits = bitwidth, 1-7) + - Byte 5: Params (bits 7-4 = axis, 0-14 naming the channel axis of the + output tensor's shape, 15 meaning one table; bits 2-0 = bitwidth, 1-7) - Byte 6: Value table channel stride (elements per channel) - Bytes 7-15: Reserved (zeros) - Bytes 16+: Value tables (concatenated, stride elements per channel) @@ -66,18 +67,25 @@ class LutAncillaryData: Attributes: lut_version: LUT format version (currently 1). bitwidth: Number of bits per index (1-7). + axis: Byte 5 bits 7-4: the channel axis (0-14), or PER_TENSOR_AXIS. value_table_stride: Number of elements per channel in value tables. value_tables: Packed value table data following the DCM. """ + # Byte 5 axis field value meaning one value table for the whole tensor. + PER_TENSOR_AXIS: ClassVar[int] = 0xF + lut_version: int = 1 bitwidth: int = 4 + axis: int = 0 value_table_stride: int = 16 value_tables: bytes = b'' def __post_init__(self): if not 1 <= self.bitwidth <= 7: raise ValueError(f"bitwidth must be 1-7, got {self.bitwidth}") + if not 0 <= self.axis <= 15: + raise ValueError(f"axis must be 0-15, got {self.axis}") if not 0 <= self.value_table_stride <= 128: raise ValueError( f"value_table_stride must be 0-128, got {self.value_table_stride}" @@ -87,7 +95,7 @@ def to_user_data(self) -> bytes: """Serialize to 12-byte user_data for DCM bytes 4-15.""" user_data = bytearray(12) user_data[0] = self.lut_version - user_data[1] = self.bitwidth & 0x07 + user_data[1] = ((self.axis & 0x0F) << 4) | (self.bitwidth & 0x07) user_data[2] = self.value_table_stride # Bytes 3-11 (DCM bytes 7-15) remain zero (reserved) return bytes(user_data) @@ -178,6 +186,34 @@ def identify_compression_axis(tensor: model_editor.Tensor) -> Optional[int]: ) +def check_channel_axis(axis: int, shape: tuple[int, ...]): + """Validates a per-channel axis against the tensor shape and the kernels. + + Args: + axis: The axis named by the spec's per_channel mode. + shape: The shape of the tensor to be compressed. + + Raises: + CompressionError: If the axis is out of range for the shape, is an + axis the kernels do not support, or does not fit the DCM axis + field. + """ + rank = len(shape) + if not 0 <= axis < rank: + raise compressor.CompressionError( + f"per_channel axis {axis} out of range for a tensor of rank {rank}" + ) + if axis not in (0, rank - 1): + raise compressor.CompressionError( + f"per_channel axis {axis} unsupported: the kernels support " + f"axis 0 and the last axis only" + ) + if axis > 14: + raise compressor.CompressionError( + f"per_channel axis {axis} does not fit the DCM axis field (0-14)" + ) + + def pack_indices(indices: np.ndarray, bitwidth: int) -> bytes: """Packs indices into a bytearray using bitwidth-sized fields. @@ -253,8 +289,21 @@ def compress( raise compressor.CompressionError("Tensor has no data to compress") spec_bitwidth = method.index_bitwidth - axis = identify_compression_axis(tensor) - compressed = compress_array(tensor.array, axis) + + match method.mode: + case None: + compress_axis = identify_compression_axis(tensor) + case spec.PerTensor(): + compress_axis = None + case spec.PerChannel(axis=axis): + check_channel_axis(axis, tensor.shape) + compress_axis = axis + case _: + raise compressor.CompressionError( + f"unknown compression mode: {method.mode!r}" + ) + + compressed = compress_array(tensor.array, compress_axis) actual_bitwidth = compressed.index_bitwidth if actual_bitwidth > spec_bitwidth: raise compressor.CompressionError( @@ -277,9 +326,15 @@ def compress( value_tables_bytes = pack_lookup_tables(compressed.lookup_tables, table_len) # Build ancillary data + dcm_axis = ( + LutAncillaryData.PER_TENSOR_AXIS + if compress_axis is None + else compress_axis + ) lut_data = LutAncillaryData( lut_version=1, bitwidth=spec_bitwidth, + axis=dcm_axis, value_table_stride=table_len, value_tables=value_tables_bytes, ) diff --git a/tensorflow/lite/micro/compression/lut_test.py b/tensorflow/lite/micro/compression/lut_test.py index 0fb50e4490b..0fd890e5b90 100644 --- a/tensorflow/lite/micro/compression/lut_test.py +++ b/tensorflow/lite/micro/compression/lut_test.py @@ -370,6 +370,114 @@ def test_compress_no_data_raises(self): compressor_instance.compress(tensor, method) +class TestCompressionMode(unittest.TestCase): + """Tests for the explicit per-tensor/per-channel compression mode.""" + + def _compress(self, tensor, mode, bitwidth=4): + method = spec.LookUpTableCompression(index_bitwidth=bitwidth, mode=mode) + return lut.LutCompressor().compress(tensor, method) + + def _per_tensor_quantized(self): + return model_editor.Tensor( + shape=(4,), + dtype=tflite.TensorType.INT8, + data=np.array([1, 2, 3, 4], dtype=np.int8), + quantization=model_editor.Quantization(scales=1.0, zero_points=0), + ) + + def _per_channel_quantized(self): + # Rows hold one unique value each, so per-channel compression along + # axis 0 yields 4 single-entry tables. + return model_editor.Tensor( + shape=(4, 2), + dtype=tflite.TensorType.INT8, + data=np.array([[1, 1], [5, 5], [9, 9], [13, 13]], dtype=np.int8), + quantization=model_editor.Quantization( + scales=[0.1, 0.2, 0.3, 0.4], + zero_points=[0, 0, 0, 0], + axis=0, + ), + ) + + def test_no_mode_writes_axis_15(self): + """The inference path states the resolved layout in byte 5 bits 7-4.""" + result = self._compress(self._per_tensor_quantized(), mode=None) + self.assertEqual(result.ancillary_data[5], 0xF4) + + def test_per_tensor_writes_axis_15(self): + result = self._compress(self._per_tensor_quantized(), spec.PerTensor()) + self.assertEqual(result.ancillary_data[5], 0xF4) + + def test_per_channel_writes_axis(self): + result = self._compress( + self._per_channel_quantized(), spec.PerChannel(axis=0) + ) + self.assertEqual(result.ancillary_data[5], 0x04) + # One single-entry table per channel + self.assertEqual(result.ancillary_data[6], 1) + self.assertEqual(len(result.ancillary_data), 16 + 4) + + def test_per_channel_last_axis_writes_axis(self): + tensor = model_editor.Tensor( + shape=(2, 3), + dtype=tflite.TensorType.INT8, + data=np.array([[1, 5, 9], [1, 5, 9]], dtype=np.int8), + quantization=model_editor.Quantization( + scales=[0.1, 0.2, 0.3], + zero_points=[0, 0, 0], + axis=1, + ), + ) + result = self._compress(tensor, spec.PerChannel(axis=1)) + self.assertEqual(result.ancillary_data[5], 0x14) + self.assertEqual(result.ancillary_data[6], 1) + self.assertEqual(len(result.ancillary_data), 16 + 3) + + def test_per_channel_on_unquantized_tensor(self): + """An explicit axis works without quantization.""" + tensor = model_editor.Tensor( + shape=(2, 2), + dtype=tflite.TensorType.INT8, + data=np.array([[1, 1], [5, 5]], dtype=np.int8), + ) + result = self._compress(tensor, spec.PerChannel(axis=0), bitwidth=1) + self.assertEqual(result.ancillary_data[5], 0x01) + self.assertEqual(result.ancillary_data[6], 1) + + def test_per_tensor_overrides_quantization(self): + result = self._compress(self._per_channel_quantized(), spec.PerTensor()) + self.assertEqual(result.ancillary_data[5], 0xF4) + # One table holding all 4 unique values + self.assertEqual(result.ancillary_data[6], 4) + + def test_per_channel_axis_differs_from_quantization(self): + result = self._compress( + self._per_channel_quantized(), spec.PerChannel(axis=1) + ) + self.assertEqual(result.ancillary_data[5], 0x14) + + def test_axis_out_of_range_raises(self): + tensor = self._per_channel_quantized() + with self.assertRaises(compressor.CompressionError): + self._compress(tensor, spec.PerChannel(axis=2)) + with self.assertRaises(compressor.CompressionError): + self._compress(tensor, spec.PerChannel(axis=-1)) + + def test_middle_axis_raises(self): + """The kernels support axis 0 and the last axis only.""" + tensor = model_editor.Tensor( + shape=(2, 2, 2), + dtype=tflite.TensorType.INT8, + data=np.zeros((2, 2, 2), dtype=np.int8), + ) + with self.assertRaises(compressor.CompressionError): + self._compress(tensor, spec.PerChannel(axis=1)) + + def test_unknown_mode_raises(self): + with self.assertRaises(compressor.CompressionError): + self._compress(self._per_tensor_quantized(), mode=object()) + + class TestLutAncillaryData(unittest.TestCase): """Tests for LutAncillaryData.""" @@ -385,9 +493,19 @@ def test_to_user_data_format(self): self.assertEqual(len(user_data), 12) self.assertEqual(user_data[0], 1) # lut_version - self.assertEqual(user_data[1], 4) # bitwidth + self.assertEqual(user_data[1], 4) # axis 0, bitwidth 4 self.assertEqual(user_data[2], 16) # stride + def test_to_user_data_axis_field(self): + """The axis occupies byte 5 bits 7-4, above the bitwidth.""" + lut_data = lut.LutAncillaryData( + bitwidth=3, axis=lut.LutAncillaryData.PER_TENSOR_AXIS + ) + self.assertEqual(lut_data.to_user_data()[1], 0xF3) + + lut_data = lut.LutAncillaryData(bitwidth=3, axis=2) + self.assertEqual(lut_data.to_user_data()[1], 0x23) + def test_bitwidth_validation(self): """Bitwidth must be 1-7.""" with self.assertRaises(ValueError): @@ -395,6 +513,13 @@ def test_bitwidth_validation(self): with self.assertRaises(ValueError): lut.LutAncillaryData(bitwidth=8) + def test_axis_validation(self): + """Axis must fit the 4-bit field.""" + with self.assertRaises(ValueError): + lut.LutAncillaryData(axis=-1) + with self.assertRaises(ValueError): + lut.LutAncillaryData(axis=16) + def test_stride_validation(self): """Stride must be 0-128.""" with self.assertRaises(ValueError): diff --git a/tensorflow/lite/micro/compression/spec.py b/tensorflow/lite/micro/compression/spec.py index 8d3a4d308db..f71f96f916d 100644 --- a/tensorflow/lite/micro/compression/spec.py +++ b/tensorflow/lite/micro/compression/spec.py @@ -23,6 +23,7 @@ """ from dataclasses import dataclass +from typing import Optional, Union import yaml EXAMPLE_YAML_SPEC = """ @@ -33,12 +34,15 @@ compression: - lut: index_bitwidth: 4 + per_channel: + axis: 0 - subgraph: 0 tensor: 55 compression: - lut: index_bitwidth: 2 + per_tensor: """ # This example is checked in this module's unit test. @@ -56,15 +60,34 @@ class Tensor: compression: list[CompressionMethod] +@dataclass +class PerTensor: + """One value table for the whole tensor.""" + + +@dataclass +class PerChannel: + """One value table per slice along an axis. + + Attributes: + axis: The axis of the tensor's shape that gives the channel count. + """ + + axis: int + + @dataclass class LookUpTableCompression(CompressionMethod): """LUT compression using lookup tables. Attributes: index_bitwidth: Number of bits per index (1-7). + mode: PerTensor or PerChannel. None means the compressor infers the + mode from the tensor's quantization. """ index_bitwidth: int + mode: Optional[Union[PerTensor, PerChannel]] = None @dataclass @@ -95,10 +118,37 @@ def __init__(self, message="error parsing spec", wrapped_exception=None): self.original_exception = wrapped_exception +def _parse_lut(lut: dict) -> LookUpTableCompression: + """Parse a lut compression entry from its YAML dict.""" + has_per_tensor = "per_tensor" in lut + has_per_channel = "per_channel" in lut + if has_per_tensor and has_per_channel: + raise ParseError( + "lut: per_tensor and per_channel are contradictory; give at most one" + ) + + if has_per_tensor: + if lut["per_tensor"] is not None: + raise ParseError("lut: per_tensor takes no value") + mode = PerTensor() + elif has_per_channel: + per_channel = lut["per_channel"] + if not isinstance(per_channel, dict) or "axis" not in per_channel: + raise ParseError("lut: per_channel requires an axis") + axis = per_channel["axis"] + if not isinstance(axis, int) or isinstance(axis, bool) or axis < 0: + raise ParseError("lut: per_channel axis must be a non-negative integer") + mode = PerChannel(axis=axis) + else: + mode = None + + return LookUpTableCompression(index_bitwidth=lut["index_bitwidth"], mode=mode) + + def _parse_compression_method(comp: dict) -> CompressionMethod: """Parse a single compression method from YAML dict.""" if "lut" in comp: - return LookUpTableCompression(index_bitwidth=comp["lut"]["index_bitwidth"]) + return _parse_lut(comp["lut"]) elif "huffman" in comp: return HuffmanCompression() elif "pruning" in comp: diff --git a/tensorflow/lite/micro/compression/spec_builder.py b/tensorflow/lite/micro/compression/spec_builder.py index 0df26253b74..509d025bb0c 100644 --- a/tensorflow/lite/micro/compression/spec_builder.py +++ b/tensorflow/lite/micro/compression/spec_builder.py @@ -28,7 +28,7 @@ .build()) """ -from typing import List, Optional +from typing import List, Optional, Union from . import spec @@ -41,17 +41,23 @@ def __init__(self, subgraph: int, tensor: int, parent_builder: 'SpecBuilder'): self.compression_methods: List[spec.CompressionMethod] = [] self._parent = parent_builder - def with_lut(self, index_bitwidth: int) -> 'SpecBuilder': + def with_lut( + self, + index_bitwidth: int, + mode: Optional[Union[spec.PerTensor, spec.PerChannel]] = None, + ) -> 'SpecBuilder': """Add LUT compression to this tensor. Args: index_bitwidth: Number of bits for the LUT index (e.g., 4 for 16 values) + mode: spec.PerTensor or spec.PerChannel. None leaves the + choice to inference from the tensor's quantization. Returns: The parent SpecBuilder for method chaining """ self.compression_methods.append( - spec.LookUpTableCompression(index_bitwidth=index_bitwidth) + spec.LookUpTableCompression(index_bitwidth=index_bitwidth, mode=mode) ) return self._parent diff --git a/tensorflow/lite/micro/compression/spec_builder_test.py b/tensorflow/lite/micro/compression/spec_builder_test.py index 05f8cb86e69..fc3692bb7f8 100644 --- a/tensorflow/lite/micro/compression/spec_builder_test.py +++ b/tensorflow/lite/micro/compression/spec_builder_test.py @@ -80,6 +80,31 @@ def test_single_tensor(self): self.assertEqual(result[0].tensor, 42) self.assertEqual(result[0].compression[0].index_bitwidth, 16) + def test_mode_passes_through(self): + """The mode argument reaches the spec object unchanged.""" + result = ( + spec_builder.SpecBuilder() + .add_tensor(subgraph=0, tensor=1) + .with_lut(index_bitwidth=4, mode=spec.PerChannel(axis=0)) + .add_tensor(subgraph=0, tensor=2) + .with_lut(index_bitwidth=2, mode=spec.PerTensor()) + .build() + ) + + self.assertEqual(result[0].compression[0].mode, spec.PerChannel(axis=0)) + self.assertEqual(result[1].compression[0].mode, spec.PerTensor()) + + def test_mode_defaults_to_none(self): + """Omitting the mode leaves the choice to inference.""" + result = ( + spec_builder.SpecBuilder() + .add_tensor(subgraph=0, tensor=1) + .with_lut(index_bitwidth=4) + .build() + ) + + self.assertIsNone(result[0].compression[0].mode) + def test_tensor_without_compression(self): """Test that tensors can be added without compression methods.""" builder = spec_builder.SpecBuilder() diff --git a/tensorflow/lite/micro/compression/spec_test.py b/tensorflow/lite/micro/compression/spec_test.py index 3a740a30879..ab32a479afa 100644 --- a/tensorflow/lite/micro/compression/spec_test.py +++ b/tensorflow/lite/micro/compression/spec_test.py @@ -21,36 +21,85 @@ spec.Tensor( subgraph=0, tensor=42, - compression=[spec.LookUpTableCompression(index_bitwidth=4)], + compression=[ + spec.LookUpTableCompression( + index_bitwidth=4, mode=spec.PerChannel(axis=0) + ) + ], ), spec.Tensor( subgraph=0, tensor=55, - compression=[spec.LookUpTableCompression(index_bitwidth=2)], + compression=[ + spec.LookUpTableCompression(index_bitwidth=2, mode=spec.PerTensor()) + ], ), ] -class TestLoadYaml(unittest.TestCase): - def testExampleSpec(self): - result = spec.parse_yaml(spec.EXAMPLE_YAML_SPEC) - self.assertEqual(result, EXPECTED_PYTHON_SPEC) +def _lut_spec(*lut_lines: str) -> str: + """Returns a one-tensor spec whose lut entry holds the given lines.""" + lines = [ + "tensors:", + " - subgraph: 0", + " tensor: 0", + " compression:", + " - lut:", + ] + lines += [" " * 10 + line for line in lut_lines] + return "\n".join(lines) + "\n" - def testMalformedYAML(self): - bad = spec.EXAMPLE_YAML_SPEC + " & foobar: 0" - self.assertRaises(spec.ParseError, lambda: spec.parse_yaml(bad)) - def testUnexpectedType(self): - bad = spec.EXAMPLE_YAML_SPEC + " - subgraph: 'foobar'" - self.assertRaises(spec.ParseError, lambda: spec.parse_yaml(bad)) +class TestLutMode(unittest.TestCase): + """Tests for parsing the per_tensor/per_channel choice.""" - def testMissingFields(self): - bad = spec.EXAMPLE_YAML_SPEC + " - foobar: 0" - self.assertRaises(spec.ParseError, lambda: spec.parse_yaml(bad)) + def testMissingModeIsNone(self): + result = spec.parse_yaml(_lut_spec("index_bitwidth: 4")) + self.assertIsNone(result[0].compression[0].mode) - def testIgnoreExtraKeys(self): - result = spec.parse_yaml(spec.EXAMPLE_YAML_SPEC + "foobar: 0") - self.assertEqual(result, EXPECTED_PYTHON_SPEC) + def testBothModesRaise(self): + bad = _lut_spec( + "index_bitwidth: 4", + "per_tensor:", + "per_channel:", + " axis: 0", + ) + with self.assertRaisesRegex(spec.ParseError, "contradictory"): + spec.parse_yaml(bad) + + def testPerTensorWithPayloadRaises(self): + bad = _lut_spec( + "index_bitwidth: 4", + "per_tensor: 1", + ) + with self.assertRaisesRegex(spec.ParseError, "no value"): + spec.parse_yaml(bad) + + def testPerChannelWithoutAxisRaises(self): + bad = _lut_spec( + "index_bitwidth: 4", + "per_channel:", + ) + with self.assertRaisesRegex(spec.ParseError, "axis"): + spec.parse_yaml(bad) + + def testNegativeAxisRaises(self): + bad = _lut_spec( + "index_bitwidth: 4", + "per_channel:", + " axis: -1", + ) + with self.assertRaisesRegex(spec.ParseError, "non-negative"): + spec.parse_yaml(bad) + + def testNonIntegerAxisRaises(self): + bad = _lut_spec( + "index_bitwidth: 4", + "per_channel:", + " axis: zero", + ) + with self.assertRaisesRegex(spec.ParseError, "non-negative"): + spec.parse_yaml(bad) if __name__ == "__main__": diff --git a/tensorflow/lite/micro/docs/compression.md b/tensorflow/lite/micro/docs/compression.md index 54e1533dc1c..f5331ba2a25 100644 --- a/tensorflow/lite/micro/docs/compression.md +++ b/tensorflow/lite/micro/docs/compression.md @@ -300,18 +300,23 @@ tensors: compression: - lut: index_bitwidth: 4 + per_channel: + axis: 0 - subgraph: 0 tensor: 10 compression: - lut: index_bitwidth: 4 + per_channel: + axis: 0 - subgraph: 0 tensor: 11 compression: - lut: index_bitwidth: 2 + per_tensor: - subgraph: 0 tensor: 22 @@ -321,6 +326,14 @@ tensors: ``` Note that each tensor can have a different bit width (1 through 7 bits). +A `lut` entry may state its compression mode, one of `per_channel` or +`per_tensor`. `per_channel` builds one value table per channel, along the +given axis of the tensor's shape. A bare `per_tensor:` builds one value +table for the whole tensor. An entry without a mode, like tensor 22 above, +takes the mode from the tensor's quantization: per-channel along the +quantized axis when the tensor has one scale per channel, otherwise +per-tensor. + Once the `YAML` specification is ready, compress the model using the following: ``` bazel run -s tensorflow/lite/micro/compression:compress -- --input=binned.tflite --output=compressed.tflite --spec=spec.yaml