From 707e078cc8055e3ccf30c88e44e143c265983300 Mon Sep 17 00:00:00 2001 From: asteinh Date: Mon, 5 Oct 2026 06:17:15 +0200 Subject: [PATCH] feat: compile TFLite LSTM in float32 with its states --- scripts/crossrepo_contract.py | 56 ++++++++++++ scripts/gen_tflite_fixtures.py | 41 ++++++++- scripts/generate_schema_package.py | 1 + src/tigris/capabilities.py | 3 + src/tigris/dtypes.py | 1 + src/tigris/emitters/binary/defs.py | 3 + src/tigris/emitters/binary/writer.py | 7 +- src/tigris/frontends/tflite.py | 81 ++++++++++++++---- src/tigris/loaders/onnx/normalize.py | 48 +++++++++++ .../schema/operator-capabilities-v1.json | 13 +++ src/tigris/schema/tigris-plan-v10.json | 2 + .../tflite/ops/float_logistic_tails.npz | Bin 0 -> 2316 bytes .../tflite/ops/float_logistic_tails.tflite | Bin 0 -> 656 bytes tests/fixtures/tflite/ops/float_lstm.npz | Bin 0 -> 870 bytes tests/fixtures/tflite/ops/float_lstm.tflite | Bin 0 -> 2020 bytes .../tflite/ops/float_lstm_time_major_clip.npz | Bin 0 -> 1119 bytes .../ops/float_lstm_time_major_clip.tflite | Bin 0 -> 1744 bytes tests/test_frontend_tflite.py | 19 ++++ 18 files changed, 256 insertions(+), 19 deletions(-) create mode 100644 tests/fixtures/tflite/ops/float_logistic_tails.npz create mode 100644 tests/fixtures/tflite/ops/float_logistic_tails.tflite create mode 100644 tests/fixtures/tflite/ops/float_lstm.npz create mode 100644 tests/fixtures/tflite/ops/float_lstm.tflite create mode 100644 tests/fixtures/tflite/ops/float_lstm_time_major_clip.npz create mode 100644 tests/fixtures/tflite/ops/float_lstm_time_major_clip.tflite diff --git a/scripts/crossrepo_contract.py b/scripts/crossrepo_contract.py index be1370d..1badc29 100644 --- a/scripts/crossrepo_contract.py +++ b/scripts/crossrepo_contract.py @@ -895,6 +895,61 @@ def _svdf_case() -> ContractCase: inputs={"x": x}, expected_operators=("Svdf",)) +def _lstm_case() -> ContractCase: + """A float LSTM over three steps from its zero initial state, against ONNX's + own LSTM operator, which orders the gates i, o, f, c.""" + rng = np.random.default_rng(86) + batch, steps, features, units = 2, 3, 3, 4 + w_in = rng.normal(0.0, 0.5, (4, units, features)).astype(np.float32) + w_rec = rng.normal(0.0, 0.5, (4, units, units)).astype(np.float32) + bias = rng.normal(0.0, 0.3, (4, units)).astype(np.float32) + x = rng.uniform(-2.0, 2.0, (batch, steps, features)).astype(np.float32) + shape_s, shape_y = [batch, units], [batch, steps, units] + weights = [numpy_helper.from_array(w, f"w{k}") + for k, w in enumerate([*w_in, *w_rec, *bias])] + weights += [numpy_helper.from_array(np.zeros(shape_s, np.float32), name) + for name in ("hidden_initial", "cell_initial")] + states = [helper.make_tensor_value_info(name, TensorProto.FLOAT, shape_s) + for name in ("hidden_out", "cell_out")] + compile_graph = helper.make_graph( + [helper.make_node("Lstm", ["x", *(f"w{k}" for k in range(12)), "hidden_in", "cell_in"], + ["y", "hidden_out", "cell_out"], domain="tigris", time_major=0, + cell_clip=0.0)], + "lstm", + [helper.make_tensor_value_info("x", TensorProto.FLOAT, list(x.shape)), + *(helper.make_tensor_value_info(name, TensorProto.FLOAT, shape_s) + for name in ("hidden_in", "cell_in"))], + [helper.make_tensor_value_info("y", TensorProto.FLOAT, shape_y), *states], + weights, + value_info=[helper.make_tensor_value_info("y", TensorProto.FLOAT, shape_y), *states]) + 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": "hidden_in", "output": "hidden_out", "initial": "hidden_initial"}, + {"input": "cell_in", "output": "cell_out", "initial": "cell_initial"}])}) + onnx_order = [0, 3, 1, 2] + reference = _model( + "lstm_reference", + # ONNX Runtime's LSTM runs time-major only. + [helper.make_node("Transpose", ["x"], ["steps_first"], perm=[1, 0, 2]), + helper.make_node("LSTM", ["steps_first", "W", "R", "B"], ["sequence"], hidden_size=units), + helper.make_node("Squeeze", ["sequence", "direction_axis"], ["squeezed"]), + helper.make_node("Transpose", ["squeezed"], ["y"], perm=[1, 0, 2])], + [helper.make_tensor_value_info("x", TensorProto.FLOAT, list(x.shape))], + [helper.make_tensor_value_info("y", TensorProto.FLOAT, shape_y)], + [numpy_helper.from_array(w_in[onnx_order].reshape(1, 4 * units, features), "W"), + numpy_helper.from_array(w_rec[onnx_order].reshape(1, 4 * units, units), "R"), + numpy_helper.from_array(np.concatenate( + [bias[onnx_order].reshape(-1), np.zeros(4 * units, np.float32)])[None], "B"), + numpy_helper.from_array(np.asarray([1], np.int64), "direction_axis")], + opset=14) + return ContractCase(name="lstm", compile_model=compile_model, reference_model=reference, + inputs={"x": x}, + # Rank-3 boundaries are stored channels-last; Lstm reads model order. + expected_operators=("Transpose", "Lstm", "Transpose")) + + 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) @@ -6970,6 +7025,7 @@ def _run_gate(runtime: Path, work_dir: Path) -> None: _qdq_convtranspose_2d_tiled_case(), _convtranspose_2d_partial_edge_case(), _svdf_case(), + _lstm_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 04e14db..848034a 100644 --- a/scripts/gen_tflite_fixtures.py +++ b/scripts/gen_tflite_fixtures.py @@ -172,6 +172,8 @@ def _binary(fn, a, b): "float_div_broadcast": [(-3.0, 3.0), _DIVISOR], # Non-negative inputs put the int8 input zero point near -128. "cumsum_offset": [(0.0, 3.0)], "cumsum_offset_exclusive_reverse": [(0.0, 3.0)], + # Past both of the float logistic's cutoffs, -9 and about 16.6. + "float_logistic_tails": [(-24.0, 24.0)], } @@ -555,6 +557,35 @@ def _svdf(batch, features, units, rank, memory, activation, quantized=False): options, tensors, op_inputs, [5], [0], [5]) +def _lstm(time_major, batch, steps, features, units, cell_clip=0.0): + """UNIDIRECTIONAL_SEQUENCE_LSTM without peepholes, projection or layer + normalization, which TFLite Micro does not run; its hidden and cell states + are variable tensors. Seeded weights, gates in TFLite's order i, f, c, o.""" + rng = np.random.default_rng(steps * 100 + features * 10 + units) + T = schema.TensorType + shape = (steps, batch, features) if time_major else (batch, steps, features) + out = (steps, batch, units) if time_major else (batch, steps, units) + tensors = [(shape, T.FLOAT32, None, None, 0, False)] + for columns in (features, units): + for _ in range(4): + tensors.append(((units, columns), T.FLOAT32, + rng.normal(0.0, 0.5, (units, columns)).astype(np.float32), None, 0, False)) + for _ in range(4): + tensors.append(((units,), T.FLOAT32, rng.normal(0.0, 0.3, (units,)).astype(np.float32), + None, 0, False)) + tensors += [((batch, units), T.FLOAT32, None, None, 0, True), + ((batch, units), T.FLOAT32, None, None, 0, True), + (out, T.FLOAT32, None, None, 0, False)] + options = schema.UnidirectionalSequenceLSTMOptionsT() + options.fusedActivationFunction = schema.ActivationFunctionType.TANH + options.cellClip = cell_clip + options.timeMajor = time_major + op_inputs = [0, *range(1, 9), -1, -1, -1, *range(9, 13), -1, -1, 13, 14, -1, -1, -1, -1] + return _one_operator(schema.BuiltinOperator.UNIDIRECTIONAL_SEQUENCE_LSTM, + schema.BuiltinOptions.UnidirectionalSequenceLSTMOptions, options, + tensors, op_inputs, [15], [0], [15]) + + # Models no converter writes, built operator by operator. _RELU = schema.ActivationFunctionType.RELU _NONE = schema.ActivationFunctionType.NONE @@ -563,10 +594,13 @@ def _svdf(batch, features, units, rank, memory, activation, quantized=False): "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), + "float_lstm": lambda: _lstm(False, 1, 3, 4, 5), + "float_lstm_time_major_clip": lambda: _lstm(True, 2, 3, 3, 4, cell_clip=0.8), } -# 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"} +# TFLite's float SVDF and LSTM compute 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", "float_lstm", + "float_lstm_time_major_clip"} # The tier-1 cases again, converted without quantization. _FLOAT_TIER1 = ( "max_pool_valid", "max_pool_same", "avg_pool_valid", "avg_pool_same", "concat_channels", @@ -582,6 +616,7 @@ def _svdf(batch, features, units, rank, memory, activation, quantized=False): *ACTIVATIONS, *INDEXING, *BOOLEAN, ) FLOAT_MODELS["float_l2_pool"] = CASES["float_l2_pool"] +FLOAT_MODELS["float_logistic_tails"] = _unary(tf.sigmoid, (1, 64)) FLOAT_MODELS["float_variable_window"] = _unary(lambda x: _WINDOW(x), (1, 2)) for _name in _FLOAT_TIER1: FLOAT_MODELS[f"float_{_name}"] = CASES[_name] diff --git a/scripts/generate_schema_package.py b/scripts/generate_schema_package.py index 92e0543..18e0c87 100644 --- a/scripts/generate_schema_package.py +++ b/scripts/generate_schema_package.py @@ -35,6 +35,7 @@ def schema_package() -> dict[str, object]: "comparison_requant": defs.OP_ATTR_COMPARISON_REQUANT, "constants": defs.OP_ATTR_CONSTANTS, "svdf": defs.OP_ATTR_SVDF, + "lstm": defs.OP_ATTR_LSTM, "epsilon": defs.OP_ATTR_EPSILON, "pads": defs.OP_ATTR_PADS, "pool_rounding": defs.OP_ATTR_POOL_ROUNDING, diff --git a/src/tigris/capabilities.py b/src/tigris/capabilities.py index b2533c9..18a1127 100644 --- a/src/tigris/capabilities.py +++ b/src/tigris/capabilities.py @@ -113,9 +113,11 @@ class KernelCapabilities: "ReduceAll", "Split", "Svdf", + "Lstm", }) _S8_REFERENCE_OPERATORS = _FLOAT_REFERENCE_OPERATORS - frozenset({ + "Lstm", "L2Pool", "Neg", "Exp", @@ -250,6 +252,7 @@ class KernelCapabilities: "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",), + "Lstm": ("hidden and cell state kept between runs; every gate, no peepholes, projection or layer normalization; tanh cell activation; float32; 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 1e2611f..679d677 100644 --- a/src/tigris/dtypes.py +++ b/src/tigris/dtypes.py @@ -48,6 +48,7 @@ class AuxiliaryDType: # 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,)), + "Lstm": 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")}, }) diff --git a/src/tigris/emitters/binary/defs.py b/src/tigris/emitters/binary/defs.py index e33f4ca..a31a1e7 100644 --- a/src/tigris/emitters/binary/defs.py +++ b/src/tigris/emitters/binary/defs.py @@ -74,11 +74,13 @@ # 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_LSTM = 16 # int32 time_major (0 or 1), then float32 cell clip (0: none) OP_ATTR_KINDS = ( OP_ATTR_COMPARISON_REQUANT, OP_ATTR_CONSTANTS, OP_ATTR_SVDF, + OP_ATTR_LSTM, OP_ATTR_TRANSPOSE_PERM, OP_ATTR_EPSILON, OP_ATTR_ALPHA, @@ -187,6 +189,7 @@ "Sum": 83, "ReduceAll": 84, "Svdf": 85, + "Lstm": 86, } OP_TYPE_UNKNOWN = 255 diff --git a/src/tigris/emitters/binary/writer.py b/src/tigris/emitters/binary/writer.py index f6df2e9..cd69e7d 100644 --- a/src/tigris/emitters/binary/writer.py +++ b/src/tigris/emitters/binary/writer.py @@ -43,6 +43,7 @@ OP_ATTR_COMPARISON_REQUANT, OP_ATTR_CONSTANTS, OP_ATTR_SVDF, + OP_ATTR_LSTM, OP_ATTR_ALPHA, OP_ATTR_BINARY_REQUANT, OP_ATTR_CONSTANT_OPERAND, @@ -942,6 +943,10 @@ def _build_op_attributes( if op.op_type == "Svdf": records.append((op_index, OP_ATTR_SVDF, _svdf_payload(ag, op))) continue + if op.op_type == "Lstm": + records.append((op_index, OP_ATTR_LSTM, struct.pack( + " tuple[int, int]: diff --git a/src/tigris/frontends/tflite.py b/src/tigris/frontends/tflite.py index fb32da6..5deee77 100644 --- a/src/tigris/frontends/tflite.py +++ b/src/tigris/frontends/tflite.py @@ -204,7 +204,8 @@ 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} + held = {op.inputs[slot] for op in operators if op.kind in _STATEFUL + for slot in _STATEFUL[op.kind]} for index, tensor in enumerate(tensors): if not tensor.variable: continue @@ -212,7 +213,7 @@ def unsupported(data: bytes) -> list[str]: 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") + "SVDF or LSTM") for index in sorted(held): if not tensors[index].variable: reasons.append(f"state {tensors[index].name!r} is not a variable tensor") @@ -302,7 +303,7 @@ def _activation(code: int) -> str: "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} +_STATEFUL = {"SVDF": (4,), "UNIDIRECTIONAL_SEQUENCE_LSTM": (18, 19)} _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") @@ -358,6 +359,8 @@ def _operator_reason(op: _Operator, tensors: list[_Tensor]) -> str: return f"{ins[position].type} run-time indices; the runtime takes int32" if op.kind == "SVDF": return _svdf_reason(op, ins, outs) + if op.kind == "UNIDIRECTIONAL_SEQUENCE_LSTM": + return _lstm_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] @@ -507,6 +510,30 @@ def _svdf_reason(op: _Operator, ins: list, outs: list[_Tensor]) -> str: return "" +# Operand positions of the LSTM's gate weights and biases, gates in order i, f, c, o. +_LSTM_CONSTANTS = (*range(1, 9), *range(12, 16)) + + +def _lstm_reason(op: _Operator, ins: list, outs: list[_Tensor]) -> str: + """What TFLite Micro runs: every gate present, no peepholes, projection or + layer normalization, and a tanh cell activation.""" + ins = ins + [None] * (24 - len(ins)) + if any(ins[i] is None for i in (0, *_LSTM_CONSTANTS, 18, 19)): + return "a missing gate, which TFLite Micro requires" + if any(ins[i] is not None for i in (9, 10, 11, 16, 17, 20, 21, 22, 23)): + return "peepholes, projection or layer normalization, which TFLite Micro does not run" + if len(ins[0].shape) != 3: + return "input of rank other than 3" + if _activation(op.option(0, "b")) != "tanh": + return f"cell activation {_activation(op.option(0, 'b'))}" + if any(not _is_constant(ins[i]) for i in _LSTM_CONSTANTS): + return "weights and biases must be constant" + tensors = [ins[i] for i in (0, *_LSTM_CONSTANTS, 18, 19)] + outs + if all(t.type == "FLOAT32" for t in tensors): + return "" + return "int8 LSTM does not convert yet" + + 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] @@ -740,6 +767,9 @@ def convert(self, op: _Operator) -> None: if kind == "SVDF": self.finish(self._svdf(op, tag), outs[0]) return + if kind == "UNIDIRECTIONAL_SEQUENCE_LSTM": + self.finish(self._lstm(op, tag), outs[0]) + return if kind in ("CONV_2D", "DEPTHWISE_CONV_2D"): y = self._conv(op, tag) elif kind == "TRANSPOSE_CONV": @@ -1171,6 +1201,28 @@ def _svdf(self, op: _Operator, tag: str) -> str: self.held_state[state]["output"] = kept return y + def _lstm(self, op: _Operator, tag: str) -> str: + """UNIDIRECTIONAL_SEQUENCE_LSTM in the compiler's own form, run in + TFLite's axis order, its hidden and cell states passed in and out.""" + b = self.b + ins = op.inputs + x = self.last_to_last(self.value(ins[0]), 3, tag + "_last", ins[0], True) + constants = [b.constant(self.tensors[ins[i]].array(), f"{tag}_w{i}") for i in _LSTM_CONSTANTS] + y = b.unique(tag + "_sequence") + kept = [b.unique(tag + "_hidden"), b.unique(tag + "_cell")] + b.nodes.append(helper.make_node( + "Lstm", [x, *constants, *(self.held_state[i]["input"] for i in ins[18:20])], + [y, *kept], domain="tigris", time_major=int(op.option(3, "?", False)), + cell_clip=float(op.option(1, "f", 0.0)))) + out = self.tensors[op.outputs[0]] + self.value_info.append(helper.make_tensor_value_info(y, TensorProto.FLOAT, list(out.shape))) + for index, name in zip(ins[18:20], kept): + held = self.tensors[index] + self.value_info.append(helper.make_tensor_value_info( + name, _STATE_TYPES[held.type], list(held.shape))) + self.held_state[index]["output"] = name + return self.tflite_order_out(y, op.outputs[0], tag) + 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.""" @@ -1396,18 +1448,17 @@ def to_onnx(data: bytes, name: str) -> onnx.ModelProto: _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 index in [op.inputs[slot] for op in operators if op.kind in _STATEFUL + for slot in _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):], diff --git a/src/tigris/loaders/onnx/normalize.py b/src/tigris/loaders/onnx/normalize.py index 7622797..86b756e 100644 --- a/src/tigris/loaders/onnx/normalize.py +++ b/src/tigris/loaders/onnx/normalize.py @@ -57,6 +57,7 @@ def normalize(ag: AnalyzedGraph) -> AnalyzedGraph: declared_outputs = list(ag.model_outputs) ag = _adopt_tflite_cumsum(ag) ag = _adopt_svdf(ag) + ag = _adopt_lstm(ag) ag = _normalize_arg_outputs(ag) ag = _drop_inference_identities(ag) ag = _lower_legacy_softmax(ag) @@ -91,6 +92,7 @@ def normalize(ag: AnalyzedGraph) -> AnalyzedGraph: ag = _absorb_activations(ag) ag = _fold_split_into_its_weight(ag) ag = _assign_tensor_layouts(ag) + ag = _align_state_layouts(ag) ag = _normalize_concat_axis(ag) ag = _mark_untileable_broadcasts(ag) ag = _gather_one_index_to_split(ag) @@ -284,6 +286,7 @@ def _strip_metadata_inputs(ag: AnalyzedGraph) -> AnalyzedGraph: # the reduction axis and compute something else. _LINEAR_LAYOUT_OPS = frozenset({ "MatMul", + "Lstm", }) @@ -475,6 +478,19 @@ def _agreed_layout(ag: AnalyzedGraph, op: OpNode) -> Layout | None: return Layout.SPATIAL if spatial > linear else Layout.LINEAR +def _align_state_layouts(ag: AnalyzedGraph) -> AnalyzedGraph: + """A state leaves in the layout it entered in. Below rank 3 the two layouts + store the same bytes, so the output takes the input's.""" + for port in ag.state_ports: + if port.output is None: + continue + entered = ag.tensors[ag.model_inputs[port.input]] + left = ag.tensors[ag.model_outputs[port.output]] + if len(entered.shape) < 3: + left.layout = entered.layout + return ag + + def _assign_tensor_layouts(ag: AnalyzedGraph) -> AnalyzedGraph: """Give every tensor a layout and convert where producer and consumer differ. @@ -2947,6 +2963,38 @@ def _adopt_svdf(ag: AnalyzedGraph) -> AnalyzedGraph: return ag +def _adopt_lstm(ag: AnalyzedGraph) -> AnalyzedGraph: + """The compiler's own LSTM: a sequence [batch, steps, features], or + [steps, batch, features] when time_major, then the input and recurrent + weights and the biases of gates i, f, c, o, then hidden and cell states + [batch, units]; it writes the hidden sequence and both states back.""" + for op in ag.ops: + if op.op_type != "tigris::Lstm": + continue + op.op_type = "Lstm" + if len(op.inputs) != 15 or len(op.outputs) != 3: + raise ValueError("Lstm requires fifteen inputs and three outputs") + x, y = ag.tensors[op.inputs[0]], ag.tensors[op.outputs[0]] + weights = [ag.weight_data.get(name) for name in op.inputs[1:13]] + if any(w is None for w in weights): + raise ValueError("Lstm requires constant weights and biases") + if len(x.shape) != 3: + raise ValueError("Lstm requires a rank-3 input") + time_major = int(op.attrs.get("time_major", 0)) + steps, batch = (x.shape[0], x.shape[1]) if time_major else (x.shape[1], x.shape[0]) + features = x.shape[2] + units = weights[0].shape[0] + expected = [(units, features)] * 4 + [(units, units)] * 4 + [(units,)] * 4 + sequence = (steps, batch, units) if time_major else (batch, steps, units) + states = [ag.tensors[name] for name in (*op.inputs[13:], *op.outputs[1:])] + if ([w.shape for w in weights] != expected or tuple(y.shape) != sequence + or any(tuple(s.shape) != (batch, units) for s in states)): + raise ValueError("Lstm shapes do not agree") + if float(op.attrs.get("cell_clip", 0.0)) < 0.0: + raise ValueError("Lstm cell clip must not be negative") + 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 4c8d8d3..e69616f 100644 --- a/src/tigris/schema/operator-capabilities-v1.json +++ b/src/tigris/schema/operator-capabilities-v1.json @@ -1124,6 +1124,19 @@ "reference": "native", "s8_ref": "native" } + }, + { + "constraints": [ + "hidden and cell state kept between runs; every gate, no peepholes, projection or layer normalization; tanh cell activation; float32; untiled" + ], + "opcode": 86, + "operator": "Lstm", + "routes": { + "cmsis-nn": "unsupported", + "esp-nn": "unsupported", + "reference": "native", + "s8_ref": "unsupported" + } } ], "schema_version": 10, diff --git a/src/tigris/schema/tigris-plan-v10.json b/src/tigris/schema/tigris-plan-v10.json index 839f4b4..021260a 100644 --- a/src/tigris/schema/tigris-plan-v10.json +++ b/src/tigris/schema/tigris-plan-v10.json @@ -20,6 +20,7 @@ "constants": 14, "cumsum_options": 11, "epsilon": 2, + "lstm": 16, "movement": 12, "pads": 5, "pool_rounding": 8, @@ -73,6 +74,7 @@ "LessOrEqual": 75, "Log": 45, "LogSoftmax": 57, + "Lstm": 86, "MatMul": 25, "Max": 41, "MaxPool": 5, diff --git a/tests/fixtures/tflite/ops/float_logistic_tails.npz b/tests/fixtures/tflite/ops/float_logistic_tails.npz new file mode 100644 index 0000000000000000000000000000000000000000..26dd905be3c72f539c0abdd93b8c3073542e505a GIT binary patch literal 2316 zcmZvec{CJy8^>pCF(&&`BSw*;jNLG<2~DN!*|Ls^u`kIoVXO_ZMse+flw_NP5Xrtx z_GB4jZDiktaCz^2&$+$leSYWpeb0H$^ZoProUg71EuA<30AM`s^8mxWTCbyjiW$HK zaBy|^dW4pQxw>P405(7n-SO6O!T!;Vwq2PqpELDqF2yks^QsK&Jh=<^ZlBAI9l!}Z zxK-WEVLD_*z92+gewBKitx`SGG*YNAciHQC(E%>A>3z0ThrM2o@Q8Y}P=>agi+X=x z_6lXZgioMu`N9$Wo@gQSSVUWsJJV2bqcXc~iDqPY1|7nD>ChAkJJic`Z(7*-H0vjKKmv4ngAqqT8j=>WOM!O>Vr z!~own(WGfERuIZtCg(mh&KcFqT*e*qYudO@c+?g<*Sec%IvS-LWx+(6WvL%`YzdCP znjZP)vqPx8{fKDpj`wckdXLMmVggGEBzf|10;`6Upp6o{Q~((V3Fs>h+*blpz($`h z@0@D_oaX9OlR^)Nkcgq*ENt$)e%v@0WXg5WQFAIrgU`P1Drd}KvCT(m=6%*7>uf)1 z1a-jr<~k+1lYAfXhb542M>?CjaL(ShTD9hihSh^j#ti0ZeTgS>Db;6Cbra;^%}R#J z{6S7Tk}1=txf@pY4>>EyXd#O8NYLdLy01%9+1E|z#Jy)VEAPct<1vD%Ui7v?X3)ie zL40MM#mj}KtRCfxj9E-qc5#Nl)mCs&kt1)!mgA(e5tH{E=;x2#4B3Hpnoa_Xd@^)x z<=f_T!M2t{st?GdEJ?t3z{e@dWZ&>HLHri7H!E0bxZf-&FEs9jVX( z*(e_=C`)2Hex(%CDnfD=1 zW6Kv%tEKn?4Q9d~N{rjAop^Lg|HIxp$Z+OL-zQi7Je(md3YI**Ze5(XP6W&JCq?*Q z<-OZ=r3bUkpOa_=;yin2E|wYaxdpstVI#T`CfelD7qS;iR;dYMc~utH2QvZMeKe3> zRMdLF7r?Mj6^}Fvfy8Qy)mx%F&VUC6YxK3h$M8-1J1hTD`VNm7$`LF#q-SB}(Mx+` z{?WY`#J@$O+B>9;F|{C)Z_{3tOK#ub=Phu_h*8ZNL$UhsBaqp2D>tAI=ijmQiNkVl zCB(CbxAt6U2B@n-n_EA;-=B7{<|<^^zWq@B>EbF{ z_czoy=QqRuiyH7TYTUdY{a@CM(EiPwv_;-|KDPEY-4@-V%mmYP^Ovr91lL5noZ?(U zW)j-q?(;u+v(V2IKpZ|AU836bPkAxIJTGZ>AtDG69s#lwNkxFlK@Wg120Viw8`4*{ zH*>2MRZ|$awOyo~CxVCSq-a@sbhpI06_X1oTidbEX1&!G=9kds#pUSr%JK1E0ksFA zNZQfUM+-}#*UuG>w~KvC^2tksD;SIhrt8`RWOf{V-yDo#%8f=KKjZ~*a}xT>Y_&A` zC2x$otXShSGgF*_98UJ4vh6RZ@2F^h6lJQ>EN%11Gy48_bNLr=v!hh<;^-eb`(LQ# z+i-cQGrVD+D$S(~2f)H;fb6Ao{b8%Ae(CAv#Bbn0z;21s4HrwG!;#9ZY<1Mchx#2GuO99+;~kU z4Hr8q?__lLhXjmrpM3FJzXB3I*oJkhSc|+t0n4Yi%U89mcc@_Y(4$OgRYj`b5BFJ# zn7k4KcZ9sAsnvc)e{)4Sp$Ca7Y!LWf^4<}s#J2P;+r~7Yw^?PyuC|t4w}f9(0H;dI zuR+`UY`+diw&|Qk_(elk>?*G{xf@(F6X|!PRtd^<&ujNT{{-813tHD8JXZN&Tt1-w zS)Fp&`rU3p^hsxzs&#a7?&9oKT*!32Q(ef#h~m^7j2#q7xp%=UnZE5ynHHPO0Lr&; zM(qZ8r)Z&au{WMRJh}a};=X?W&qZ9fbJ?nTsTBIuBP#gG<9&LS07Z(oS&DozUk9~T zN;ZX*16|+A$evZEFRMdiZnkPRqEoFPD|?=%#!d$V$WDFadx()i-7*AuNC2>O7FPz!~wLk>a1wl$Xz zXcWM2UoJyndR*@@;+^oHSnyOhb=D7jeNTGgJ`E_B$!uXMPC3wxT9?yt_~KE9OCRK- z5D^imq~wWpihA1`{m8_5vB3#$7BmKF(X$WS_^64BXa4I;DU!o zRpZFfk&&A_ba5W!B>xQ`6{*1}i?3B~nAD726F0PP^|)mgX{^FGAbrh0HVO~rDF%d7 z-Krk$d#aRU`(bDJ_PQ%;68z7m>o~t9i+lh2QK9^vI9@l!*MwXUwb<*h>~ujTFUvlY zT$iK$;fXZ2$ub&oDfh5_Wzq;W361xO({wtD{1(zc6xM01+w!PAxFk3#nUC&1@5`@{ zm7(WZ%IPmp-dw$;bTO{V+}otKY+ggx+Xqd literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/float_logistic_tails.tflite b/tests/fixtures/tflite/ops/float_logistic_tails.tflite new file mode 100644 index 0000000000000000000000000000000000000000..63c2989acea1c9b160a4a74055e562c3bfd050eb GIT binary patch literal 656 zcmX|9Jxjw-6g|~gOAWPPkq#XSb`G`{CkIQdSm=k;B7#Cynju3|sHrIUE1Vr29sCt8 z&TfK#!NI}Fc+N|c8xHsLo%im&DFGNBc6JIVql6*~$irqV1B*5oix5y(Qb3*9TLPBp zOK74_Sj1--hSI_Y(d_q5TZ5B!zc=c(hK+;9ut8qiaA%koM_fDc{LyUg2j0XT-MX_I zZ|Y-(`4{4ucp@GM)juxoxul&ZBpKuqUd1XVnl`aCZJm?YNspH4^WuE>{o{J;OK(vx zs#AT1lr2lelBhOl-v@>gkLS=Nsonl@gzzT#oG8cFFt5mYa2nJ@D1=c_QFgz}k6}qUOXb$y>T0Z*tnAb&D6wn-(7uG-ZbP{P-!8xP;2R&uIFwc8RB} zh%I6IDWzt`)n>(Y%!*5GmKE1lt`vrKFLS;4zOrxS+f!n_egEau#=m#7m+@VkTK)F$ zw9}R+_NCq4?&NpE@lwbEt&hvsG(T|N$=o03`u_6w-n-Lx-Mz8zQ`(xj-LrT9J-Q_C z_m=eU-#5pdyYb%EIH!L0r9*q4*Vf;BU-Rqsy!WTi|DE(Y?R|LU*OdK{t6nr7w!d~; zeVmq7zn?`UaY*;3E)Z-1UY+cGuoP~zE#FN{(n0A==ZI4cWy;jKmHxJS)#qRpP%#n z@9f*H`neJyz@J6v(bC+l~ZqAe{*LWKQJN~nRJ;^lQ$^sg3>WJBq@W*25`_LC3bX; jpg07@4=Ap9P&9@BMUZ1Oz?+o~q<{$si-GiACJ+w*q9TH& literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/float_lstm.tflite b/tests/fixtures/tflite/ops/float_lstm.tflite new file mode 100644 index 0000000000000000000000000000000000000000..d8c5cb3d981ee7d8e7acf7c22804c9765508ab98 GIT binary patch literal 2020 zcmY*a3s96*6u!jj5-x~=4qIzlq9Tt3@d3hrE|BsVS9B!bk9?(xsA$lLibG=#C^3Sr zMu>!f*aHicRCND)(bCMxVH9@K_T?|lD#eBZt2oO}Os)|WB% zL7Xv+F;C{hyqJ!;;fw}H3`dOlX&5WhGUkf?K{SVB4Wbe!q7W^wF}4lihA_D>HXKon zx?IFcL^Q&H@ImMhxRiyzc=aj@5sml`L36~LyeYwe?9vx zc$|6+I%69JJDjBK31w<~FMfsky=e89;&9Vv=D;xjc1n-Y$ ziMO{0@}`C#!1I$mV#dkm%Es;;;?~?`?&|5u1L{KLnDJFg!o>=6T)eAX*92$XU)tiQ z_$GJ7^KY}W)XBBQk!c=2a^suDkbVAh{^>=N7`{l$8)ua$(}D}c?sM*ZN?tIZWjPqB z`%uplul)*1cRZjia;=obS+etJiI`ZfmD2-P@`5#;N;zv2J{}iD_||@~rI)?R?qO>A zu^d>gw5!YhOOs(oqGal!d@<%(y{vdR8`d2=4h4fE;QUg?Ck1Yhb?Zlqd(Vu}uy};x zuP;?^w(21@_muh1BlF~oz4b7*V}xv~-6wJ$7x4&}9pZk=1*P3?ZzjgK<#{&z66+0? z#dpP)_d3MZZ|a26RaFY}_lWlTz4Fo8L9)3rKxv%*w~~COC~DjMJhAqYmpUhK9Cw>B z6^g3os9m~h#Z3-?^59W2?v?HAXz%<{UqNXGsPQTD;gH5`LRnrNoySc*f0n@NUK?vHzzWSUvD-(Uq)&^GrD0Jx-gP$cTxoV)i zb)Zr4Ok|3U4?*!_sj^lubkFg;j&tsa0SKx=J%k7L$#&1Ur|q#t zEj<)RIvR^c_#z%4aS2g|*o>Ho@JEo7ZK$PKix8=Z5Cr*W!MrEiv=_BhPt!vuSX*3X zZ7Bu|I)RM~X((Q&c08kaNlQ8}>sb>|^;&D)muS*cOwM&-ebz0-d?AivcWTBt_RY5Q z#5ws;2VZ9Ad(P4pM~H)Ou=8}@w4QuUk8t|kX6NZn9DKMDh0@ zUr7s`JP~j#VC}p~QFCIJ&)+kJlJ`~Q+#67KJB=Kl6SX~*vCbCQq$Q{8>QI7F;!R*ue>)*}&X>se> z-FNfLE=e~2{wHe}jCcB!|Ls*^h3&$|=i zwp%x}JvU|D+pWcRD>uEL5^^!FQPwy|dY<9U?I%^=Y~6b+ar3{ACe|}gJ}vumI`S>k z*Eiq(y-V9#zx~*$XD=>Y-WGQ9dv1BF$&INWjNe@KuFRDFvM%BMzo+VNY~B^W^IdUb z>&mY@TQ2X~e9BC1X}*8M7p3TZ&wlJ(dAGLAQvc7la~2Epvwv?tk+(A9_dWmqobRV* zh(#7>&3b?Kif!`u)RQ8+cR!l`CM?H3!QA}Ct$R6j*Y<9jJ~@{$z?+dtml-u*gEB2B y19L-?0+?(7rw625j;;}uu)v81n1FatH0}UuMNVb`-mGjO1x!F#45Y6ygLnYN02WdJ literal 0 HcmV?d00001 diff --git a/tests/fixtures/tflite/ops/float_lstm_time_major_clip.tflite b/tests/fixtures/tflite/ops/float_lstm_time_major_clip.tflite new file mode 100644 index 0000000000000000000000000000000000000000..c5654d80422eba941a9918314293ac0d83982139 GIT binary patch literal 1744 zcmY+FdrVVT9LEn>!9nF=Ac9nR=vduwSRS@u@9(OadxZeo$g+`UL!*Qd%wApGmIl}+ zoeoqsm~JRp65(W1A-QM1|C;9Z8d%xfF z_?`RvK^kMM>cO%?#u8WxOJYX00^b<$iQyBDmNC{A!x(lf8JUY_Y#QlB9Eb@Sk3tV( zN3NsiB;rCGh#fH@DTom%N8qf#jyopqQ7&YAmEkaLbyur$ z`;Y~sbLB9b7$Z-GX|*r!1k=GCgKhqD<@q5eEQV>VSF&L&w|TI)-V3eUt;(9kz<%3s zzwyciNy_L|Og%5bh0#xVeX2vO^?xCs+<1vUaLx<+hjQglKmHXaPuRqnKgW3T_t3(~o@S$QAzRQsBNe z<9$`_QXu`84r!AwTdc^;lATf=UINL{`k_~)zs;By%pVD+A3cIom#`1#m7@O{|akEZ_v z*ZCH{%iJT?ud9G@SEE>;YLqkb4vT?=SP^Zri_(1_zHxWD9R2zTXDkM@*nq!*e513E zJkt*DzPrJ?Zi!brO2zyCn#GNoW1o0rRB9(|4p+iof*N1$7G$3@(biPYif3!EPMLqrPa9WDhp2if5>cd0e zFp4}BZ^&DudV;4ChW+M*Vio2QOh-v4BHtA79h literal 0 HcmV?d00001 diff --git a/tests/test_frontend_tflite.py b/tests/test_frontend_tflite.py index 3b40576..f9b10e7 100644 --- a/tests/test_frontend_tflite.py +++ b/tests/test_frontend_tflite.py @@ -303,6 +303,25 @@ def mark(tensors, operators): assert any("is not the state of one SVDF" in r for r in reasons) +@pytest.mark.parametrize(("edit", "reason"), [ + (lambda ops: ops[0].inputs.__setitem__(1, -1), "a missing gate, which TFLite Micro requires"), + (lambda ops: ops[0].inputs.__setitem__(9, 5), "peepholes, projection or layer normalization"), +]) +def test_an_lstm_tflite_micro_cannot_run_is_refused(monkeypatch, edit, reason): + def change(tensors, operators): + operators[0].inputs = list(operators[0].inputs) + edit(operators) + reasons = _read_edited(monkeypatch, "float_lstm", change) + assert any(reason in r for r in reasons) + + +def test_an_lstm_cell_activation_other_than_tanh_is_refused(monkeypatch): + def relu(tensors, operators): + operators[0].option = lambda slot, fmt, default=0: 1 if slot == 0 else default + reasons = _read_edited(monkeypatch, "float_lstm", relu) + assert any("cell activation relu" 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"]