diff --git a/scripts/gen_tflite_fixtures.py b/scripts/gen_tflite_fixtures.py index 848034a..5d8d70c 100644 --- a/scripts/gen_tflite_fixtures.py +++ b/scripts/gen_tflite_fixtures.py @@ -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 @@ -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. diff --git a/src/tigris/capabilities.py b/src/tigris/capabilities.py index 18a1127..72917d0 100644 --- a/src/tigris/capabilities.py +++ b/src/tigris/capabilities.py @@ -117,7 +117,6 @@ class KernelCapabilities: }) _S8_REFERENCE_OPERATORS = _FLOAT_REFERENCE_OPERATORS - frozenset({ - "Lstm", "L2Pool", "Neg", "Exp", @@ -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",), diff --git a/src/tigris/emitters/binary/defs.py b/src/tigris/emitters/binary/defs.py index a31a1e7..d0ca3f2 100644 --- a/src/tigris/emitters/binary/defs.py +++ b/src/tigris/emitters/binary/defs.py @@ -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, diff --git a/src/tigris/emitters/binary/writer.py b/src/tigris/emitters/binary/writer.py index cd69e7d..0a31b49 100644 --- a/src/tigris/emitters/binary/writer.py +++ b/src/tigris/emitters/binary/writer.py @@ -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("