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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 56 additions & 0 deletions scripts/crossrepo_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
41 changes: 38 additions & 3 deletions scripts/gen_tflite_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)],
}


Expand Down Expand Up @@ -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
Expand All @@ -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",
Expand All @@ -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]
Expand Down
1 change: 1 addition & 0 deletions scripts/generate_schema_package.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
3 changes: 3 additions & 0 deletions src/tigris/capabilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,9 +113,11 @@ class KernelCapabilities:
"ReduceAll",
"Split",
"Svdf",
"Lstm",
})

_S8_REFERENCE_OPERATORS = _FLOAT_REFERENCE_OPERATORS - frozenset({
"Lstm",
"L2Pool",
"Neg",
"Exp",
Expand Down Expand Up @@ -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",),
Expand Down
1 change: 1 addition & 0 deletions src/tigris/dtypes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")},
})
Expand Down
3 changes: 3 additions & 0 deletions src/tigris/emitters/binary/defs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -187,6 +189,7 @@
"Sum": 83,
"ReduceAll": 84,
"Svdf": 85,
"Lstm": 86,
}
OP_TYPE_UNKNOWN = 255

Expand Down
7 changes: 6 additions & 1 deletion src/tigris/emitters/binary/writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
"<if", int(op.attrs.get("time_major", 0)), float(op.attrs.get("cell_clip", 0.0)))))
continue
if op.op_type in {"Resize", "ResizeLinear"}:
scales = op.attrs.get("resize_scales")
if scales is not None:
Expand Down Expand Up @@ -1051,7 +1056,7 @@ def _build_op_attributes(


# Operators whose constants are listed by OP_ATTR_CONSTANTS, not weight and bias.
_MANY_CONSTANTS = {"Svdf"}
_MANY_CONSTANTS = {"Svdf", "Lstm"}


def _resolve_weight_bias(op: OpNode, weight_idx: dict[str, int]) -> tuple[int, int]:
Expand Down
81 changes: 66 additions & 15 deletions src/tigris/frontends/tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -204,15 +204,16 @@ 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
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")
"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")
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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":
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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):],
Expand Down
Loading
Loading