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
9 changes: 8 additions & 1 deletion scripts/gen_tflite_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,8 @@ def _binary(fn, a, b):
"div": [(-3.0, 3.0), _DIVISOR], "float_div": [(-3.0, 3.0), _DIVISOR], "float_floor_div": [(-3.0, 3.0), _DIVISOR],
"float_floor_mod": [(-3.0, 3.0), _DIVISOR], "float_div_constant_first": [_DIVISOR],
"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)],
}


Expand Down Expand Up @@ -252,6 +254,9 @@ def rewrite(model: bytes) -> bytes:
"cumsum": _unary(lambda x: tf.math.cumsum(x, 2), _MAP),
"cumsum_exclusive_reverse": _unary(
lambda x: tf.math.cumsum(x, -1, exclusive=True, reverse=True), _MAP),
"cumsum_offset": _unary(lambda x: tf.math.cumsum(x, 1), _MAP),
"cumsum_offset_exclusive_reverse": _unary(
lambda x: tf.math.cumsum(x, 2, exclusive=True, reverse=True), _MAP),
}
CASES.update(ACTIVATIONS)
# TFLite Micro runs L2_POOL_2D in float only; the converter emits it for no
Expand Down Expand Up @@ -425,7 +430,9 @@ 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),
"not_equal": _int8_island(), "add_n": _int8_island()})
"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):
Expand Down
2 changes: 1 addition & 1 deletion src/tigris/analysis/validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -319,7 +319,7 @@ def validate_operator_support(ag: AnalyzedGraph) -> OperatorSupportValidation:
or quants[0].zero_point[0] != quants[1].zero_point[0]):
reasons.append("ReduceMax/ReduceMin requires identical input and output quantization")
elif op.op_type == "CumSum":
if int(quants[0].zero_point[0]) != 0:
if int(quants[0].zero_point[0]) != 0 and not op.attrs.get("tflite_seeded"):
reasons.append("int8 CumSum requires input zero point 0 to preserve ONNX prefix sums")
if float(quants[0].scale[0]) / float(quants[1].scale[0]) >= 2**19:
reasons.append("int8 CumSum output multiplier must be smaller than one")
Expand Down
17 changes: 10 additions & 7 deletions src/tigris/frontends/tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -343,10 +343,6 @@ def _operator_reason(op: _Operator, tensors: list[_Tensor]) -> str:
axes = sorted({a % len(ins[0].shape) for a in ins[1].array().reshape(-1).tolist()})
if not axes or axes != list(range(axes[0], axes[-1] + 1)):
return "axes that are not adjacent"
if op.kind == "CUMSUM" and outs[0].type == "INT8" and ins[0].zero_point[0] != 0:
# TFLite Micro seeds the int8 sum with the input zero point, which
# ONNX's CumSum over dequantized values cannot express.
return "int8 input zero point other than 0"
if op.kind == "SOFTMAX" and op.option(0, "f", 1.0) != 1.0:
return "beta other than 1"
if op.kind == "FULLY_CONNECTED" and op.option(1, "b") != 0:
Expand Down Expand Up @@ -928,9 +924,16 @@ def _reduction(self, op: _Operator, tag: str) -> str:
tag + "_rows_q")
axis = _onnx_axis(1, 3)
if op.kind == "CUMSUM":
y = b.node("CumSum", [x, b.constant(np.asarray(axis, np.int64), tag + "_axis")],
tag + "_scan", exclusive=int(op.option(0, "?", False)),
reverse=int(op.option(1, "?", False)))
operands = [x, b.constant(np.asarray(axis, np.int64), tag + "_axis")]
options = {"exclusive": int(op.option(0, "?", False)),
"reverse": int(op.option(1, "?", False))}
if source.type == "INT8":
# TFLite seeds the int8 sum with the input zero point, which
# ONNX's CumSum over dequantized values does not; the
# compiler's own CumSum states it.
y = self.custom("CumSum", operands, tag + "_scan", _onnx_shape(rows), **options)
else:
y = b.node("CumSum", operands, tag + "_scan", **options)
kept = rows
elif op.kind == "SUM":
y = b.node("ReduceSum", [x, b.constant(np.asarray([axis], np.int64), tag + "_axes")],
Expand Down
12 changes: 12 additions & 0 deletions src/tigris/loaders/onnx/normalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
def normalize(ag: AnalyzedGraph) -> AnalyzedGraph:
"""Apply all normalization passes in sequence."""
declared_outputs = list(ag.model_outputs)
ag = _adopt_tflite_cumsum(ag)
ag = _normalize_arg_outputs(ag)
ag = _drop_inference_identities(ag)
ag = _lower_legacy_softmax(ag)
Expand Down Expand Up @@ -2873,6 +2874,17 @@ def constant(op, position):
return ag


def _adopt_tflite_cumsum(ag: AnalyzedGraph) -> AnalyzedGraph:
"""The compiler's own CumSum states TFLite's int8 semantics, which seed the
sum with the scaled input zero point; ONNX's CumSum sums dequantized
values, which differs whenever that zero point is not 0."""
for op in ag.ops:
if op.op_type == "tigris::CumSum":
op.op_type = "CumSum"
op.attrs["tflite_seeded"] = 1
return ag


def _normalize_arg_outputs(ag: AnalyzedGraph) -> AnalyzedGraph:
# An index cast to int32 as the model output is the index itself, stored
# as the runtime writes it.
Expand Down
7 changes: 4 additions & 3 deletions tests/fixtures/tflite/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,10 @@ RESHAPE into the SQUEEZE and EXPAND_DIMS operators, which the converter never
emits itself. The converter keeps int8 ELU, CUMSUM, DYNAMIC_UPDATE_SLICE,
NOT_EQUAL and ADD_N in float after DEQUANTIZEs, so `elu`, `cumsum`,
`cumsum_exclusive_reverse`, `dynamic_update_slice`, `not_equal` and `add_n`
are rewritten to run the int8 operator itself: the CUMSUM cases with the input
zero point at 0, and the update with the operand's quantization, since that
kernel copies raw bytes. `float_l2_pool` recodes an AVERAGE_POOL_2D as
are rewritten to run the int8 operator itself: `cumsum` and
`cumsum_exclusive_reverse` with the input zero point at 0, the `cumsum_offset`
cases with the converter's own zero point near -128, and the update with the
operand's quantization, since that kernel copies raw bytes. `float_l2_pool` recodes an AVERAGE_POOL_2D as
L2_POOL_2D, and the `embedding_lookup` cases recode a GATHER on axis 0 as
EMBEDDING_LOOKUP; the converter emits neither. `select_v2` recodes the
converter's SELECT, which TFLite Micro does not register, as the SELECT_V2 it
Expand Down
Binary file added tests/fixtures/tflite/ops/cumsum_offset.npz
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/cumsum_offset.tflite
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/float_cumsum_offset.npz
Binary file not shown.
Binary file added tests/fixtures/tflite/ops/float_cumsum_offset.tflite
Binary file not shown.
Binary file not shown.
Binary file not shown.
10 changes: 4 additions & 6 deletions tests/test_frontend_tflite.py
Original file line number Diff line number Diff line change
Expand Up @@ -212,9 +212,7 @@ 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 test_an_int8_cumsum_off_a_zero_input_zero_point_is_refused():
_, graphs = tflite._read((FIXTURES / "ops" / "cumsum.tflite").read_bytes())
_, tensors, _, _, operators = graphs[0]
assert tflite._operator_reason(operators[0], tensors) == ""
tensors[operators[0].inputs[0]].zero_point = np.array([3])
assert tflite._operator_reason(operators[0], tensors) == "int8 input zero point other than 0"
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"]
assert [node.domain for node in scans] == ["tigris"]
11 changes: 11 additions & 0 deletions tests/test_operator_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -980,6 +980,17 @@ def _reduction_graph(kind, *, axes=(1,), shape=(2, 5, 3), keep=True, quantized=F
return graph


def test_int8_cumsum_off_a_zero_input_zero_point_needs_tflite_semantics():
"""ONNX's CumSum sums dequantized values; only the TFLite form, which seeds
the sum with the scaled input zero point, runs off a zero point of 0."""
graph = _reduction_graph("CumSum", quantized=True)
graph.tensors["x"].quant = QuantParam(scale=np.array([0.125], np.float32),
zero_point=np.array([-128], np.int8))
assert "zero point 0" in validate_operator_support(graph).describe()
graph.ops[0].attrs["tflite_seeded"] = 1
assert validate_operator_support(graph).supported


@pytest.mark.parametrize("kind", ["ReduceMax", "ReduceMin", "ReduceSum", "CumSum"])
@pytest.mark.parametrize("axes,shape", [([], (2, 5, 3)), ([0, 1], (2, 5, 3)),
([3], (2, 5, 3)), ([1], (2, 5, 3, 4))])
Expand Down
Loading