From c08a5f100e69d92ab25406d18df2df758e02f8e7 Mon Sep 17 00:00:00 2001 From: Shinichiro Hamaji Date: Wed, 5 Apr 2023 15:48:27 +0900 Subject: [PATCH] [QNN] Convert fake quantized take to quantized op --- .../transform/fake_quantization_to_integer.py | 11 +++++++++++ .../test_pass_fake_quantization_to_integer.py | 14 ++++++++++++++ 2 files changed, 25 insertions(+) diff --git a/python/tvm/relay/transform/fake_quantization_to_integer.py b/python/tvm/relay/transform/fake_quantization_to_integer.py index 84b1f33e98cc..7375a4f3c0a0 100644 --- a/python/tvm/relay/transform/fake_quantization_to_integer.py +++ b/python/tvm/relay/transform/fake_quantization_to_integer.py @@ -622,3 +622,14 @@ def unary(expr, type_map): register_unary_qnn("tanh", relay.qnn.op.tanh) register_unary_qnn("abs", relay.qnn.op.abs) register_unary_qnn("log", relay.qnn.op.log) + + +@register_fake_quantization_to_integer("take") +def take(expr, type_map): + """Rewrite a take op""" + arg = expr.args[0] + indices = expr.args[1] + t = type_map[arg] + + out = relay.op.take(arg, indices, **expr.attrs) + return [out, t] diff --git a/tests/python/relay/test_pass_fake_quantization_to_integer.py b/tests/python/relay/test_pass_fake_quantization_to_integer.py index d384635e42e5..f349e0979395 100644 --- a/tests/python/relay/test_pass_fake_quantization_to_integer.py +++ b/tests/python/relay/test_pass_fake_quantization_to_integer.py @@ -1100,5 +1100,19 @@ def test_fq_qat_intermediate_infertype(): compare_expected_fq_qat_to_int(expr, expected_expr, [x_np]) +def test_fake_quantize_take(): + x = relay.var("x", shape=[33, 11], dtype="int8") + indices_np = np.random.randint(0, 33, size=[37], dtype="int32") + indices = relay.const(indices_np) + + x = relay.qnn.op.dequantize(x, relay.const(2.0), relay.const(114)) + op = relay.op.take(x, indices, axis=0) + op = relay.qnn.op.quantize(op, relay.const(2.0), relay.const(114), out_dtype="uint8") + + x_np = np.random.randint(-25, 25, size=[33, 11], dtype="int8") + + compare_fq_to_int(op, [x_np]) + + if __name__ == "__main__": tvm.testing.main()