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 0000000..26dd905 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_logistic_tails.npz differ 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 0000000..63c2989 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_logistic_tails.tflite differ diff --git a/tests/fixtures/tflite/ops/float_lstm.npz b/tests/fixtures/tflite/ops/float_lstm.npz new file mode 100644 index 0000000..f53fdab Binary files /dev/null and b/tests/fixtures/tflite/ops/float_lstm.npz differ diff --git a/tests/fixtures/tflite/ops/float_lstm.tflite b/tests/fixtures/tflite/ops/float_lstm.tflite new file mode 100644 index 0000000..d8c5cb3 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_lstm.tflite differ diff --git a/tests/fixtures/tflite/ops/float_lstm_time_major_clip.npz b/tests/fixtures/tflite/ops/float_lstm_time_major_clip.npz new file mode 100644 index 0000000..d89514e Binary files /dev/null and b/tests/fixtures/tflite/ops/float_lstm_time_major_clip.npz differ 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 0000000..c5654d8 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_lstm_time_major_clip.tflite differ 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"]