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
14 changes: 13 additions & 1 deletion src/relax/transform/static_plan_block_memory.cc
Original file line number Diff line number Diff line change
Expand Up @@ -935,6 +935,18 @@ class StorageAllocationRewriter : public ExprMutator {
private:
using ExprMutator::VisitExpr_;

Expr VisitExpr_(const SeqExprNode* seq) final {
// A storage var is only visible in the scope it is emitted in, such as an if branch.
// Forget the vars emitted in this scope on exit, so that a token first used inside a
// branch gets a new `alloc_storage` where it is reused in another branch or after the if.
// A token shared by several scopes is allocated in each at its final size, which
// `RequestReuse` may have enlarged.
auto saved_token2storage_var = token2storage_var_;
Expr ret = ExprMutator::VisitExpr_(seq);
token2storage_var_ = std::move(saved_token2storage_var);
return ret;
}

Expr VisitExpr_(const CallNode* call) final {
static const Op alloc_tensor_op = Op::Get("relax.builtin.alloc_tensor");
static const Op mem_alloc_storage = Op::Get("relax.memory.alloc_storage");
Expand Down Expand Up @@ -1030,7 +1042,7 @@ class StorageAllocationRewriter : public ExprMutator {
std::unordered_map<const ExprNode*, StorageToken> alloc_tensor2token_;
/*! \brief The mapping from each binding block to the storage tokens that are create inside. */
std::unordered_map<const BindingBlockNode*, std::vector<const StorageTokenNode*>> block2tokens_;
/*! \brief The mapping from each token to its corresponding storage var in each function. */
/*! \brief The mapping from each token to its storage var visible in the current scope. */
std::unordered_map<const StorageTokenNode*, Var> token2storage_var_;
};

Expand Down
109 changes: 109 additions & 0 deletions tests/python/relax/test_transform_static_plan_block_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -1877,5 +1877,114 @@ def main(x: R.Tensor((10,), dtype="float32")) -> R.Tensor((10,), dtype="float32"
tvm.ir.assert_structural_equal(after, Expected)


def _count_alloc_storage(func):
alloc_storage_op = tvm.ir.Op.get("relax.memory.alloc_storage")
count = 0

def visit(expr):
nonlocal count
if isinstance(expr, relax.Call) and expr.op.same_as(alloc_storage_op):
count += 1

relax.analysis.post_order_visit(func, visit)
return count


def test_if_branches_do_not_share_storage_var():
@I.ir_module
class Before:
@Ts.prim_func
def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")):
T.evaluate(0)

@R.function
def main(
cond: R.Tensor((), dtype="bool"), x: R.Tensor((2, 3), dtype="float32")
) -> R.Tensor((2, 3), dtype="float32"):
R.func_attr({"relax.force_pure": True})
cls = Before
if cond:
alloc = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(x, alloc)
out = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(alloc, out)
z = out
else:
alloc1 = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(x, alloc1)
out1 = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(alloc1, out1)
z = out1
return z

after = relax.transform.StaticPlanBlockMemory()(Before)
assert relax.analysis.check_well_formed(after)
assert _count_alloc_storage(after["main"]) == 2


def test_if_branch_storage_not_reused_after_if():
@I.ir_module
class Before:
@Ts.prim_func
def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")):
T.evaluate(0)

@R.function
def main(
cond: R.Tensor((), dtype="bool"), x: R.Tensor((2, 3), dtype="float32")
) -> R.Tensor((2, 3), dtype="float32"):
R.func_attr({"relax.force_pure": True})
cls = Before
if cond:
z = x
else:
alloc = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(x, alloc)
out = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(alloc, out)
z = out
alloc1 = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(z, alloc1)
out1 = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(alloc1, out1)
return out1

after = relax.transform.StaticPlanBlockMemory()(Before)
assert relax.analysis.check_well_formed(after)
assert _count_alloc_storage(after["main"]) == 2


def test_if_branches_share_storage_allocated_before_if():
@I.ir_module
class Before:
@Ts.prim_func
def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")):
T.evaluate(0)

@R.function
def main(
cond: R.Tensor((), dtype="bool"), x: R.Tensor((2, 3), dtype="float32")
) -> R.Tensor((2, 3), dtype="float32"):
R.func_attr({"relax.force_pure": True})
cls = Before
alloc = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(x, alloc)
out = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(alloc, out)
if cond:
alloc1 = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(out, alloc1)
z = out
else:
alloc2 = R.builtin.alloc_tensor(R.shape([2, 3]), "float32", 0)
cls.exp(out, alloc2)
z = out
return z

after = relax.transform.StaticPlanBlockMemory()(Before)
assert relax.analysis.check_well_formed(after)
assert _count_alloc_storage(after["main"]) == 1


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