diff --git a/scripts/gen_tflite_fixtures.py b/scripts/gen_tflite_fixtures.py index 5c99e88..4142ed7 100644 --- a/scripts/gen_tflite_fixtures.py +++ b/scripts/gen_tflite_fixtures.py @@ -348,6 +348,22 @@ def rewrite(model: bytes) -> bytes: if _xla is not None: INDEXING["dynamic_update_slice"] = _binary( lambda x, u: _xla.dynamic_update_slice(x, u, tf.constant([0, 3, 1, 0])), _MAP, (1, 2, 3, 4)) +# Indices supplied at run time: int32 inputs drawn from [low, high), keyed by +# case and input position. Dynamic update starts reach past both ends, where +# TFLite clamps them. +INDEX_INPUTS = {"gather_runtime": {1: (0, 6)}, "gather_nd_runtime": {1: (0, 6)}, + "embedding_lookup_runtime": {1: (0, 6)}, + "dynamic_update_slice_runtime": {2: (-2, 7)}} +INDEXING.update({ + "gather_runtime": ([_MAP, (3,)], lambda x, i: tf.gather(x, i, axis=2)), + "gather_nd_runtime": ([(6, 6, 4), (3, 2)], lambda x, i: tf.gather_nd(x, i)), + "embedding_lookup_runtime": ([(6, 4), (4,)], lambda x, i: tf.gather(x, i)), + "arg_max_gather": _unary(lambda x: tf.gather(x, tf.argmax(x, 0, output_type=tf.int32)), + (6, 4)), +}) +if _xla is not None: + INDEXING["dynamic_update_slice_runtime"] = ( + [_MAP, (1, 2, 3, 4), (4,)], lambda x, u, s: _xla.dynamic_update_slice(x, u, s)) # A Keras LSTM unrolled over its time steps, which the converter writes as # plain operators; seeded weights keep the case reproducible. _LSTM = tf.keras.layers.LSTM(4, return_sequences=True, unroll=True, @@ -452,6 +468,7 @@ def _embedding_lookup(model: bytes) -> bytes: REWRITES["embedding_lookup"] = _embedding_lookup +REWRITES["embedding_lookup_runtime"] = _embedding_lookup # The tier-1 cases again, converted without quantization. _FLOAT_TIER1 = ( "max_pool_valid", "max_pool_same", "avg_pool_valid", "avg_pool_same", "concat_channels", @@ -477,19 +494,24 @@ def _embedding_lookup(model: bytes) -> bytes: REWRITES.update({"elu": _int8_island(), "cumsum": _int8_island(True), "cumsum_exclusive_reverse": _int8_island(True), "dynamic_update_slice": _int8_island(shared=True), + "dynamic_update_slice_runtime": _int8_island(shared=True), "not_equal": _int8_island(), "add_n": _int8_island(), "cumsum_offset": _int8_island(), "cumsum_offset_exclusive_reverse": _int8_island()}) -def _convert(fn, shapes, ranges, rng, float_io=False, quantize=True, trackable=None): - specs = [tf.TensorSpec(shape, tf.float32) for shape in shapes] +def _convert(fn, shapes, ranges, rng, float_io=False, quantize=True, trackable=None, + index_inputs=None): + index_inputs = index_inputs or {} + specs = [tf.TensorSpec(shape, tf.int32 if i in index_inputs else tf.float32) + for i, shape in enumerate(shapes)] concrete = tf.function(fn).get_concrete_function(*specs) def representative(): for _ in range(64): - yield [rng.uniform(lo, hi, shape).astype(np.float32) - for shape, (lo, hi) in zip(shapes, ranges)] + yield [rng.integers(*index_inputs[i], shape, dtype=np.int32) if i in index_inputs + else rng.uniform(lo, hi, shape).astype(np.float32) + for i, (shape, (lo, hi)) in enumerate(zip(shapes, ranges))] converter = tf.lite.TFLiteConverter.from_concrete_functions( [concrete], trackable if trackable is not None else tf.function(fn)) @@ -503,8 +525,10 @@ def representative(): return converter.convert() -def _inputs(details, rng, value_range=None): +def _inputs(details, rng, value_range=None, index_range=None): shape = (SAMPLES, *details["shape"]) + if index_range is not None: + return rng.integers(*index_range, shape, dtype=np.int32) if details["dtype"] == np.float32: return rng.uniform(*(value_range or (-3.0, 3.0)), shape).astype(np.float32) if value_range is None: @@ -521,7 +545,8 @@ def generate(name: str) -> bool: ranges = RANGES.get(name) model = _convert(fn, shapes, ranges or [(-3.0, 3.0)] * len(shapes), rng, float_io=name in FLOAT_BOUNDARIES, quantize=name not in FLOAT_MODELS, - trackable=TRACKABLES.get(name.removeprefix("float_"))) + trackable=TRACKABLES.get(name.removeprefix("float_")), + index_inputs=INDEX_INPUTS.get(name.removeprefix("float_"))) if name in REWRITES: model = REWRITES[name](model) micro_interpreter = micro.Interpreter.from_bytes(model, arena_size=1024 * 1024) @@ -543,7 +568,8 @@ def generate(name: str) -> bool: # TFLite's reference resolver lacks some operators (CEIL, ELU, int8 # CUMSUM); TFLite Micro's outputs are recorded unchecked for those. reference = None - inputs = [_inputs(found, rng, ranges[i] if ranges else None) + index_inputs = INDEX_INPUTS.get(name.removeprefix("float_"), {}) + inputs = [_inputs(found, rng, ranges[i] if ranges else None, index_inputs.get(i)) for i, found in enumerate(details)] if name == "div": # TFLite refuses a divisor whose raw byte is 0; keep its value nonzero too. diff --git a/src/tigris/frontends/tflite.py b/src/tigris/frontends/tflite.py index 97359a8..c088994 100644 --- a/src/tigris/frontends/tflite.py +++ b/src/tigris/frontends/tflite.py @@ -191,13 +191,21 @@ def unsupported(data: bytes) -> list[str]: reasons.append(f"subgraph {position}: only constant variable initial values convert") for _, tensors, inputs, outputs, operators in graphs[:1]: indices = {i for op in operators if op.kind in _INDEX for i in op.outputs} - consumed = {i for op in operators for i in op.inputs} + slots = {op.inputs[_INDEX_SLOTS[op.kind]] for op in operators if op.kind in _INDEX_SLOTS} + index_inputs = {i for i in inputs if tensors[i].type == "INT32"} for index in list(inputs) + list(outputs): - if tensors[index].type not in ("INT8", "FLOAT32", "BOOL") and index not in indices: + if (tensors[index].type not in ("INT8", "FLOAT32", "BOOL") and index not in indices + and index not in index_inputs): reasons.append(f"model boundary {tensors[index].name!r} is {tensors[index].type}") - for index in sorted(indices): - if index not in outputs or index in consumed: - reasons.append(f"index {tensors[index].name!r} is not only a model output") + for index in sorted(indices | index_inputs): + uses = [op for op in operators if index in op.inputs] + 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") + for index in sorted(slots - indices - index_inputs): + if not _is_constant(tensors[index]): + reasons.append(f"indices {tensors[index].name!r} are computed by an operator " + "other than ARG_MAX or ARG_MIN") grouped: dict[str, list[_Operator]] = {} for op in operators: reason = _operator_reason(op, tensors) @@ -281,8 +289,10 @@ def _activation(code: int) -> str: "REDUCE_ALL": ((0,), True)} # 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, only as model outputs. +# Index outputs: int32 positions, as model outputs or the indices of the operators below. _INDEX = {"ARG_MAX": "ArgMax", "ARG_MIN": "ArgMin"} +# The operand holding indices or start positions, which may be computed at run time. +_INDEX_SLOTS = {"GATHER": 1, "GATHER_ND": 1, "EMBEDDING_LOOKUP": 0, "DYNAMIC_UPDATE_SLICE": 2} # Reductions over one run of adjacent axes, and the prefix sum over one axis. _REDUCTIONS = {"REDUCE_MAX": "ReduceMax", "REDUCE_MIN": "ReduceMin", "SUM": "ReduceSum", "REDUCE_ALL": "ReduceAll"} @@ -325,7 +335,10 @@ def _operator_reason(op: _Operator, tensors: list[_Tensor]) -> str: return "" for position in _CONSTANT_OPERANDS.get(op.kind, ()): if position < len(ins) and ins[position] is not None and not _is_constant(ins[position]): - return f"input {position} must be a constant" + if position != _INDEX_SLOTS.get(op.kind): + return f"input {position} must be a constant" + if ins[position].type != "INT32": + return f"{ins[position].type} run-time indices; the runtime takes int32" 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] @@ -452,7 +465,8 @@ def _gather_run(op: _Operator, ins: list[_Tensor]): def _data_movement_reason(op: _Operator, ins: list[_Tensor], outs: list[_Tensor]) -> str: - data = [t for t in ins if not _is_constant(t)] + slot = _INDEX_SLOTS.get(op.kind) + data = [t for i, t in enumerate(ins) if not _is_constant(t) and i != slot] if not _same_quantization([*data, *outs]): return "inputs and outputs quantized differently" if op.kind == "STRIDED_SLICE": @@ -800,7 +814,8 @@ def convert(self, op: _Operator) -> None: fill = float((int(self.ints(ins[2])[0]) - source.zero_point[0]) * source.scale[0]) y = b.node("Pad", [self.value(ins[0]), pads, b.constant(np.float32(fill), tag + "_fill")], tag, mode="constant") - elif kind == "GATHER" and _gather_run(op, [self.tensors[i] for i in ins]) is None: + elif kind == "GATHER" and (not _is_constant(self.tensors[ins[1]]) + or _gather_run(op, [self.tensors[i] for i in ins]) is None): y = self._gather("Gather", ins[0], ins[1], op.option(0, "i"), op, tag) elif kind == "EMBEDDING_LOOKUP": y = self._gather("Gather", ins[1], ins[0], 0, op, tag) @@ -826,6 +841,13 @@ def convert(self, op: _Operator) -> None: axes = b.constant(np.asarray(self.ints(ins[1]), np.int64), tag + "_axes") y = self.tflite_order_out(self.custom("ReverseV2", [x, axes], tag + "_reversed", out.shape), outs[0], tag) + elif kind == "DYNAMIC_UPDATE_SLICE" and not _is_constant(self.tensors[ins[2]]): + # Run-time starts are in TFLite's axis order, so the update runs there. + rank = len(out.shape) + x = self.last_to_last(self.value(ins[0]), rank, tag + "_last", ins[0], True) + u = self.last_to_last(self.value(ins[1]), rank, tag + "_update", ins[1], True) + y = self.tflite_order_out(self.custom("DynamicUpdateSlice", [x, u, self.value(ins[2])], + tag + "_updated", out.shape), outs[0], tag) elif kind == "DYNAMIC_UPDATE_SLICE": rank = len(out.shape) starts = [self.ints(ins[2])[a] for a in _to_first(rank)] @@ -1061,7 +1083,8 @@ def _gather(self, kind: str, data: int, ids: int, axis, op: _Operator, tag: str) TFLite's axis order, where its axis and index tuples are stated.""" rank = len(self.tensors[data].shape) x = self.last_to_last(self.value(data), rank, tag + "_last", data, True) - indices = self.b.constant(np.asarray(self.tensors[ids].array(), np.int64), tag + "_indices") + indices = (self.b.constant(np.asarray(self.tensors[ids].array(), np.int64), tag + "_indices") + if _is_constant(self.tensors[ids]) else self.value(ids)) attributes = {} if axis is None else {"axis": axis % rank} batch = op.option(1, "i") if op.kind == "GATHER" else 0 if batch: @@ -1266,7 +1289,7 @@ def to_onnx(data: bytes, name: str) -> onnx.ModelProto: b._names.add(tensor.name) onnx_inputs.append(helper.make_tensor_value_info( tensor.name, boundary_type[tensor.type], _onnx_shape(tensor.shape))) - if tensor.type in ("FLOAT32", "BOOL"): + if tensor.type in ("FLOAT32", "BOOL", "INT32"): converter.values[index] = tensor.name else: scale = b.constant(np.float32(tensor.scale[0]), tensor.name + "_scale") diff --git a/tests/fixtures/tflite/ops/arg_max_gather.npz b/tests/fixtures/tflite/ops/arg_max_gather.npz new file mode 100644 index 0000000..55d139b Binary files /dev/null and b/tests/fixtures/tflite/ops/arg_max_gather.npz differ diff --git a/tests/fixtures/tflite/ops/arg_max_gather.tflite b/tests/fixtures/tflite/ops/arg_max_gather.tflite new file mode 100644 index 0000000..7fc2360 Binary files /dev/null and b/tests/fixtures/tflite/ops/arg_max_gather.tflite differ diff --git a/tests/fixtures/tflite/ops/dynamic_update_slice_runtime.npz b/tests/fixtures/tflite/ops/dynamic_update_slice_runtime.npz new file mode 100644 index 0000000..97fd44d Binary files /dev/null and b/tests/fixtures/tflite/ops/dynamic_update_slice_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/dynamic_update_slice_runtime.tflite b/tests/fixtures/tflite/ops/dynamic_update_slice_runtime.tflite new file mode 100644 index 0000000..f39018c Binary files /dev/null and b/tests/fixtures/tflite/ops/dynamic_update_slice_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/embedding_lookup_runtime.npz b/tests/fixtures/tflite/ops/embedding_lookup_runtime.npz new file mode 100644 index 0000000..1953a95 Binary files /dev/null and b/tests/fixtures/tflite/ops/embedding_lookup_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/embedding_lookup_runtime.tflite b/tests/fixtures/tflite/ops/embedding_lookup_runtime.tflite new file mode 100644 index 0000000..185413e Binary files /dev/null and b/tests/fixtures/tflite/ops/embedding_lookup_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/float_arg_max_gather.npz b/tests/fixtures/tflite/ops/float_arg_max_gather.npz new file mode 100644 index 0000000..e557f5d Binary files /dev/null and b/tests/fixtures/tflite/ops/float_arg_max_gather.npz differ diff --git a/tests/fixtures/tflite/ops/float_arg_max_gather.tflite b/tests/fixtures/tflite/ops/float_arg_max_gather.tflite new file mode 100644 index 0000000..c74ed56 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_arg_max_gather.tflite differ diff --git a/tests/fixtures/tflite/ops/float_dynamic_update_slice_runtime.npz b/tests/fixtures/tflite/ops/float_dynamic_update_slice_runtime.npz new file mode 100644 index 0000000..90e2221 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_dynamic_update_slice_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/float_dynamic_update_slice_runtime.tflite b/tests/fixtures/tflite/ops/float_dynamic_update_slice_runtime.tflite new file mode 100644 index 0000000..2d66a24 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_dynamic_update_slice_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/float_embedding_lookup_runtime.npz b/tests/fixtures/tflite/ops/float_embedding_lookup_runtime.npz new file mode 100644 index 0000000..eef7bd5 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_embedding_lookup_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/float_embedding_lookup_runtime.tflite b/tests/fixtures/tflite/ops/float_embedding_lookup_runtime.tflite new file mode 100644 index 0000000..16287e9 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_embedding_lookup_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/float_gather_nd_runtime.npz b/tests/fixtures/tflite/ops/float_gather_nd_runtime.npz new file mode 100644 index 0000000..3a06f12 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_gather_nd_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/float_gather_nd_runtime.tflite b/tests/fixtures/tflite/ops/float_gather_nd_runtime.tflite new file mode 100644 index 0000000..085228d Binary files /dev/null and b/tests/fixtures/tflite/ops/float_gather_nd_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/float_gather_runtime.npz b/tests/fixtures/tflite/ops/float_gather_runtime.npz new file mode 100644 index 0000000..d2e59ff Binary files /dev/null and b/tests/fixtures/tflite/ops/float_gather_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/float_gather_runtime.tflite b/tests/fixtures/tflite/ops/float_gather_runtime.tflite new file mode 100644 index 0000000..3622403 Binary files /dev/null and b/tests/fixtures/tflite/ops/float_gather_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/gather_nd_runtime.npz b/tests/fixtures/tflite/ops/gather_nd_runtime.npz new file mode 100644 index 0000000..92b5a94 Binary files /dev/null and b/tests/fixtures/tflite/ops/gather_nd_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/gather_nd_runtime.tflite b/tests/fixtures/tflite/ops/gather_nd_runtime.tflite new file mode 100644 index 0000000..5970fc3 Binary files /dev/null and b/tests/fixtures/tflite/ops/gather_nd_runtime.tflite differ diff --git a/tests/fixtures/tflite/ops/gather_runtime.npz b/tests/fixtures/tflite/ops/gather_runtime.npz new file mode 100644 index 0000000..80e3124 Binary files /dev/null and b/tests/fixtures/tflite/ops/gather_runtime.npz differ diff --git a/tests/fixtures/tflite/ops/gather_runtime.tflite b/tests/fixtures/tflite/ops/gather_runtime.tflite new file mode 100644 index 0000000..6eb8b9f Binary files /dev/null and b/tests/fixtures/tflite/ops/gather_runtime.tflite differ diff --git a/tests/test_frontend_tflite.py b/tests/test_frontend_tflite.py index 5b62e19..287491b 100644 --- a/tests/test_frontend_tflite.py +++ b/tests/test_frontend_tflite.py @@ -253,6 +253,33 @@ def test_a_reduction_over_axes_that_are_not_adjacent_is_refused(): assert any("axes that are not adjacent" in reason for reason in tflite.unsupported(bytes(data))) +def _read_edited(monkeypatch, name, edit): + """The fixture `name` as the frontend reads it, after `edit(tensors, operators)`.""" + data = (FIXTURES / "ops" / f"{name}.tflite").read_bytes() + model, graphs = tflite._read(data) + edit(graphs[0][1], graphs[0][4]) + monkeypatch.setattr(tflite, "_read", lambda _: (model, graphs)) + return tflite.unsupported(data) + + +def test_run_time_indices_come_from_a_model_input_or_an_arg_max(monkeypatch): + """Indices stay positions: an index tensor feeds only index operands, and + indices computed at run time come from ARG_MAX or ARG_MIN.""" + def swap(tensors, operators): + gather = next(op for op in operators if op.kind == "GATHER") + gather.inputs = [gather.inputs[1], gather.inputs[0]] + reasons = _read_edited(monkeypatch, "float_arg_max_gather", swap) + assert any("feeds an operand other than indices" in r for r in reasons) + assert any("other than ARG_MAX or ARG_MIN" in r for r in reasons) + + +def test_int64_run_time_indices_are_refused(monkeypatch): + def widen(tensors, operators): + tensors[operators[0].inputs[1]].type = "INT64" + reasons = _read_edited(monkeypatch, "float_gather_runtime", widen) + assert any("INT64 run-time indices; the runtime takes int32" 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"]