From ecb5698e3a1b9405092a215a4fb24fee6e125f58 Mon Sep 17 00:00:00 2001 From: Anders Lie Date: Thu, 10 Sep 2026 22:15:30 -0700 Subject: [PATCH] [BugFix][TIRx] Legalize BF16 conversions inside Let and Bind values ComputeLegalizer promoted the value of a Let expression or a Bind statement without visiting it, so a conversion from a BF16 buffer bound to a variable survived the pass. The paired storage pass then retyped the buffer to uint16 and the surviving cast became an integer-to-float conversion of the raw bits. Every other handler in the class visits its operands before promoting; Let and Bind now do the same. Tests cover a Bind statement (a typed assignment in script) and a Let expression whose value is a BF16 load cast to float32. --- .../transform/unsupported_dtype_legalize.cc | 4 +- .../test_tir_transform_bf16_legalize.py | 68 +++++++++++++++++++ 2 files changed, 70 insertions(+), 2 deletions(-) diff --git a/src/tirx/transform/unsupported_dtype_legalize.cc b/src/tirx/transform/unsupported_dtype_legalize.cc index 51196904c88a..69a41550ac56 100644 --- a/src/tirx/transform/unsupported_dtype_legalize.cc +++ b/src/tirx/transform/unsupported_dtype_legalize.cc @@ -268,7 +268,7 @@ class ComputeLegalizer : public StmtExprMutator { } PrimExpr VisitExpr_(const LetNode* op) final { - PrimExpr value = PromoteToTarget(op->value); + PrimExpr value = PromoteToTarget(this->VisitExpr(op->value)); Var var = op->var; if (value.dtype() != op->value.dtype()) { var = op->var.copy_with_dtype(op->value.dtype()); @@ -298,7 +298,7 @@ class ComputeLegalizer : public StmtExprMutator { DEFINE_BIOP_EXPR_LEGALIZE(NENode, operator!=); Stmt VisitStmt_(const BindNode* op) final { - PrimExpr value = PromoteToTarget(op->value); + PrimExpr value = PromoteToTarget(this->VisitExpr(op->value)); Var var = op->var; if (value.dtype() != op->value.dtype()) { var = op->var.copy_with_dtype(op->value.dtype()); diff --git a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py index fdaa51622b6b..e8837734d82c 100644 --- a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py +++ b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py @@ -448,6 +448,74 @@ def main( tvm.ir.assert_structural_equal(after_storage, BindTarget(target)(after_storage_legalize())) +def test_bf16_bind_value_will_legalize(): + """A conversion inside a bound value is legalized like any other.""" + + def get_before(): + @tvm.script.ir_module + class Before: + @T.prim_func + def main(Aptr: T.handle("bfloat16"), Cptr: T.handle("float32")): + T.func_attr({"global_symbol": "main"}) + A = T.decl_buffer((100,), "bfloat16", data=Aptr) + C = T.decl_buffer((100,), "float32", data=Cptr) + for i in T.grid(100): + b: T.float32 = T.Cast("float32", A[i]) + C[i] = b * T.float32(2) + + return Before + + def after_compute_legalize(): + @tvm.script.ir_module + class After: + @T.prim_func + def main(Aptr: T.handle("bfloat16"), Cptr: T.handle("float32")): + T.func_attr({"global_symbol": "main"}) + A = T.decl_buffer((100,), "bfloat16", data=Aptr) + C = T.decl_buffer((100,), "float32", data=Cptr) + for i in T.grid(100): + b: T.float32 = bf16tof32(A[i]) + C[i] = b * T.float32(2) + + return After + + target = Target("nvidia/geforce-rtx-2080-ti") + before = BindTarget(target)(get_before()) + after_compute = tvm.tirx.transform.BF16ComputeLegalize()(before) + tvm.ir.assert_structural_equal(after_compute, BindTarget(target)(after_compute_legalize())) + + +def test_bf16_let_value_will_legalize(): + """A conversion inside a Let expression's value is legalized like any other.""" + A = tvm.tirx.decl_buffer((100,), "bfloat16", name="A") + C = tvm.tirx.decl_buffer((100,), "float32", name="C") + i = tvm.tirx.Var("i", "int32") + x = tvm.tirx.Var("x", "float32") + value = tvm.tirx.Cast("float32", tvm.tirx.BufferLoad(A, [i])) + store = tvm.tirx.BufferStore(C, tvm.tirx.Let(x, value, x * tvm.tirx.const(2.0, "float32")), [i]) + loop = tvm.tirx.For(i, 0, 100, tvm.tirx.ForKind.SERIAL, store) + body = tvm.tirx.SeqStmt([tvm.tirx.DeclBuffer(A), tvm.tirx.DeclBuffer(C), loop]) + func = tvm.tirx.PrimFunc([A.data, C.data], body).with_attr("global_symbol", "main") + target = Target("nvidia/geforce-rtx-2080-ti") + before = BindTarget(target)(tvm.IRModule({"main": func})) + after_compute = tvm.tirx.transform.BF16ComputeLegalize()(before) + + lets, bf16_casts = [], [] + + def visit(node): + if isinstance(node, tvm.tirx.Let): + lets.append(node) + if isinstance(node, tvm.tirx.Cast) and "bfloat16" in ( + str(node.dtype), + str(node.value.dtype), + ): + bf16_casts.append(node) + + tvm.tirx.stmt_functor.post_order_visit(after_compute["main"].body, visit) + assert len(lets) == 1 + assert bf16_casts == [] + + if __name__ == "__main__": test_bf16_storage_compute_scope_will_legalize() test_bf16_storage_compute_scope_wont_legalize()