From 95b355c5145a48530d633bd9701db6c904d1cd51 Mon Sep 17 00:00:00 2001 From: Javier de Jesus Date: Wed, 30 Sep 2026 23:41:42 +0000 Subject: [PATCH] [Fix][Relax] Scope planned storage to if branches Fixes #20191. `StaticPlanBlockMemory` plans with one token pool per function, and its rewriter emits `R.memory.alloc_storage` only at the first use of a token. A token released in one `if` branch can be reused in the other branch or after the `if`, where that storage variable is not defined. The planned module is then ill-formed, and the VM segfaults for `cond=False` in the issue script. The rewriter now restores its token-to-storage-variable map when it leaves a `SeqExpr`, so a token first used in a branch gets a new `alloc_storage` wherever it is reused outside that branch. Storage allocated before the `if` is still shared with the branches. Tests: - `python -m pytest -q tests/python/relax/test_transform_static_plan_block_memory.py` (`32 passed`). Two of the three new tests fail without the change; the third pins that storage allocated before the `if` stays shared, and passes either way. - The issue script without `exec_mode="bytecode"` (`relax.build` no longer takes it), built with the default pipeline for llvm, returns `x` for `cond=True` and `cond=False`; before the change `cond=False` segfaults. - `pre-commit run --files src/relax/transform/static_plan_block_memory.cc tests/python/relax/test_transform_static_plan_block_memory.py` --- .../transform/static_plan_block_memory.cc | 14 ++- ...test_transform_static_plan_block_memory.py | 109 ++++++++++++++++++ 2 files changed, 122 insertions(+), 1 deletion(-) diff --git a/src/relax/transform/static_plan_block_memory.cc b/src/relax/transform/static_plan_block_memory.cc index 340caaf5d348..df11b47f1802 100644 --- a/src/relax/transform/static_plan_block_memory.cc +++ b/src/relax/transform/static_plan_block_memory.cc @@ -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"); @@ -1030,7 +1042,7 @@ class StorageAllocationRewriter : public ExprMutator { std::unordered_map alloc_tensor2token_; /*! \brief The mapping from each binding block to the storage tokens that are create inside. */ std::unordered_map> 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 token2storage_var_; }; diff --git a/tests/python/relax/test_transform_static_plan_block_memory.py b/tests/python/relax/test_transform_static_plan_block_memory.py index c8a23ca3c5df..9b4e9c2941d1 100644 --- a/tests/python/relax/test_transform_static_plan_block_memory.py +++ b/tests/python/relax/test_transform_static_plan_block_memory.py @@ -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()