Skip to content
Open
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: 9 additions & 0 deletions python/tvm/ir/type.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 1 addition & 6 deletions python/tvm/relax/transform/legalize_ops/qdq.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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":
Expand Down
47 changes: 47 additions & 0 deletions tests/python/relax/test_transform_legalize_ops_qdq.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()