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
8 changes: 4 additions & 4 deletions python/tvm/tirx/stmt_functor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down
43 changes: 43 additions & 0 deletions tests/python/tirx/transform/test_stmt_functor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading