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
40 changes: 33 additions & 7 deletions scripts/gen_tflite_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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",
Expand All @@ -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))
Expand All @@ -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:
Expand All @@ -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)
Expand All @@ -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.
Expand Down
45 changes: 34 additions & 11 deletions src/tigris/frontends/tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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"}
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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":
Expand Down Expand Up @@ -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)
Expand All @@ -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)]
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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")
Expand Down
Binary file added tests/fixtures/tflite/ops/arg_max_gather.npz
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/arg_max_gather.tflite
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/float_arg_max_gather.npz
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/gather_nd_runtime.npz
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/gather_nd_runtime.tflite
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/gather_runtime.npz
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/gather_runtime.tflite
Binary file not shown.
27 changes: 27 additions & 0 deletions tests/test_frontend_tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down
Loading