From 85156cc9a77cab0edd79011219047a9e468cb9ca Mon Sep 17 00:00:00 2001 From: Andrew Fish Date: Wed, 30 Sep 2026 11:48:32 -0700 Subject: [PATCH 1/2] [Fix][Relax] Dequantize fp8 as a float, not an integer dequantize_compute chose its intermediate dtype with matches_code(FLOAT, BFLOAT). Those are DLPack codes 2 and 4, while every narrow float carries its own code instead -- Float8E3M4 at 7 through Float4E2M1FN at 17. An fp8 input therefore took the int32 branch and was truncated to an integer before the zero point was subtracted and the scale applied, so a stored -0.9375 became 0. This is reachable through the public API: relax.quantize already accepts float8_e4m3fn and float8_e5m2 as output dtypes. It stays hidden wherever the stored value is already integral, which at E4M3's spacing covers most of the top of the range. Ask the dtype whether it is a float instead of enumerating codes, via a PrimType.is_float that forwards to the one tvm_ffi already provides. The float code list keeps a single home, so a future narrow float needs no edit here. quantize is unaffected -- it decides the same question with "int" in out_dtype, which excludes the fp8 dtypes correctly. --- python/tvm/ir/type.py | 9 +++++++++ python/tvm/relax/transform/legalize_ops/qdq.py | 7 +------ 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/python/tvm/ir/type.py b/python/tvm/ir/type.py index e0fbf6cd099a..daeb6f34b862 100644 --- a/python/tvm/ir/type.py +++ b/python/tvm/ir/type.py @@ -98,6 +98,15 @@ def __hash__(self): def __str__(self): return str(self.dtype) + @property + def is_float(self) -> bool: + """Return whether this type stores floating-point values. + + True for every float width, including the narrow ones, which each carry + their own DLPack dtype code. + """ + return self.dtype.is_float + def matches_code(self, *codes) -> bool: """Return whether this type has any of the given DLPack dtype codes.""" type_code = self.dtype.type_code diff --git a/python/tvm/relax/transform/legalize_ops/qdq.py b/python/tvm/relax/transform/legalize_ops/qdq.py index 9cea92314681..03a1abf9fef3 100644 --- a/python/tvm/relax/transform/legalize_ops/qdq.py +++ b/python/tvm/relax/transform/legalize_ops/qdq.py @@ -20,7 +20,6 @@ import tvm from tvm import te, tirx from tvm.ir import Call -from tvm.runtime import DataTypeCode from ...block_builder import BlockBuilder from ...expr import Expr @@ -142,11 +141,7 @@ def dequantize_compute(*indices): zp_value = zp[(0,) * len(zp.shape)] else: zp_value = zp[indices[axis]] - dtype = ( - "float32" - if data.dtype.matches_code(DataTypeCode.FLOAT, DataTypeCode.BFLOAT) - else "int32" - ) + dtype = "float32" if data.dtype.is_float else "int32" sub = data[indices].astype(dtype) - zp_value out = sub * scale_value.astype("float32") if out_dtype == "float32": From 7ce4651f496a71f79cc4372863d722ecdf453733 Mon Sep 17 00:00:00 2001 From: Andrew Fish Date: Wed, 30 Sep 2026 11:48:37 -0700 Subject: [PATCH 2/2] [Relax] Test fp8 dequantize legalization Every dequantize case in this file starts from int8, and the fp8 coverage in test_op_qdq.py stops at type inference, so nothing checked the TIR generated for an fp8 input. --- .../relax/test_transform_legalize_ops_qdq.py | 47 +++++++++++++++++++ 1 file changed, 47 insertions(+) diff --git a/tests/python/relax/test_transform_legalize_ops_qdq.py b/tests/python/relax/test_transform_legalize_ops_qdq.py index 86d766afc8c0..c1078bb5916b 100644 --- a/tests/python/relax/test_transform_legalize_ops_qdq.py +++ b/tests/python/relax/test_transform_legalize_ops_qdq.py @@ -580,5 +580,52 @@ def main(data: R.Tensor((2, 4), dtype="int8")) -> R.Tensor((2, 4), dtype="float1 tvm.ir.assert_structural_equal(mod, Expected) +def test_dequantize_float8_e4m3fn_to_fp32(): + @tvm.script.ir_module + class Dequantize: + @R.function + def main( + data: R.Tensor((2, 4), "float8_e4m3fn"), + scale: R.Tensor((2,), "float32"), + zp: R.Tensor((2,), "float16"), + ) -> R.Tensor((2, 4), "float32"): + out = R.dequantize(data, scale, zp, axis=0, out_dtype="float32") + return out + + @tvm.script.ir_module + class Expected: + @Ts.prim_func(private=True) + def dequantize( + A: T.Buffer((T.int64(2), T.int64(4)), "float8_e4m3fn"), + B: T.Buffer((T.int64(2),), "float32"), + C: T.Buffer((T.int64(2),), "float16"), + dequantized: T.Buffer((T.int64(2), T.int64(4)), "float32"), + ): + T.func_attr({"tirx.noalias": True}) + # with Ts.sblock("root"): + for i0, i1 in T.grid(T.int64(2), T.int64(4)): + with Ts.sblock("dequantized"): + v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) + Ts.reads(A[v_i0, v_i1], C[v_i0], B[v_i0]) + Ts.writes(dequantized[v_i0, v_i1]) + dequantized[v_i0, v_i1] = ( + T.Cast("float32", A[v_i0, v_i1]) - T.Cast("float32", C[v_i0]) + ) * B[v_i0] + + @R.function + def main( + data: R.Tensor((2, 4), dtype="float8_e4m3fn"), + scale: R.Tensor((2,), dtype="float32"), + zp: R.Tensor((2,), dtype="float16"), + ) -> R.Tensor((2, 4), dtype="float32"): + out = R.call_tir( + Expected.dequantize, (data, scale, zp), out_ty=R.Tensor((2, 4), dtype="float32") + ) + return out + + mod = LegalizeOps()(Dequantize) + tvm.ir.assert_structural_equal(mod, Expected) + + if __name__ == "__main__": tvm.testing.main()