Skip to content
Merged
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
36 changes: 32 additions & 4 deletions src/tirx/transform/ir_utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/scope_stack.h>
#include <tvm/s_tir/stmt.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/layout.h>
#include <tvm/tirx/stmt_functor.h>
#include <tvm/tirx/transform.h>
Expand Down Expand Up @@ -526,6 +527,33 @@ class IRConvertSSA final : public StmtExprMutator {
Var new_var;
};

/*! \brief Check whether a buffer uses a variable in any remapped field. */
static bool BufferDependsOnVar(const Buffer& buffer, const VarNode* var) {
if (buffer->data.get() == var) return true;

auto uses_var = [var](const PrimExpr& expr) {
return expr.defined() && UsesVar(expr, [var](const VarNode* node) { return node == var; });
};
if (uses_var(buffer->elem_offset)) return true;
for (const PrimExpr& dim : buffer->shape) {
if (uses_var(dim)) return true;
}
for (const PrimExpr& stride : buffer->strides) {
if (uses_var(stride)) return true;
}
if (buffer->layout.has_value()) {
if (const auto* tile_layout = buffer->layout.value().as<TileLayoutNode>()) {
for (const Iter& iter : tile_layout->shard) {
if (uses_var(iter->extent) || uses_var(iter->stride)) return true;
}
for (const Iter& iter : tile_layout->replica) {
if (uses_var(iter->extent) || uses_var(iter->stride)) return true;
}
}
}
return false;
}

/*! \brief Create a new variable with the same name and type as the original. */
static Var MakeNewVar(const Var& old_var) { return Var(old_var->name, old_var->ty); }

Expand All @@ -542,7 +570,7 @@ class IRConvertSSA final : public StmtExprMutator {
var_remap_[old_var.get()].pop_back();
for (auto& kv : buf_remap_) {
std::vector<Buffer>& buffers = kv.second;
if (buffers.size() && (buffers.back()->data.get() == new_var.get())) {
if (buffers.size() && BufferDependsOnVar(buffers.back(), new_var.get())) {
buffers.pop_back();
}
}
Expand All @@ -561,7 +589,7 @@ class IRConvertSSA final : public StmtExprMutator {
var_remap_[remap.old_var.get()].pop_back();
for (auto& kv : buf_remap_) {
std::vector<Buffer>& buffers = kv.second;
if (buffers.size() && (buffers.back()->data.get() == remap.new_var.get())) {
if (buffers.size() && BufferDependsOnVar(buffers.back(), remap.new_var.get())) {
buffers.pop_back();
}
}
Expand Down Expand Up @@ -598,7 +626,7 @@ class IRConvertSSA final : public StmtExprMutator {
parent->var_remap_[remap.old_var.get()].pop_back();
for (auto& kv : parent->buf_remap_) {
std::vector<Buffer>& buffers = kv.second;
if (buffers.size() && (buffers.back()->data.get() == remap.new_var.get())) {
if (buffers.size() && BufferDependsOnVar(buffers.back(), remap.new_var.get())) {
buffers.pop_back();
}
}
Expand All @@ -622,7 +650,7 @@ class IRConvertSSA final : public StmtExprMutator {
parent->var_remap_[remap.old_var.get()].pop_back();
for (auto& kv : parent->buf_remap_) {
std::vector<Buffer>& buffers = kv.second;
if (buffers.size() && (buffers.back()->data.get() == remap.new_var.get())) {
if (buffers.size() && BufferDependsOnVar(buffers.back(), remap.new_var.get())) {
buffers.pop_back();
}
}
Expand Down
24 changes: 24 additions & 0 deletions tests/python/tirx-transform/test_tir_transform_convert_ssa.py
Original file line number Diff line number Diff line change
Expand Up @@ -535,5 +535,29 @@ def test_shared_shape_var_in_buffer_map_and_alloc_buffer():
tvm.ir.assert_structural_equal(after["main"], before)


def test_reused_loop_var_in_decl_buffer_elem_offset():
"""Remap a buffer whose elem_offset depends on an SSA-renamed loop var."""
loop_var = tirx.Var("loop_var", "int32")
buffer = tirx.decl_buffer(
(128,),
"float32",
"buffer",
elem_offset=loop_var * 128,
scope="shared.dyn",
)
loop = tirx.For(
loop_var,
0,
128,
tirx.ForKind.SERIAL,
tirx.DeclBuffer(buffer, tirx.Evaluate(tirx.BufferLoad(buffer, [0]))),
)
func = tirx.PrimFunc([buffer.data], tirx.SeqStmt([loop, loop, loop]))

after = tvm.tirx.transform.ConvertSSA()(tvm.IRModule.from_expr(func))

tvm.tirx.analysis.verify_well_formed(after["main"], assert_mode=True)


if __name__ == "__main__":
tvm.testing.main()
Loading