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
44 changes: 33 additions & 11 deletions scripts/gen_tflite_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -557,25 +557,45 @@ 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):
def _lstm(time_major, batch, steps, features, units, cell_clip=0.0, quantized=False):
"""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."""
are variable tensors. Seeded weights, gates in TFLite's order i, f, c, o.
The int8 form keeps its cell in int16 at a power-of-two scale."""
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)]
x_scale, hidden_scale, hidden_zero = 6.0 / 255, 1.0 / 128, 3
if quantized:
tensors = [(shape, T.INT8, None, x_scale, 2, False)]
else:
tensors = [(shape, T.FLOAT32, None, None, 0, False)]
input_scales = []
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)]
w = rng.normal(0.0, 0.5, (units, columns)).astype(np.float32)
if quantized:
scale = float(np.abs(w).max()) / 127
input_scales.append(scale)
tensors.append(((units, columns), T.INT8, _quantize(w, scale, np.int8), scale, 0, False))
else:
tensors.append(((units, columns), T.FLOAT32, w, None, 0, False))
for gate in range(4):
b = rng.normal(0.0, 0.3, (units,)).astype(np.float32)
if quantized:
scale = float(np.float32(x_scale) * np.float32(input_scales[gate]))
tensors.append(((units,), T.INT32, _quantize(b, scale, np.int32), scale, 0, False))
else:
tensors.append(((units,), T.FLOAT32, b, None, 0, False))
if quantized:
tensors += [((batch, units), T.INT8, None, hidden_scale, hidden_zero, True),
((batch, units), T.INT16, None, 2.0 ** -11, 0, True),
(out, T.INT8, None, hidden_scale, hidden_zero, False)]
else:
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
Expand All @@ -596,6 +616,8 @@ def _lstm(time_major, batch, steps, features, units, cell_clip=0.0):
"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),
"lstm": lambda: _lstm(False, 1, 3, 4, 5, quantized=True),
"lstm_time_major_clip": lambda: _lstm(True, 2, 3, 3, 4, cell_clip=0.8, quantized=True),
}
# 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.
Expand Down
3 changes: 1 addition & 2 deletions src/tigris/capabilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,6 @@ class KernelCapabilities:
})

_S8_REFERENCE_OPERATORS = _FLOAT_REFERENCE_OPERATORS - frozenset({
"Lstm",
"L2Pool",
"Neg",
"Exp",
Expand Down Expand Up @@ -252,7 +251,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",),
"Lstm": ("hidden and cell state kept between runs; every gate, no peepholes, projection or layer normalization; tanh cell activation; int8 with an int16 cell state; 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
3 changes: 2 additions & 1 deletion src/tigris/emitters/binary/defs.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,8 @@
# 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_LSTM = 16 # int32 time_major (0 or 1), float32 cell clip (0: none); int8
# adds zero points, cell power, clip and 11 multiplier pairs

OP_ATTR_KINDS = (
OP_ATTR_COMPARISON_REQUANT,
Expand Down
35 changes: 33 additions & 2 deletions src/tigris/emitters/binary/writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -920,6 +920,38 @@ def _svdf_payload(ag: AnalyzedGraph, op: OpNode) -> bytes:
return struct.pack("<6i", rank, int(op.attrs["state_zero_point"]), *first, *second)


def _lstm_payload(ag: AnalyzedGraph, op: OpNode) -> bytes:
"""time_major and the cell clip; an int8 Lstm adds the input and hidden
zero points, the cell scale's power of two, the clip in cell units, then
per gate i, f, c, o the input and recurrent (multiplier, shift) into the
2^-12 gate scale, and the forget, input and output products' pairs. Each
is formed as TFLite Micro's Prepare forms it."""
payload = struct.pack("<if", int(op.attrs.get("time_major", 0)),
float(op.attrs.get("cell_clip", 0.0)))
if not ag.is_quantized:
return payload
x = ag.tensors[op.inputs[0]].quant
hidden = ag.tensors[op.inputs[13]].quant
x_scale, h_scale = float(np.float32(x.scale[0])), float(np.float32(hidden.scale[0]))
weights = [float(np.float32(scale)) for scale in op.attrs["weight_scales"]]
cell = np.float32(op.attrs["cell_scale"])
f32 = np.float32
power = int(_round_half_away(float(f32(np.log(cell)) * (f32(1.0) / f32(np.log(f32(2.0)))))))
clip = float(np.float32(op.attrs.get("cell_clip", 0.0)))
clipped = int(min(max(clip / float(cell), -32768.0), 32767.0))
gate, nonlinear = 2.0 ** -12, 2.0 ** -15
pairs = []
for k in range(4):
pairs += _compute_multiplier_shift(x_scale * weights[k] / gate)
pairs += _compute_multiplier_shift(h_scale * weights[4 + k] / gate)
pairs += _compute_multiplier_shift(nonlinear * float(cell) / float(cell))
pairs += _compute_multiplier_shift(nonlinear * nonlinear / float(cell))
pairs += _compute_multiplier_shift(nonlinear * nonlinear / h_scale)
return payload + struct.pack(
f"<4i{len(pairs)}i", int(x.zero_point[0]), int(hidden.zero_point[0]), power, clipped,
*pairs)


def _build_op_attributes(
ag: AnalyzedGraph, tensor_idx: dict[str, int],
quant_idx_map: dict[str, int] | None = None,
Expand All @@ -944,8 +976,7 @@ def _build_op_attributes(
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)))))
records.append((op_index, OP_ATTR_LSTM, _lstm_payload(ag, op)))
continue
if op.op_type in {"Resize", "ResizeLinear"}:
scales = op.attrs.get("resize_scales")
Expand Down
51 changes: 41 additions & 10 deletions src/tigris/frontends/tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -304,7 +304,7 @@ def _activation(code: int) -> str:
"REDUCE_ALL": ((0,), True)}
# Operators that keep state in a variable tensor, at this operand.
_STATEFUL = {"SVDF": (4,), "UNIDIRECTIONAL_SEQUENCE_LSTM": (18, 19)}
_STATE_TYPES = {"FLOAT32": TensorProto.FLOAT, "INT16": TensorProto.INT16}
_STATE_TYPES = {"FLOAT32": TensorProto.FLOAT, "INT16": TensorProto.INT16, "INT8": TensorProto.INT8}
# 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.
Expand Down Expand Up @@ -531,7 +531,16 @@ def _lstm_reason(op: _Operator, ins: list, outs: list[_Tensor]) -> str:
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"
x, hidden, cell, y = ins[0], ins[18], ins[19], outs[0]
if (x.type != "INT8" or hidden.type != "INT8" or y.type != "INT8"
or any(ins[i].type != "INT8" for i in range(1, 9))
or any(ins[i].type != "INT32" for i in range(12, 16))):
return "activations must be all int8 or all float32"
if cell.type != "INT16" or cell.zero_point.size != 1 or cell.zero_point[0] != 0:
return "cell state other than symmetric int16"
if any(len(t.scale) != 1 for t in (x, hidden, cell, y, *(ins[i] for i in range(1, 9)))):
return "operands must be quantized per tensor"
return ""


def _data_movement_reason(op: _Operator, ins: list[_Tensor], outs: list[_Tensor]) -> str:
Expand Down Expand Up @@ -1208,20 +1217,32 @@ def _lstm(self, op: _Operator, tag: str) -> str:
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]
attributes = {"time_major": int(op.option(3, "?", False)),
"cell_clip": float(op.option(1, "f", 0.0))}
if self.tensors[ins[0]].type == "INT8":
# Integer weights pass as stored, their scales stated; the int16
# cell state stays raw.
attributes.update(weight_scales=[float(self.tensors[ins[i]].scale[0]) for i in range(1, 9)],
cell_scale=float(self.tensors[ins[19]].scale[0]))
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))))
"Lstm", [x, *constants, self.value(ins[18]), self.held_state[ins[19]]["input"]],
[y, *kept], domain="tigris", **attributes))
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)))
if held.type == "INT8":
self.value_info.append(helper.make_tensor_value_info(
name, TensorProto.FLOAT, list(held.shape)))
scale, point = self.held_state[index]["quantization"]
name = b.node("QuantizeLinear", [name, scale, point], name + "_q")
else:
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)
return self.tflite_order_out(self.held(y, op.outputs[0], tag + "_q"), 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
Expand Down Expand Up @@ -1454,10 +1475,20 @@ def to_onnx(data: bytes, name: str) -> onnx.ModelProto:
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]),
# An int8 variable starts at its zero point and is read through its
# quantization; an int16 one is passed to its operator as stored.
zero = int(tensor.zero_point[0]) if tensor.type == "INT8" else 0
initial = b.constant(np.full(tensor.shape, zero, _NUMPY[tensor.type]),
name.removesuffix("_in") + "_initial")
converter.held_state[index] = {"input": name, "output": None}
converter.values[index] = name
if tensor.type == "INT8":
scale = b.constant(np.float32(tensor.scale[0]), name + "_scale")
point = b.constant(np.array(zero, np.int8), name + "_zero_point")
converter.values[index] = b.node("DequantizeLinear", [name, scale, point],
name + "_float")
converter.held_state[index]["quantization"] = (scale, point)
else:
converter.values[index] = name
state.append({"input": name, "output": None, "initial": initial})
for op in operators:
converter.convert(op)
Expand Down
3 changes: 3 additions & 0 deletions src/tigris/loaders/onnx/normalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -2992,6 +2992,9 @@ def _adopt_lstm(ag: AnalyzedGraph) -> AnalyzedGraph:
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")
if "weight_scales" in op.attrs and (len(op.attrs["weight_scales"]) != 8
or "cell_scale" not in op.attrs):
raise ValueError("an integer Lstm states its eight weight scales and its cell scale")
return ag


Expand Down
8 changes: 4 additions & 4 deletions src/tigris/schema/operator-capabilities-v1.json
Original file line number Diff line number Diff line change
Expand Up @@ -1127,15 +1127,15 @@
},
{
"constraints": [
"hidden and cell state kept between runs; every gate, no peepholes, projection or layer normalization; tanh cell activation; float32; untiled"
"hidden and cell state kept between runs; every gate, no peepholes, projection or layer normalization; tanh cell activation; int8 with an int16 cell state; untiled"
],
"opcode": 86,
"operator": "Lstm",
"routes": {
"cmsis-nn": "unsupported",
"esp-nn": "unsupported",
"cmsis-nn": "fallback:s8_ref",
"esp-nn": "fallback:s8_ref",
"reference": "native",
"s8_ref": "unsupported"
"s8_ref": "native"
}
}
],
Expand Down
Binary file added tests/fixtures/tflite/ops/lstm.npz
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/lstm.tflite
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/lstm_time_major_clip.npz
Binary file not shown.
Binary file not shown.
16 changes: 16 additions & 0 deletions tests/test_frontend_tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -328,6 +328,22 @@ def relu(tensors, operators):
assert any("cell activation relu" in r for r in reasons)


def test_an_int8_lstm_cell_is_symmetric_int16(monkeypatch):
def offset(tensors, operators):
tensors[operators[0].inputs[19]].zero_point = np.asarray([5], np.int64)
reasons = _read_edited(monkeypatch, "lstm", offset)
assert any("cell state other than symmetric int16" in r for r in reasons)


def test_an_int8_lstm_with_per_channel_weights_is_refused(monkeypatch):
"""TFLite Micro's LSTM reads one multiplier per gate projection."""
def per_channel(tensors, operators):
weight = tensors[operators[0].inputs[1]]
weight.scale = np.repeat(weight.scale, weight.shape[0])
reasons = _read_edited(monkeypatch, "lstm", per_channel)
assert any("quantized per tensor" 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"]
Expand Down
Loading