diff --git a/python/tvm/tirx/stmt_functor.py b/python/tvm/tirx/stmt_functor.py index db418f18b6ca..1f9755b22a38 100644 --- a/python/tvm/tirx/stmt_functor.py +++ b/python/tvm/tirx/stmt_functor.py @@ -363,14 +363,14 @@ def visit_scope_id_def_stmt_(self, op): def visit_op_call_(self, op): """Visitor implementation for TilePrimitiveCall.""" for arg in op.args: - if tvm.ir.is_prim_expr(arg): + if isinstance(arg, tvm.ir.Expr): self.visit_expr(arg) elif isinstance(arg, tvm.tirx.Stmt): self.visit_stmt(arg) elif isinstance(arg, tvm.tirx.BufferRegion): self.visit_buffer_region_(arg) for value in op.config.values(): - if tvm.ir.is_prim_expr(value): + if isinstance(value, tvm.ir.Expr): self.visit_expr(value) elif isinstance(value, tvm.tirx.Stmt): self.visit_stmt(value) @@ -842,7 +842,7 @@ def visit_op_call_(self, op): args_changed = False for arg in op.args: - if tvm.ir.is_prim_expr(arg): + if isinstance(arg, tvm.ir.Expr): new_arg = self.visit_expr(arg) elif isinstance(arg, tvm.tirx.Stmt): new_arg = self.visit_stmt(arg) @@ -859,7 +859,7 @@ def visit_op_call_(self, op): new_config = {} config_changed = False for key, value in op.config.items(): - if tvm.ir.is_prim_expr(value): + if isinstance(value, tvm.ir.Expr): new_value = self.visit_expr(value) elif isinstance(value, tvm.tirx.Stmt): new_value = self.visit_stmt(value) diff --git a/tests/python/tirx/transform/test_stmt_functor.py b/tests/python/tirx/transform/test_stmt_functor.py index 8d602cce3560..fc7deb29e13b 100644 --- a/tests/python/tirx/transform/test_stmt_functor.py +++ b/tests/python/tirx/transform/test_stmt_functor.py @@ -1183,6 +1183,49 @@ def op_call_with_config(A: T.Buffer((10,), "int32"), B: T.Buffer((10,), "int32") ) +def test_op_call_pointer_config_visited_and_mutated(): + """Pointer-valued config expressions participate in Python traversal.""" + + @T.prim_func + def copy_async( + A: T.Buffer((8,), "float16"), + B: T.Buffer((8,), "float16"), + mbar: T.Buffer((1,), "uint64"), + ): + Tx.copy_async(B[:], A[:], dispatch="tma_auto", mbar=T.address_of(mbar[0])) + + op_call = copy_async.body + assert isinstance(op_call, tir.TilePrimitiveCall) + mbar_buffer = copy_async.buffer_map[copy_async.params[2]] + + class LoadCollector(StmtExprVisitor): + def __init__(self): + super().__init__() + self.buffers = [] + + def visit_buffer_load_(self, op): + self.buffers.append(op.buffer) + return super().visit_buffer_load_(op) + + collector = LoadCollector() + collector.visit_stmt(op_call) + assert any(buffer.same_as(mbar_buffer) for buffer in collector.buffers) + + replacement = tir.decl_buffer((1,), "uint64", name="replacement") + + class ReplaceMbarLoad(StmtExprMutator): + def visit_buffer_load_(self, op): + new_op = super().visit_buffer_load_(op) + if op.buffer.same_as(mbar_buffer): + return tir.BufferLoad(replacement, new_op.indices, new_op.predicate) + return new_op + + updated = ReplaceMbarLoad().visit_stmt(op_call) + mbar_load = updated.config["mbar"].args[0] + assert isinstance(mbar_load, tir.BufferLoad) + assert mbar_load.buffer.same_as(replacement) + + def test_op_call_nested_config_visited_and_substituted(): """Nested selector arrays participate in the core visitor and mutator.""" from tvm.tirx.stmt_functor import post_order_visit, substitute