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": 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()