Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 60 additions & 5 deletions tensorflow/lite/micro/compression/lut.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -58,26 +58,34 @@ 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)

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}"
Expand All @@ -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)
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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(
Expand All @@ -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,
)
Expand Down
127 changes: 126 additions & 1 deletion tensorflow/lite/micro/compression/lut_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand All @@ -385,16 +493,33 @@ 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):
lut.LutAncillaryData(bitwidth=0)
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):
Expand Down
52 changes: 51 additions & 1 deletion tensorflow/lite/micro/compression/spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
"""

from dataclasses import dataclass
from typing import Optional, Union
import yaml

EXAMPLE_YAML_SPEC = """
Expand All @@ -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.

Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading