From f90719640870b63f417995d3c2a3c687f8b8117d Mon Sep 17 00:00:00 2001 From: asteinh Date: Sun, 4 Oct 2026 20:51:07 +0200 Subject: [PATCH] feat: compile TFLite SVDF with its state --- scripts/crossrepo_contract.py | 52 ++++++++ scripts/gen_tflite_fixtures.py | 116 ++++++++++++++++-- scripts/generate_schema_package.py | 3 + src/tigris/analysis/validation.py | 3 +- src/tigris/capabilities.py | 2 + src/tigris/dtypes.py | 24 +++- src/tigris/emitters/binary/defs.py | 7 ++ src/tigris/emitters/binary/writer.py | 54 ++++++-- src/tigris/emitters/codegen.py | 6 +- src/tigris/frontends/tflite.py | 103 +++++++++++++++- src/tigris/graph/ir.py | 10 ++ src/tigris/loaders/onnx/normalize.py | 35 ++++++ .../schema/operator-capabilities-v1.json | 13 ++ src/tigris/schema/tigris-plan-v10.json | 4 + src/tigris/utils.py | 2 +- tests/fixtures/tflite/ops/float_svdf.npz | Bin 0 -> 609 bytes tests/fixtures/tflite/ops/float_svdf.tflite | Bin 0 -> 808 bytes .../tflite/ops/float_svdf_rank2_relu.npz | Bin 0 -> 668 bytes .../tflite/ops/float_svdf_rank2_relu.tflite | Bin 0 -> 848 bytes tests/fixtures/tflite/ops/svdf.npz | Bin 0 -> 457 bytes tests/fixtures/tflite/ops/svdf.tflite | Bin 0 -> 1016 bytes tests/fixtures/tflite/ops/svdf_batch.npz | Bin 0 -> 483 bytes tests/fixtures/tflite/ops/svdf_batch.tflite | Bin 0 -> 912 bytes tests/test_frontend_tflite.py | 33 ++++- 24 files changed, 432 insertions(+), 35 deletions(-) create mode 100644 tests/fixtures/tflite/ops/float_svdf.npz create mode 100644 tests/fixtures/tflite/ops/float_svdf.tflite create mode 100644 tests/fixtures/tflite/ops/float_svdf_rank2_relu.npz create mode 100644 tests/fixtures/tflite/ops/float_svdf_rank2_relu.tflite create mode 100644 tests/fixtures/tflite/ops/svdf.npz create mode 100644 tests/fixtures/tflite/ops/svdf.tflite create mode 100644 tests/fixtures/tflite/ops/svdf_batch.npz create mode 100644 tests/fixtures/tflite/ops/svdf_batch.tflite diff --git a/scripts/crossrepo_contract.py b/scripts/crossrepo_contract.py index dbd40f9..be1370d 100644 --- a/scripts/crossrepo_contract.py +++ b/scripts/crossrepo_contract.py @@ -11,6 +11,7 @@ import argparse import copy +import json import math import re import subprocess @@ -41,6 +42,7 @@ from tigris.emitters.binary.reader import read_binary_plan from tigris.emitters.binary.writer import emit_binary from tigris.fixtures import build_tcn +from tigris.frontends.tflite import STATE_KEY from tigris.graph.ir import Stage # The byte-level line-buffer flag decoder already exists in the compiler's @@ -844,6 +846,55 @@ def build(body): +def _svdf_case() -> ContractCase: + """A float SVDF with a fused Relu, run once from its zero initial state, so + each filter's time projection sees only the newest slot of its memory.""" + rng = np.random.default_rng(85) + batch, features, units, rank, memory = 2, 6, 3, 2, 4 + filters = units * rank + feature = rng.normal(0.0, 0.5, (filters, features)).astype(np.float32) + time = rng.normal(0.0, 0.5, (filters, memory)).astype(np.float32) + bias = rng.normal(0.0, 0.3, (units,)).astype(np.float32) + x = rng.uniform(-2.0, 2.0, (batch, features)).astype(np.float32) + shape_x, shape_y, shape_s = [batch, features], [batch, units], [batch, filters * memory] + weights = [numpy_helper.from_array(feature, "feature"), numpy_helper.from_array(time, "time"), + numpy_helper.from_array(bias, "bias"), + numpy_helper.from_array(np.zeros(shape_s, np.float32), "state_initial")] + compile_graph = helper.make_graph( + [helper.make_node("Svdf", ["x", "feature", "time", "bias", "state_in"], ["y", "state_out"], + domain="tigris", rank=rank, activation="relu")], + "svdf", + [helper.make_tensor_value_info("x", TensorProto.FLOAT, shape_x), + helper.make_tensor_value_info("state_in", TensorProto.FLOAT, shape_s)], + [helper.make_tensor_value_info("y", TensorProto.FLOAT, shape_y), + helper.make_tensor_value_info("state_out", TensorProto.FLOAT, shape_s)], + weights, + value_info=[helper.make_tensor_value_info("y", TensorProto.FLOAT, shape_y), + helper.make_tensor_value_info("state_out", TensorProto.FLOAT, shape_s)]) + compile_model = helper.make_model(compile_graph, opset_imports=[ + helper.make_opsetid("", 17), helper.make_opsetid("tigris", 1)]) + compile_model.ir_version = 9 + helper.set_model_props(compile_model, {STATE_KEY: json.dumps( + [{"input": "state_in", "output": "state_out", "initial": "state_initial"}])}) + reference = _model( + "svdf_reference", + [helper.make_node("MatMul", ["x", "feature_t"], ["projection"]), + helper.make_node("Mul", ["projection", "newest"], ["weighted"]), + helper.make_node("Reshape", ["weighted", "grouped_shape"], ["grouped"]), + helper.make_node("ReduceSum", ["grouped", "rank_axis"], ["summed"], keepdims=0), + helper.make_node("Add", ["summed", "bias"], ["biased"]), + helper.make_node("Relu", ["biased"], ["y"])], + [helper.make_tensor_value_info("x", TensorProto.FLOAT, shape_x)], + [helper.make_tensor_value_info("y", TensorProto.FLOAT, shape_y)], + [numpy_helper.from_array(np.ascontiguousarray(feature.T), "feature_t"), + numpy_helper.from_array(np.ascontiguousarray(time[:, -1]), "newest"), + numpy_helper.from_array(np.asarray([batch, units, rank], np.int64), "grouped_shape"), + numpy_helper.from_array(np.asarray([2], np.int64), "rank_axis"), + numpy_helper.from_array(bias, "bias")]) + return ContractCase(name="svdf_relu", compile_model=compile_model, reference_model=reference, + inputs={"x": x}, expected_operators=("Svdf",)) + + def _runtime_index_case(kind: str, quantized: bool, variant: int = 0, arg: str | None = None, cast: bool = False) -> ContractCase: case = _movement_case(kind, quantized, variant) @@ -6918,6 +6969,7 @@ def _run_gate(runtime: Path, work_dir: Path) -> None: _convtranspose_2d_tiled_case(), _qdq_convtranspose_2d_tiled_case(), _convtranspose_2d_partial_edge_case(), + _svdf_case(), ] covered_operators = { operator for case in cases for operator in case.expected_operators diff --git a/scripts/gen_tflite_fixtures.py b/scripts/gen_tflite_fixtures.py index 4142ed7..04e14db 100644 --- a/scripts/gen_tflite_fixtures.py +++ b/scripts/gen_tflite_fixtures.py @@ -469,6 +469,104 @@ def _embedding_lookup(model: bytes) -> bytes: REWRITES["embedding_lookup"] = _embedding_lookup REWRITES["embedding_lookup_runtime"] = _embedding_lookup +def _one_operator(code, options_type, options, tensors, op_inputs, op_outputs, inputs, outputs): + """A model of one builtin operator. `tensors` are (shape, type, data, scale, + zero point, variable) records; data None marks an activation.""" + tree = schema.ModelT() + tree.version = 3 + tree.buffers = [schema.BufferT()] + graph = schema.SubGraphT() + graph.tensors = [] + for i, (shape, kind, data, scale, zero_point, variable) in enumerate(tensors): + tensor = schema.TensorT() + tensor.name = f"t{i}".encode() + tensor.shape = np.asarray(shape, np.int32) + tensor.type = kind + tensor.isVariable = variable + tensor.buffer = len(tree.buffers) + buffer = schema.BufferT() + if data is not None: + buffer.data = np.frombuffer(np.ascontiguousarray(data).tobytes(), np.uint8) + tree.buffers.append(buffer) + if scale is not None: + quant = schema.QuantizationParametersT() + quant.scale = np.asarray([scale], np.float32) + quant.zeroPoint = np.asarray([zero_point], np.int64) + tensor.quantization = quant + graph.tensors.append(tensor) + op = schema.OperatorT() + op.opcodeIndex = 0 + op.inputs = np.asarray(op_inputs, np.int32) + op.outputs = np.asarray(op_outputs, np.int32) + op.builtinOptionsType = options_type + op.builtinOptions = options + graph.operators = [op] + graph.inputs = np.asarray(inputs, np.int32) + graph.outputs = np.asarray(outputs, np.int32) + tree.subgraphs = [graph] + opcode = schema.OperatorCodeT() + opcode.builtinCode = code + opcode.deprecatedBuiltinCode = min(code, 127) + opcode.version = 1 + tree.operatorCodes = [opcode] + builder = flatbuffers.Builder(4096) + builder.Finish(tree.Pack(builder), file_identifier=b"TFL3") + return bytes(builder.Output()) + + +def _quantize(values, scale, dtype): + info = np.iinfo(dtype) + return np.clip(np.round(values / scale), info.min, info.max).astype(dtype) + + +def _svdf(batch, features, units, rank, memory, activation, quantized=False): + """SVDF on a [batch, features] input, its state a variable tensor; seeded + weights. The int8 form keeps its state and time weights in int16. TFLite + Micro's Prepare dereferences the bias, so every case has one.""" + rng = np.random.default_rng(features * 100 + units * 10 + rank) + filters = units * rank + feature = rng.normal(0.0, 0.4, (filters, features)).astype(np.float32) + time = rng.normal(0.0, 0.4, (filters, memory)).astype(np.float32) + offsets = rng.normal(0.0, 0.3, (units,)).astype(np.float32) + options = schema.SVDFOptionsT() + options.rank = rank + options.fusedActivationFunction = activation + T = schema.TensorType + if not quantized: + tensors = [((batch, features), T.FLOAT32, None, None, 0, False), + ((filters, features), T.FLOAT32, feature, None, 0, False), + ((filters, memory), T.FLOAT32, time, None, 0, False), + ((units,), T.FLOAT32, offsets, None, 0, False), + ((batch, filters * memory), T.FLOAT32, None, None, 0, True), + ((batch, units), T.FLOAT32, None, None, 0, False)] + else: + x_scale, state_scale, out_scale = 6.0 / 255, 8.0 / 32767, 0.06 + f_scale = float(np.abs(feature).max()) / 127 + t_scale = float(np.abs(time).max()) / 32767 + b_scale = float(np.float32(state_scale) * np.float32(t_scale)) + tensors = [((batch, features), T.INT8, None, x_scale, 2, False), + ((filters, features), T.INT8, _quantize(feature, f_scale, np.int8), f_scale, 0, False), + ((filters, memory), T.INT16, _quantize(time, t_scale, np.int16), t_scale, 0, False), + ((units,), T.INT32, _quantize(offsets, b_scale, np.int32), b_scale, 0, False), + ((batch, filters * memory), T.INT16, None, state_scale, 0, True), + ((batch, units), T.INT8, None, out_scale, -3, False)] + op_inputs = [0, 1, 2, 3, 4] + return _one_operator(schema.BuiltinOperator.SVDF, schema.BuiltinOptions.SVDFOptions, + options, tensors, op_inputs, [5], [0], [5]) + + +# Models no converter writes, built operator by operator. +_RELU = schema.ActivationFunctionType.RELU +_NONE = schema.ActivationFunctionType.NONE +HANDMADE = { + "float_svdf": lambda: _svdf(1, 8, 4, 1, 5, _NONE), + "float_svdf_rank2_relu": lambda: _svdf(2, 6, 3, 2, 4, _RELU), + "svdf": lambda: _svdf(1, 8, 4, 2, 5, _NONE, quantized=True), + "svdf_batch": lambda: _svdf(2, 6, 3, 1, 4, _NONE, quantized=True), +} +# TFLite's float SVDF sums in another order than TFLite Micro, so the two +# differ in the last bits; TFLite Micro's outputs are recorded for these. +SUMMATION_ORDER = {"float_svdf", "float_svdf_rank2_relu"} # The tier-1 cases again, converted without quantization. _FLOAT_TIER1 = ( "max_pool_valid", "max_pool_same", "avg_pool_valid", "avg_pool_same", "concat_channels", @@ -540,13 +638,16 @@ def _inputs(details, rng, value_range=None, index_range=None): def generate(name: str) -> bool: """Writes the case; True when TFLite Micro deviates from the reference kernels.""" - shapes, fn = CASES[name] rng = np.random.default_rng(sum(name.encode())) ranges = RANGES.get(name) - model = _convert(fn, shapes, ranges or [(-3.0, 3.0)] * len(shapes), rng, - float_io=name in FLOAT_BOUNDARIES, quantize=name not in FLOAT_MODELS, - trackable=TRACKABLES.get(name.removeprefix("float_")), - index_inputs=INDEX_INPUTS.get(name.removeprefix("float_"))) + if name in HANDMADE: + model = HANDMADE[name]() + else: + shapes, fn = CASES[name] + model = _convert(fn, shapes, ranges or [(-3.0, 3.0)] * len(shapes), rng, + float_io=name in FLOAT_BOUNDARIES, quantize=name not in FLOAT_MODELS, + trackable=TRACKABLES.get(name.removeprefix("float_")), + index_inputs=INDEX_INPUTS.get(name.removeprefix("float_"))) if name in REWRITES: model = REWRITES[name](model) micro_interpreter = micro.Interpreter.from_bytes(model, arena_size=1024 * 1024) @@ -597,7 +698,8 @@ def generate(name: str) -> bool: continue reference_outputs.append([reference.get_tensor(d["index"]).copy() for d in reference.get_output_details()]) - has_reference = reference is not None and name not in BROKEN_REFERENCE + has_reference = (reference is not None and name not in BROKEN_REFERENCE + and name not in SUMMATION_ORDER) deviates = has_reference and any(not np.array_equal(m, r) for ms, rs in zip(micro_outputs, reference_outputs) for m, r in zip(ms, rs)) outputs = reference_outputs if deviates else micro_outputs @@ -637,7 +739,7 @@ def main(names: list[str]) -> None: for name in REFERENCE_MODELS: generate_model(name) return - deviations = [name for name in names or sorted(CASES) if generate(name)] + deviations = [name for name in names or sorted({*CASES, *HANDMADE}) if generate(name)] if deviations: print("recorded from the reference kernels: " + ", ".join(deviations)) print(f"tensorflow {tf.__version__}, tflite-micro {version('tflite-micro')}") diff --git a/scripts/generate_schema_package.py b/scripts/generate_schema_package.py index 24792c6..92e0543 100644 --- a/scripts/generate_schema_package.py +++ b/scripts/generate_schema_package.py @@ -33,6 +33,8 @@ def schema_package() -> dict[str, object]: "cumsum_options": defs.OP_ATTR_CUMSUM_OPTIONS, "movement": defs.OP_ATTR_MOVEMENT, "comparison_requant": defs.OP_ATTR_COMPARISON_REQUANT, + "constants": defs.OP_ATTR_CONSTANTS, + "svdf": defs.OP_ATTR_SVDF, "epsilon": defs.OP_ATTR_EPSILON, "pads": defs.OP_ATTR_PADS, "pool_rounding": defs.OP_ATTR_POOL_ROUNDING, @@ -66,6 +68,7 @@ def schema_package() -> dict[str, object]: "stage table authoritative; operator byte is canonical low-byte hint" ), "state_v10": "only a plan with a state section is written at schema 10", + "state_int16": "int16 (ONNX dtype 5) is stored only on state tensors", }, "section_alignment": defs.PLAN_SECTION_ALIGNMENT, "tensor_flags": { diff --git a/src/tigris/analysis/validation.py b/src/tigris/analysis/validation.py index 000c87b..bad585c 100644 --- a/src/tigris/analysis/validation.py +++ b/src/tigris/analysis/validation.py @@ -17,6 +17,7 @@ from tigris.analysis.broadcast import DENSE, PERIODIC, stored_operands from tigris.capabilities import KERNEL_CAPABILITIES, effective_operators from tigris.emitters.binary.defs import OP_TYPE_MAP +from tigris.graph.ir import state_tensor_names from tigris.dtypes import check_dtype_signatures from tigris.graph.ir import ( AnalyzedGraph, @@ -96,7 +97,7 @@ def validate_execution_dtype(ag: AnalyzedGraph) -> ExecutionDTypeValidation: [(name, tensor.dtype, tensor.is_constant, tensor.quant is not None) for name, tensor in ag.tensors.items()], [(op.op_type, op.inputs, op.outputs) for op in ag.ops], - ag.model_inputs, ag.model_outputs) + ag.model_inputs, ag.model_outputs, state_tensor_names(ag)) if signature_issues: return ExecutionDTypeValidation(dtype=None, issues=signature_issues) diff --git a/src/tigris/capabilities.py b/src/tigris/capabilities.py index 0f47018..b2533c9 100644 --- a/src/tigris/capabilities.py +++ b/src/tigris/capabilities.py @@ -112,6 +112,7 @@ class KernelCapabilities: "Sum", "ReduceAll", "Split", + "Svdf", }) _S8_REFERENCE_OPERATORS = _FLOAT_REFERENCE_OPERATORS - frozenset({ @@ -248,6 +249,7 @@ class KernelCapabilities: "ReduceMax": ("one axis of a rank-3 tensor; independent height or row bands; rank-4 spatial max with keepdims uses GlobalMaxPool; int8 quantization must match",), "ReduceMin": ("one axis of a rank-3 tensor; independent height or row bands; int8 quantization must match",), "ReduceAll": ("one axis of a rank-3 bool tensor; independent height or row bands",), + "Svdf": ("state kept between runs; int8 with int16 state and time weights; untiled",), "ReduceSum": ("one axis of a rank-3 tensor; independent height or row bands; int8 rejects reference arithmetic overflow",), "Gather": ("constant or runtime int32 indices; independent bands with constant indices; identical int8 quantization",), "GatherND": ("constant or runtime int32 indices; independent bands with constant indices; identical int8 quantization",), diff --git a/src/tigris/dtypes.py b/src/tigris/dtypes.py index bb3a3e9..1e2611f 100644 --- a/src/tigris/dtypes.py +++ b/src/tigris/dtypes.py @@ -7,6 +7,8 @@ DATA = 0 DATA_OR_BOOL = 255 +# Data, or int16 on a tensor that holds state between runs. +STATE = 254 @dataclass(frozen=True) @@ -20,9 +22,11 @@ class DTypeSignature: class AuxiliaryDType: terminal_only: bool requires_source: bool = False + state_only: bool = False -AUXILIARY_DTYPES = {6: AuxiliaryDType(False, requires_source=True), 9: AuxiliaryDType(False)} +AUXILIARY_DTYPES = {6: AuxiliaryDType(False, requires_source=True), 9: AuxiliaryDType(False), + 5: AuxiliaryDType(False, state_only=True)} OP_DTYPE_SIGNATURES = {kind: DTypeSignature() for kind in OP_TYPE_MAP} OP_DTYPE_SIGNATURES.update({ **{kind: DTypeSignature(inputs=(DATA, 6, 6)) @@ -41,13 +45,18 @@ class AuxiliaryDType: "Where": DTypeSignature(inputs=(9, DATA, DATA)), "Cast": DTypeSignature(inputs=(9, 9, 9)), "ReduceAll": DTypeSignature(inputs=(9, 9, 9), outputs=(9,)), + # Constant operands sit between the input and the state; slots skip them. + # The output role admits int16 only on the next state, a state tensor. + "Svdf": DTypeSignature(inputs=(DATA, STATE, STATE), outputs=(STATE,)), **{kind: DTypeSignature(inputs=(DATA_OR_BOOL, DATA_OR_BOOL, DATA_OR_BOOL), outputs=(DATA_OR_BOOL,)) for kind in ("Transpose", "Reshape", "Flatten")}, }) -def check_dtype_signatures(tensors, operators, model_inputs, model_outputs): - """Return data tensors by dtype and errors for (name, dtype, constant, quantized) records.""" +def check_dtype_signatures(tensors, operators, model_inputs, model_outputs, state=()): + """Return data tensors by dtype and errors for (name, dtype, constant, quantized) records. + `state` names the tensors that carry state between runs.""" + state = set(state) by_name = {name: (dtype, constant, quantized) for name, dtype, constant, quantized in tensors} data = {} issues = [] @@ -58,6 +67,10 @@ def check_dtype_signatures(tensors, operators, model_inputs, model_outputs): if policy is None: data.setdefault(dtype, []).append(name) continue + if policy.state_only: + if name not in state: + issues.append(f"ONNX dtype {dtype} tensor {name} must hold state between runs") + continue if quantized: issues.append(f"auxiliary tensor {name} cannot carry quantization") if policy.requires_source and sum(name in outputs for _, _, outputs in operators) != (0 if name in model_inputs else 1): @@ -79,6 +92,11 @@ def check_dtype_signatures(tensors, operators, model_inputs, model_outputs): continue dtype = tensor[0] expected = slots[min(position, len(slots) - 1)] + if expected == STATE: + if dtype not in {1, 3} and not (dtype == 5 and name in state): + issues.append(f"{kind} {direction} {position} ({name}) requires data " + f"or int16 state, got ONNX dtype {dtype}") + continue if expected == DATA_OR_BOOL: if dtype not in {1, 3, 9} or dtype != by_name[inputs[0]][0]: issues.append(f"{kind} must preserve its data or bool dtype") diff --git a/src/tigris/emitters/binary/defs.py b/src/tigris/emitters/binary/defs.py index 1d1cdcf..e33f4ca 100644 --- a/src/tigris/emitters/binary/defs.py +++ b/src/tigris/emitters/binary/defs.py @@ -70,9 +70,15 @@ OP_ATTR_CUMSUM_OPTIONS = 11 # uint8[2], exclusive then reverse (each 0 or 1) OP_ATTR_COMPARISON_REQUANT = 13 # int32[5]: left shift, then two multiplier/shift pairs +OP_ATTR_CONSTANTS = 14 # uint16 weight index per constant operand, in operand order; + # NO_WEIGHT for an absent optional one +OP_ATTR_SVDF = 15 # int32: rank; for int8 also the state zero point and the + # input-to-state and state-to-output multiplier/shift pairs OP_ATTR_KINDS = ( OP_ATTR_COMPARISON_REQUANT, + OP_ATTR_CONSTANTS, + OP_ATTR_SVDF, OP_ATTR_TRANSPOSE_PERM, OP_ATTR_EPSILON, OP_ATTR_ALPHA, @@ -180,6 +186,7 @@ "Cast": 82, "Sum": 83, "ReduceAll": 84, + "Svdf": 85, } OP_TYPE_UNKNOWN = 255 diff --git a/src/tigris/emitters/binary/writer.py b/src/tigris/emitters/binary/writer.py index 510d0c4..f6df2e9 100644 --- a/src/tigris/emitters/binary/writer.py +++ b/src/tigris/emitters/binary/writer.py @@ -21,6 +21,7 @@ serialized_axis_map, serialized_shape, serialized_transpose_perm, + state_tensor_names, ) from .defs import ( @@ -40,6 +41,8 @@ OP_ATTR_CUMSUM_OPTIONS, OP_ATTR_MOVEMENT, OP_ATTR_COMPARISON_REQUANT, + OP_ATTR_CONSTANTS, + OP_ATTR_SVDF, OP_ATTR_ALPHA, OP_ATTR_BINARY_REQUANT, OP_ATTR_CONSTANT_OPERAND, @@ -486,8 +489,8 @@ def _build_weights( # Preserve ONNX flatten/reshape order across the internal layout change. if name in fc_layout_shapes: arr = _permute_fc_weight_for_nhwc(arr, fc_layout_shapes[name]) - # Preserve int8/int32 dtype for quantized weights - if arr.dtype in (np.int8, np.int32, np.bool_): + # Preserve the integer dtypes of quantized weights + if arr.dtype in (np.int8, np.int16, np.int32, np.bool_): raw = arr.tobytes() else: raw = arr.astype(np.float32).tobytes() @@ -576,7 +579,7 @@ def _build_weights_compressed( arr = _transpose_weight_nhwc(arr, op_type) if name in fc_layout_shapes: arr = _permute_fc_weight_for_nhwc(arr, fc_layout_shapes[name]) - if arr.dtype in (np.int8, np.int32, np.bool_): + if arr.dtype in (np.int8, np.int16, np.int32, np.bool_): raw = arr.tobytes() else: raw = arr.astype(np.float32).tobytes() @@ -699,13 +702,10 @@ def compressed_weight_reserve_bytes(ag: AnalyzedGraph) -> int: def _state_names(ag: AnalyzedGraph) -> set[str]: - """The model inputs and outputs that carry variables, not the interface.""" - names = set() - for port in ag.state_ports: - names.add(ag.model_inputs[port.input]) - if port.output is not None: - names.add(ag.model_outputs[port.output]) - return names + return state_tensor_names(ag) + + +_STATE_DTYPES = {1: np.float32, 3: np.int8, 5: np.int16} def _build_state(ag: AnalyzedGraph, tensor_idx: dict[str, int]) -> bytes: @@ -729,7 +729,7 @@ def _build_state(ag: AnalyzedGraph, tensor_idx: dict[str, int]) -> bytes: axis_map = serialized_axis_map(len(info.shape), info.layout) perm = sorted(range(len(axis_map)), key=lambda axis: axis_map[axis]) value = np.ascontiguousarray( - np.asarray(port.initial, np.float32).reshape(info.shape).transpose(perm)) + np.asarray(port.initial, _STATE_DTYPES[info.dtype]).reshape(info.shape).transpose(perm)) data = value.tobytes() entries.append(struct.pack(" bytes: + """The rank; for int8 also the state zero point and the two requantizations, + each scale formed in float32 as TFLite Micro forms it.""" + rank = int(op.attrs["rank"]) + if not ag.is_quantized: + return struct.pack(" bytes: """Build optional, typed per-operator attributes. @@ -919,6 +935,13 @@ def _build_op_attributes( constant = _constant_operand_payload(ag, op, quant_idx_map or {}) if constant is not None: records.append((op_index, OP_ATTR_CONSTANT_OPERAND, constant)) + if op.op_type in _MANY_CONSTANTS: + indices = [(weight_idx or {}).get(name, NO_WEIGHT) for name in op.inputs + if name not in tensor_idx] + records.append((op_index, OP_ATTR_CONSTANTS, struct.pack(f"<{len(indices)}H", *indices))) + if op.op_type == "Svdf": + records.append((op_index, OP_ATTR_SVDF, _svdf_payload(ag, op))) + continue if op.op_type in {"Resize", "ResizeLinear"}: scales = op.attrs.get("resize_scales") if scales is not None: @@ -1027,6 +1050,10 @@ def _build_op_attributes( ) +# Operators whose constants are listed by OP_ATTR_CONSTANTS, not weight and bias. +_MANY_CONSTANTS = {"Svdf"} + + def _resolve_weight_bias(op: OpNode, weight_idx: dict[str, int]) -> tuple[int, int]: """Map an op's constant inputs to weight/bias indices. @@ -1123,7 +1150,8 @@ def _build_ops( spatial = _pack_spatial_attrs(op, ag.weight_data) # Resolve weight/bias indices from op's constant inputs - w_idx, b_idx = _resolve_weight_bias(op, weight_idx) + w_idx, b_idx = ((NO_WEIGHT, NO_WEIGHT) if op.op_type in _MANY_CONSTANTS + else _resolve_weight_bias(op, weight_idx)) # Determine fused activation fused_act_str = op.attrs.get("fused_activation") @@ -1639,7 +1667,7 @@ def emit_binary_bytes( op_data = _build_ops(ag, tensor_idx, weight_idx, strings, index_pool) op_attributes_data = _build_op_attributes( - ag, tensor_idx, quant_idx_map + ag, tensor_idx, quant_idx_map, weight_idx ) stage_data = bytearray(_build_stages(ag, tensor_idx, index_pool)) tile_data, stage_to_tile = _build_tile_plans(ag) diff --git a/src/tigris/emitters/codegen.py b/src/tigris/emitters/codegen.py index 66bea71..cc6e262 100644 --- a/src/tigris/emitters/codegen.py +++ b/src/tigris/emitters/codegen.py @@ -19,7 +19,7 @@ resolve_kernel_backend, ) from tigris.dtypes import check_dtype_signatures -from tigris.emitters.binary.defs import FLAG_XIP +from tigris.emitters.binary.defs import FLAG_XIP, TENSOR_FLAG_STATE from tigris.emitters.binary.reader import read_binary_plan @@ -170,7 +170,9 @@ def _plan_dtype(plan: dict) -> DTypeMode: tensor.get("quant_param_idx", 65535) != 65535) for index, tensor in enumerate(plan.get("tensors", []))], operators, - plan.get("model_inputs", []), plan.get("model_outputs", [])) + plan.get("model_inputs", []), plan.get("model_outputs", []), + [index for index, tensor in enumerate(plan.get("tensors", [])) + if tensor.get("flags", 0) & TENSOR_FLAG_STATE]) if issues: raise ValueError("Plan dtype signature mismatch: " + "; ".join(issues)) tensor_dtypes = set(by_dtype) diff --git a/src/tigris/frontends/tflite.py b/src/tigris/frontends/tflite.py index c088994..fb32da6 100644 --- a/src/tigris/frontends/tflite.py +++ b/src/tigris/frontends/tflite.py @@ -25,7 +25,7 @@ # Field slots in the TFLite schema (tensorflow/lite/schema/schema.fbs). _MODEL_VERSION, _MODEL_OPCODES, _MODEL_SUBGRAPHS, _MODEL_DESCRIPTION, _MODEL_BUFFERS = 0, 1, 2, 3, 4 _SG_TENSORS, _SG_INPUTS, _SG_OUTPUTS, _SG_OPERATORS, _SG_NAME = 0, 1, 2, 3, 4 -_T_SHAPE, _T_TYPE, _T_BUFFER, _T_NAME, _T_QUANT = 0, 1, 2, 3, 4 +_T_SHAPE, _T_TYPE, _T_BUFFER, _T_NAME, _T_QUANT, _T_VARIABLE = 0, 1, 2, 3, 4, 5 _Q_SCALE, _Q_ZERO_POINT, _Q_DIMENSION = 2, 3, 6 _B_DATA, _B_OFFSET, _B_SIZE = 0, 1, 2 _OP_OPCODE, _OP_INPUTS, _OP_OUTPUTS, _OP_OPTIONS = 0, 1, 2, 4 @@ -137,6 +137,8 @@ def __init__(self, model: _Model, table: Table): self.scale = np.asarray(quant.scalars(_Q_SCALE, "f") if quant else [], np.float32) self.zero_point = np.asarray(quant.scalars(_Q_ZERO_POINT, "q") if quant else [], np.int64) self.quantized_dimension = quant.scalar(_Q_DIMENSION, "i") if quant else 0 + # A variable tensor keeps what an operator wrote into it for the next invocation. + self.variable = bool(table.scalar(_T_VARIABLE, "B")) self._model = model @property @@ -202,6 +204,18 @@ def unsupported(data: bytes) -> list[str]: if any(op.kind not in _INDEX_SLOTS or op.inputs[_INDEX_SLOTS[op.kind]] != index for op in uses): reasons.append(f"index {tensors[index].name!r} feeds an operand other than indices") + held = {op.inputs[_STATEFUL[op.kind]] for op in operators if op.kind in _STATEFUL} + for index, tensor in enumerate(tensors): + if not tensor.variable: + continue + uses = [op for op in operators if index in op.inputs] + if (index in inputs or index in outputs or len(uses) != 1 or index not in held + or any(index in op.outputs for op in operators)): + reasons.append(f"variable tensor {tensor.name!r} is not the state of one " + "SVDF") + for index in sorted(held): + if not tensors[index].variable: + reasons.append(f"state {tensors[index].name!r} is not a variable tensor") for index in sorted(slots - indices - index_inputs): if not _is_constant(tensors[index]): reasons.append(f"indices {tensors[index].name!r} are computed by an operator " @@ -287,6 +301,9 @@ def _activation(code: int) -> str: "LOGICAL_AND": ((0, 1), True), "LOGICAL_OR": ((0, 1), True), "LOGICAL_NOT": ((0,), True), "SELECT_V2": ((0,), False), "CAST": ((0,), False), "REDUCE_ALL": ((0,), True)} +# Operators that keep state in a variable tensor, at this operand. +_STATEFUL = {"SVDF": 4} +_STATE_TYPES = {"FLOAT32": TensorProto.FLOAT, "INT16": TensorProto.INT16} # Resource variables: a handle names a variable, which is read and assigned. _VARIABLES = ("CALL_ONCE", "VAR_HANDLE", "READ_VARIABLE", "ASSIGN_VARIABLE") # Index outputs: int32 positions, as model outputs or the indices of the operators below. @@ -297,7 +314,7 @@ def _activation(code: int) -> str: _REDUCTIONS = {"REDUCE_MAX": "ReduceMax", "REDUCE_MIN": "ReduceMin", "SUM": "ReduceSum", "REDUCE_ALL": "ReduceAll"} _SUPPORTED = (*_FUSED_SLOT, *_UNARY, *_SHAPE_ONLY, *_ELEMENTWISE, *_DATA_MOVEMENT, *_REDUCTIONS, - *_INDEX, *_BOOL_SLOTS, *_VARIABLES, "ADD_N", + *_INDEX, *_BOOL_SLOTS, *_VARIABLES, *_STATEFUL, "ADD_N", "RELU6", "SOFTMAX", "LOG_SOFTMAX", "LEAKY_RELU", "PRELU", "L2_NORMALIZATION", "CUMSUM", "MEAN", "TRANSPOSE", @@ -315,7 +332,7 @@ def _activation(code: int) -> str: "REDUCE_MIN": (1,), "SUM": (1,), "CUMSUM": (1,), "GATHER_ND": (1,), "MIRROR_PAD": (1,), "REVERSE_V2": (1,), "EMBEDDING_LOOKUP": (0,), "DYNAMIC_UPDATE_SLICE": (2,), "ARG_MAX": (1,), "ARG_MIN": (1,), - "REDUCE_ALL": (1,)} + "REDUCE_ALL": (1,), "SVDF": (1, 2, 3)} def _is_constant(tensor: _Tensor) -> bool: @@ -339,6 +356,8 @@ def _operator_reason(op: _Operator, tensors: list[_Tensor]) -> str: return f"input {position} must be a constant" if ins[position].type != "INT32": return f"{ins[position].type} run-time indices; the runtime takes int32" + if op.kind == "SVDF": + return _svdf_reason(op, ins, outs) data = [t for i, t in enumerate(ins) if t is not None and i not in _CONSTANT_OPERANDS.get(op.kind, ())] if op.kind in _WEIGHTED: source, position, bias_position = _WEIGHTED[op.kind] @@ -464,6 +483,30 @@ def _gather_run(op: _Operator, ins: list[_Tensor]): return axis, indices[0], len(indices), len(ins[1].shape) == 0 +def _svdf_reason(op: _Operator, ins: list, outs: list[_Tensor]) -> str: + x, feature, time, bias, state = ins + y = outs[0] + if bias is None: + # TFLite Micro's Prepare reads the bias whether or not it is there. + return "no bias, which TFLite Micro requires" + if len(x.shape) != 2 or op.option(0, "i") <= 0: + return "input of rank other than 2" + if x.type == "FLOAT32": + if any(t.type != "FLOAT32" for t in (feature, time, bias, state, y)): + return "float32 input with operands of another dtype" + fused = _activation(op.option(1, "b")) + if fused not in ("none", "relu", "relu6"): + return f"fused {fused}" + return "" + if x.type != "INT8" or y.type != "INT8" or feature.type != "INT8" or bias.type != "INT32": + return "activations must be all int8 or all float32" + if time.type != "INT16" or state.type != "INT16": + return "int8 state; the converter writes int16" + if any(len(t.scale) != 1 for t in (x, feature, time, state, y)): + return "operands must be quantized per tensor" + return "" + + def _data_movement_reason(op: _Operator, ins: list[_Tensor], outs: list[_Tensor]) -> str: slot = _INDEX_SLOTS.get(op.kind) data = [t for i, t in enumerate(ins) if not _is_constant(t) and i != slot] @@ -565,6 +608,8 @@ def __init__(self, tensors: list[_Tensor], outputs=(), consumed=()): # latest value. self.handles: dict[int, str] = {} self.variables: dict[str, dict] = {} + # Per variable tensor, its state input and the value written for the next run. + self.held_state: dict[int, dict] = {} def value(self, index: int, rank: int | None = None) -> str: """The ONNX value of a tensor; a constant is broadcast-aligned to `rank`.""" @@ -692,6 +737,9 @@ def convert(self, op: _Operator) -> None: select_last_index=0) self.values[outs[0]] = b.node("Cast", [y], out.name, to=TensorProto.INT32) return + if kind == "SVDF": + self.finish(self._svdf(op, tag), outs[0]) + return if kind in ("CONV_2D", "DEPTHWISE_CONV_2D"): y = self._conv(op, tag) elif kind == "TRANSPOSE_CONV": @@ -1095,6 +1143,34 @@ def _gather(self, kind: str, data: int, ids: int, axis, op: _Operator, tag: str) y = self.b.node(kind, [x, indices], tag + "_gathered", **attributes) return self.tflite_order_out(y, op.outputs[0], tag) + def _svdf(self, op: _Operator, tag: str) -> str: + """SVDF in the compiler's own form, its state passed in and out. An int8 + SVDF takes its constants as stored, with their scales stated; its int16 + state stays raw. TFLite Micro applies no activation to an int8 SVDF.""" + b = self.b + x, feature, time, bias, state = op.inputs + quantized = self.tensors[x].type == "INT8" + attributes = {"rank": op.option(0, "i"), + "activation": "none" if quantized else _activation(op.option(1, "b"))} + if quantized: + attributes.update(feature_scale=float(self.tensors[feature].scale[0]), + time_scale=float(self.tensors[time].scale[0]), + state_scale=float(self.tensors[state].scale[0]), + state_zero_point=int(self.tensors[state].zero_point[0])) + constants = [b.constant(self.tensors[i].array(), f"{tag}_{part}") + for i, part in ((feature, "feature"), (time, "time"), (bias, "bias"))] + y, kept = b.unique(tag), b.unique(tag + "_state") + b.nodes.append(helper.make_node("Svdf", [self.value(x), *constants, + self.held_state[state]["input"]], + [y, kept], domain="tigris", **attributes)) + held = self.tensors[state] + self.value_info += [ + helper.make_tensor_value_info(y, TensorProto.FLOAT, + list(self.tensors[op.outputs[0]].shape)), + helper.make_tensor_value_info(kept, _STATE_TYPES[held.type], list(held.shape))] + self.held_state[state]["output"] = kept + return y + def _strided_slice(self, op: _Operator, tag: str) -> str: """A STRIDED_SLICE with any strides, as an ONNX Slice in TFLite's axis order; a dropped axis is a slice of one, reshaped away.""" @@ -1318,9 +1394,30 @@ def to_onnx(data: bytes, name: str) -> onnx.ModelProto: state.append({"input": name + "_in", "output": None, "initial": b.constant(initial.astype(np.float32).transpose( _to_first(len(shape))), name + "_initial")}) + # A variable tensor enters as a state input, zero before the first run as + # TFLite Micro resets it; the operator's write leaves as the state output. + for op in operators: + if op.kind in _STATEFUL: + index = op.inputs[_STATEFUL[op.kind]] + tensor = tensors[index] + name = b.unique(f"state{len(state)}_{tensor.name}_in") + onnx_inputs.append(helper.make_tensor_value_info( + name, _STATE_TYPES[tensor.type], list(tensor.shape))) + initial = b.constant(np.zeros(tensor.shape, _NUMPY[tensor.type]), + name.removesuffix("_in") + "_initial") + converter.held_state[index] = {"input": name, "output": None} + converter.values[index] = name + state.append({"input": name, "output": None, "initial": initial}) for op in operators: converter.convert(op) + for entry, held in zip(state[len(state) - len(converter.held_state):], + converter.held_state.values()): + entry["output"] = held["output"] state_outputs = [] + for index, held in converter.held_state.items(): + tensor = tensors[index] + state_outputs.append(helper.make_tensor_value_info( + held["output"], _STATE_TYPES[tensor.type], list(tensor.shape))) for entry, variable in zip(state, shapes): current = converter.variables[variable] if current["current"] != entry["input"]: diff --git a/src/tigris/graph/ir.py b/src/tigris/graph/ir.py index 1912b0f..d2f1c93 100644 --- a/src/tigris/graph/ir.py +++ b/src/tigris/graph/ir.py @@ -307,3 +307,13 @@ def mem_budget(self) -> int: def fast_memory_reserve_bytes(self) -> int: """Bytes held outside the activation arena. Compat accessor over ``budget``.""" return self.budget.fast_reserve + + +def state_tensor_names(ag: "AnalyzedGraph") -> set[str]: + """The model inputs and outputs that carry variables, not the interface.""" + names = set() + for port in ag.state_ports: + names.add(ag.model_inputs[port.input]) + if port.output is not None: + names.add(ag.model_outputs[port.output]) + return names diff --git a/src/tigris/loaders/onnx/normalize.py b/src/tigris/loaders/onnx/normalize.py index d002bae..7622797 100644 --- a/src/tigris/loaders/onnx/normalize.py +++ b/src/tigris/loaders/onnx/normalize.py @@ -56,6 +56,7 @@ def normalize(ag: AnalyzedGraph) -> AnalyzedGraph: """Apply all normalization passes in sequence.""" declared_outputs = list(ag.model_outputs) ag = _adopt_tflite_cumsum(ag) + ag = _adopt_svdf(ag) ag = _normalize_arg_outputs(ag) ag = _drop_inference_identities(ag) ag = _lower_legacy_softmax(ag) @@ -2912,6 +2913,40 @@ def _adopt_tflite_cumsum(ag: AnalyzedGraph) -> AnalyzedGraph: return ag +def _adopt_svdf(ag: AnalyzedGraph) -> AnalyzedGraph: + """The compiler's own SVDF: input [batch, features], constant feature + [filters, features], time [filters, memory] and bias [units] weights, and a + state [batch, filters * memory] passed in and written back.""" + for op in ag.ops: + if op.op_type != "tigris::Svdf": + continue + op.op_type = "Svdf" + if len(op.inputs) != 5 or len(op.outputs) != 2: + raise ValueError("Svdf requires five inputs and two outputs") + x, y = ag.tensors[op.inputs[0]], ag.tensors[op.outputs[0]] + feature, time, bias = (ag.weight_data.get(name) for name in op.inputs[1:4]) + state, kept = ag.tensors[op.inputs[4]], ag.tensors[op.outputs[1]] + rank = int(op.attrs.get("rank", 0)) + if feature is None or time is None or bias is None: + raise ValueError("Svdf requires constant feature, time and bias weights") + if len(x.shape) != 2 or feature.ndim != 2 or time.ndim != 2 or bias.ndim != 1: + raise ValueError("Svdf requires a rank-2 input and rank-2 weights") + batch, features = x.shape + filters, memory = time.shape + if (rank <= 0 or filters % rank or feature.shape != (filters, features) + or bias.shape != (filters // rank,) or tuple(y.shape) != (batch, filters // rank) + or tuple(state.shape) != (batch, filters * memory) or kept.shape != state.shape): + raise ValueError("Svdf shapes do not agree with its rank") + activation = op.attrs.pop("activation", "none") + if isinstance(activation, bytes): + activation = activation.decode() + if activation not in ("none", "relu", "relu6"): + raise ValueError(f"Svdf activation {activation!r} is not supported") + if activation != "none": + op.attrs["fused_activation"] = {"relu": "Relu", "relu6": "Relu6"}[activation] + return ag + + def _normalize_arg_outputs(ag: AnalyzedGraph) -> AnalyzedGraph: # An index cast to int32 is already stored as the runtime writes it. for cast in list(ag.ops): diff --git a/src/tigris/schema/operator-capabilities-v1.json b/src/tigris/schema/operator-capabilities-v1.json index 4e41ea0..4c8d8d3 100644 --- a/src/tigris/schema/operator-capabilities-v1.json +++ b/src/tigris/schema/operator-capabilities-v1.json @@ -1111,6 +1111,19 @@ "reference": "native", "s8_ref": "native" } + }, + { + "constraints": [ + "state kept between runs; int8 with int16 state and time weights; untiled" + ], + "opcode": 85, + "operator": "Svdf", + "routes": { + "cmsis-nn": "fallback:s8_ref", + "esp-nn": "fallback:s8_ref", + "reference": "native", + "s8_ref": "native" + } } ], "schema_version": 10, diff --git a/src/tigris/schema/tigris-plan-v10.json b/src/tigris/schema/tigris-plan-v10.json index 5951392..839f4b4 100644 --- a/src/tigris/schema/tigris-plan-v10.json +++ b/src/tigris/schema/tigris-plan-v10.json @@ -17,12 +17,14 @@ "clip_bounds": 4, "comparison_requant": 13, "constant_operand": 9, + "constants": 14, "cumsum_options": 11, "epsilon": 2, "movement": 12, "pads": 5, "pool_rounding": 8, "resize_scales": 10, + "svdf": 15, "transpose_perm": 1 }, "op_types": { @@ -106,6 +108,7 @@ "StridedSlice": 68, "Sub": 18, "Sum": 83, + "Svdf": 85, "Tanh": 20, "Transpose": 29, "Unsqueeze": 28, @@ -158,6 +161,7 @@ "capability_growth_requires_schema_bump": false, "scope": "wire_format", "stage_assignment_v5": "stage table authoritative; operator byte is canonical low-byte hint", + "state_int16": "int16 (ONNX dtype 5) is stored only on state tensors", "state_v10": "only a plan with a state section is written at schema 10" } } diff --git a/src/tigris/utils.py b/src/tigris/utils.py index 184b62d..f6e346e 100644 --- a/src/tigris/utils.py +++ b/src/tigris/utils.py @@ -59,7 +59,7 @@ def describe_interface(ag) -> list[tuple[str, str]]: ) rows.append((label, text)) if ports: - size = sum(port.initial.size * 4 for port in ports) + size = sum(ag.tensors[ag.model_inputs[port.input]].size_bytes for port in ports) rows.append(("State", f"{len(ports)} variable{'s' if len(ports) != 1 else ''}, " f"{size} bytes kept across runs")) return rows diff --git a/tests/fixtures/tflite/ops/float_svdf.npz b/tests/fixtures/tflite/ops/float_svdf.npz new file mode 100644 index 0000000000000000000000000000000000000000..0bd58d47948fb2715ee92e02bb93d7f62fbf2d4f GIT binary patch literal 609 zcmWIWW@gc4U|`??Vnv3`KdB7=p@5q~gdsDpptL03KrgSLl954xfq@aI3J5MjD2CZ@ z#9v7ZoIDY5EMV=tNl|lRmgFs6kT*GP(YnP8=1q$a37RrPe180tNnAqZ-e)xZSi8j2 zRm7Gs{ghI(;%c+v5}RekwUz4$^3ncKhnKh3{)f628B`==)_{=IbX+V9Ef zM%FT)Dl_7<6$%*`0#HI@;+K4BSZMG7LnFVm1SL2QK!c%&I>Dg^4Gz7?z0PH?{%-u{ z6fgbFn{E27jc4wpzE&*1?dM#5^xgtz{x@CHx8LUGt-bx<>3H#fqvN*UZi|+_nXGPl z({26(%MERw7a4&*VPw)}MvYl;v;m`#8|q#V)c}rCq1cP4vvTS*TjBlc(8q&LW7P!D);-}j zg`FJktJW<#r8jk*Pjk~#w#UEg@%T;WSbHlEKdf%B`wysS(Y^Pw)E=}S>#;N4PD^Zr zo9`cWG9PBt1g2NN`R9sJ&9@oG2CDVaaR6SbkWZki#;@JjOez=WK z?&wWzt9q%I^L>2dcCY%J$)pkw=GC2Vuhi9hisv?GbBDX%*-&Kvx0WPIRSW>C$uP+H`LhZX;ERmQvfz*{D3HBQ9yj; z0qZ@8F=vs@=msrTILA|OYhJW99Y!3u*bE?ceTZ~AdY%Q>f*f&33>()RHYR`Hc)b%-uH4Zuondy<45jX(u X@cA`seSwdd`#+dD9#^AE<5T5-(&?v6 literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/float_svdf_rank2_relu.npz b/tests/fixtures/tflite/ops/float_svdf_rank2_relu.npz new file mode 100644 index 0000000000000000000000000000000000000000..0bfb63be42f6b823276a74c0f7b4dc2f0e9bc0c5 GIT binary patch literal 668 zcmWIWW@gc4U|`??VnqhwldWq1p@5q~gdsDpptL03KrgSLl954x!GRH|3J64@^z1j{ zucQS|o(MP=uy)?0s5vo9@|G^ho1C_2-QorFrp1Q@O_?D+KYq$2E}?SoGn#&^UE=8~ zVoR8QN~u|Kxmj_g&9dU!%Jqa%Z|Ya2-=$B)-@M*;D%n2u`PEvn*4uBIovy$4ucZou?H>+j;>+4Kw6`(O+#17vV=FITiEP=X{PZxJ2MOPj^LsxT zM!&y$jNPy7giOWvH%Ina%LTE=)f;Z?Kf8PDbme1594nb;H`i_GJ*=t7+@5?t??%1Y z-MsocHi_{syARav@?ZAOdik7v|BngzwcpfN;k#qKy>$M8IZnShzlXp59(}&%L8jjE zf4SdZeLJ6WW@~z>&F=VXC*B;m@7=?mr@JSOpCJGxL@#AptHVN+2NcKnw#7j6p!H$RQEn&B_LnWdg!tAT7)U;sF5Qc>6^F literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/float_svdf_rank2_relu.tflite b/tests/fixtures/tflite/ops/float_svdf_rank2_relu.tflite new file mode 100644 index 0000000000000000000000000000000000000000..16d8ee9290c2a38d6edcc35efc7ed0ce10f54a14 GIT binary patch literal 848 zcmYLHU1$?Q5T59n2778AVkss?u@9n%Rijcgy9b3@Erke;1zUnr6{<9nhrH-ZYm10f zu>L_HT8bk2P!us$i@n{_K4?qDJQaVDB4|OQ;17lxJ!keV*SmZ-v$Hdk+3(v*0Eq2w z3Ii;M0Qf-%A4W7>0j>Z#JOKJ)0Q^B2B8KZdh&$+SMg$N#LPJddv#cEAlx118ruc`? z8;!GPK0TR_#uo}*IpIz?n0!sZOp7|^Bb3$C3=4ih z3iDpg1@^6amNj1-=fThh2%2jfb#t~XQzR7*&b5+d0d%d!K z98&n7ogo0loe5fh#9;2^0lG84v;@VW6QB-^qK-q=pbpg&Q@nS6O-@#h&3jh4!oojR zyg+v`GU+m-h9D?-7636f)OHZnzz8Cd!V_I1$V(vSf}G9+QVRtQj4OaFWDf>-v$BCC NnSihuNIQUy1^~H%iVy$* literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/svdf.tflite b/tests/fixtures/tflite/ops/svdf.tflite new file mode 100644 index 0000000000000000000000000000000000000000..42fb419521b3fa38459e961d6b287beefdd01523 GIT binary patch literal 1016 zcmaJ=OK1~O6uoJtNooeu(3p;C7D=&KOQ2~JTf`QF1viQWDMD>lQfw(*6hG^_aA&}U zKSUP_ZbXDwx)6y_T~vxLy0KId3nh!R{z8!wYC7XNFB6-MUgpeu_q}(|J@?L=u!szw zIF-Nv+baQS7F|4;Yw!dfW)6{@Qv{E=ff7&vZUIRk0O)`Q#4XE80-sFFqR)8a>4-DE zugkeKlHQ)anEAbB@1?OxR`G*epHj9xTo+x37+F4SoL_deFJv z?$5;zSTk8)C0%UHez-lGmBr}K3lEnU%%RGm(zAPKSNzwXtb`9vM$aHO^QM6Hg^kkw z*v#|0?^pLkgFjv_igond=$ct5?t~J*3U^xkzK{1egQu7JJ|BDgZoFgei@8?uxy_DP zzr*m|JYG>;nyL#eP_^iPJEWFUO1xy1x$9}MvLA=>buiZ;St&59eSj}DNv?eah z8vw>O+Q6WEuCwo~skYCN#<)D34B!Fs@W%k2okWkRBM$H!tSN?92J8H~)J>ip)JWs} z#+xgH%y;?lt{!qVL#-jLsumA!0Apbe@+SdmhPjz5Z}W4Hob{%B$JHY$2Lfv~cj8b8 z8W=+k=7cp+ua5P4y=n(9az%Q}J&FsrmpS9jT zHCv~P`!LnV$N$lFDd!zfTs^y7;M1& literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/svdf_batch.npz b/tests/fixtures/tflite/ops/svdf_batch.npz new file mode 100644 index 0000000000000000000000000000000000000000..789434695f2a7a3434fc666ff72a7987782012f9 GIT binary patch literal 483 zcmWIWW@gc4U|`??Vnv3;g2+4np@5q~gdsDpptL03KrgSLl954xVFOSR2-H9*hS_h# zUr7s`JP~j#VC}p~QFCIJCtZX1jCLk;Z(otZe0ot#hy#N3J literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/svdf_batch.tflite b/tests/fixtures/tflite/ops/svdf_batch.tflite new file mode 100644 index 0000000000000000000000000000000000000000..081f12164d0efa9ded6de3327087a32d2270b942 GIT binary patch literal 912 zcmaJ!1Lrd1)P)iLF1cCPYzUzDGQiu1wd(XY+e&;*q-mr)yCT99Ez&zp?ujt~!T!Sa@ zM6$qSlL$U>0z1GaFbnhoen1B_VBau|EHDnx=eqAGCWi+vU+L~bWzJj=V!a&M`zVZ+ zZj9}e{q?!(d@LSkPH)%BTV{+8G1_$Ovnwz)|E^s{-ZGAdr#q2~+SgP*8)B$Y|5dK3 zoLpN1R)7@1IO=4LZTPcFw^<{#z?wORYg2SUlV;3$j&PX3pnN+<_JcKA`y6RJmxq%C zTtFWFBEZdSit7)(1s6TDiSY8m#f6*U zS(uId0{}I{oXnLs`MF2VMpM2mJ*;xzpp}Dn*A5LlLk{YP*{N5@dZS*ogCDt8FLP1F zg;zzNeOfSisEa;K=}MS9f9dMH$PFp3zIxr%an_)V`!LnVlYeyGpS$CVD_ZBO-GK*v z^h+z=1n-^iz5^HqRByl?$6Qx03wMeb_J;Lax=jsOshuemi?v*69QO$8J@KB%<;KXa Kk}IU<68Zr=oOh)F literal 0 HcmV?d00001 diff --git a/tests/test_frontend_tflite.py b/tests/test_frontend_tflite.py index 287491b..3b40576 100644 --- a/tests/test_frontend_tflite.py +++ b/tests/test_frontend_tflite.py @@ -182,16 +182,16 @@ def _with_operator_replaced(data: bytes, kind: str, code: int) -> bytes: def test_an_unsupported_operator_is_refused_by_name(tmp_path): - svdf = tflite._BUILTIN_OPERATORS.index("SVDF") - path = tmp_path / "kws_svdf.tflite" - path.write_bytes(_with_operator_replaced(KWS.read_bytes(), "FULLY_CONNECTED", svdf)) + sparse = tflite._BUILTIN_OPERATORS.index("EMBEDDING_LOOKUP_SPARSE") + path = tmp_path / "kws_sparse.tflite" + path.write_bytes(_with_operator_replaced(KWS.read_bytes(), "FULLY_CONNECTED", sparse)) report = inspect_file(path) assert report["format"] == "tflite" - assert any("SVDF: not supported" in reason for reason in report["unsupported"]) + assert any("EMBEDDING_LOOKUP_SPARSE: not supported" in reason for reason in report["unsupported"]) result = CliRunner().invoke(cli, ["compile", str(path), "-m", "16K", "-o", str(tmp_path / "x.tgrs")]) assert result.exit_code != 0 - assert "SVDF: not supported" in result.output + assert "EMBEDDING_LOOKUP_SPARSE: not supported" in result.output assert not (tmp_path / "x.tgrs").exists() @@ -280,6 +280,29 @@ def widen(tensors, operators): assert any("INT64 run-time indices; the runtime takes int32" in r for r in reasons) +def test_an_svdf_without_bias_is_refused(monkeypatch): + """TFLite Micro's SVDF Prepare reads the bias whether or not it is there.""" + def drop(tensors, operators): + operators[0].inputs = [*operators[0].inputs[:3], -1, operators[0].inputs[4]] + reasons = _read_edited(monkeypatch, "float_svdf", drop) + assert any("no bias, which TFLite Micro requires" in r for r in reasons) + + +def test_an_int8_svdf_keeps_int16_state(monkeypatch): + def narrow(tensors, operators): + for position in (2, 4): + tensors[operators[0].inputs[position]].type = "INT8" + reasons = _read_edited(monkeypatch, "svdf", narrow) + assert any("int8 state; the converter writes int16" in r for r in reasons) + + +def test_a_variable_tensor_is_only_svdf_state(monkeypatch): + def mark(tensors, operators): + tensors[operators[0].inputs[0]].variable = True + reasons = _read_edited(monkeypatch, "float_svdf", mark) + assert any("is not the state of one SVDF" in r for r in reasons) + + def test_an_int8_cumsum_keeps_tflite_semantics_in_the_compilers_own_form(): model = tflite.to_onnx((FIXTURES / "ops" / "cumsum_offset.tflite").read_bytes(), "cumsum") scans = [node for node in model.graph.node if node.op_type == "CumSum"]