This is an automated email from the ASF dual-hosted git repository.

tlopex pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git


The following commit(s) were added to refs/heads/main by this push:
     new 979cc089cb [Fix][Metal][WebGPU] Preserve variable bindings when 
resolving allocation bounds (#20337)
979cc089cb is described below

commit 979cc089cb1bb0a5616e8472c2d4c5ab9ba9b7f8
Author: Akaash Parthasarathy <[email protected]>
AuthorDate: Sun Sep 20 21:45:15 2026 -0700

    [Fix][Metal][WebGPU] Preserve variable bindings when resolving allocation 
bounds (#20337)
    
    Track primitive `Bind` expressions in a per-function analyzer so
    allocation bounds remain provable after common subexpression
    elimination. Continue rejecting extents whose only upper bound is the
    integer type limit.
---
 src/backend/metal/codegen/codegen_metal.cc         |  11 ++-
 src/backend/metal/codegen/codegen_metal.h          |   2 +
 src/backend/webgpu/codegen/codegen_webgpu.cc       |  11 ++-
 src/backend/webgpu/codegen/codegen_webgpu.h        |   3 +
 tests/python/codegen/test_target_codegen_metal.py  |  79 ++++++++++++++-
 tests/python/codegen/test_target_codegen_webgpu.py | 110 +++++++++++++++++++++
 6 files changed, 209 insertions(+), 7 deletions(-)

diff --git a/src/backend/metal/codegen/codegen_metal.cc 
b/src/backend/metal/codegen/codegen_metal.cc
index fc22e24896..5325cefe6d 100644
--- a/src/backend/metal/codegen/codegen_metal.cc
+++ b/src/backend/metal/codegen/codegen_metal.cc
@@ -27,7 +27,7 @@
 #include <tvm/ffi/container/map.h>
 #include <tvm/ffi/reflection/registry.h>
 #include <tvm/runtime/logging.h>
-#include <tvm/sym/analyzer.h>
+#include <tvm/tirx/analysis.h>
 #include <tvm/tirx/transform.h>
 
 #include <algorithm>
@@ -67,6 +67,7 @@ Var GetSimdgroupBufferVar(const Expr& data) {
 
 void CodeGenMetal::InitFuncState(const PrimFunc& f) {
   CodeGenC::InitFuncState(f);
+  analyzer_ = sym::Analyzer();
   // analyze the data;
   for (Var arg : f->params) {
     if (arg->ty.as<PointerTypeNode>()) {
@@ -331,6 +332,11 @@ void CodeGenMetal::PrintStorageScope(const std::string& 
scope, std::ostream& os)
 }
 
 void CodeGenMetal::Dispatch_(const BindNode* op) {
+  // Stateful reads cannot be substituted after the underlying state changes.
+  if (auto prim_value = op->value.as<PrimExpr>();
+      prim_value && SideEffect(prim_value.value()) <= CallEffectKind::kPure) {
+    analyzer_->Bind(op->var, prim_value.value());
+  }
   const auto* pointer_type = op->var->ty.as<PointerTypeNode>();
   if (pointer_type == nullptr || pointer_type->storage_scope.empty()) {
     return CodeGenC::Dispatch_(op);
@@ -362,11 +368,10 @@ void CodeGenMetal::Dispatch_(const AllocBufferNode* op) {
   this->PrintIndent();
   // Compute a compile-time upper bound on the number of buffer elements.
   size_t constant_size = 1;
-  sym::Analyzer analyzer;
   for (const auto& dim : op->buffer->shape) {
     const auto* dim_imm = dim.as<IntImmNode>();
     int64_t dim_size =
-        dim_imm ? static_cast<int64_t>(dim_imm->value) : 
analyzer->const_int_bound(dim)->max_value;
+        dim_imm ? static_cast<int64_t>(dim_imm->value) : 
analyzer_->const_int_bound(dim)->max_value;
     if (dim_imm == nullptr) {
       // An integer dtype's intrinsic maximum is not a program-derived 
allocation bound.
       TVM_FFI_ICHECK(dim_size != sym::ConstIntBound::kPosInf)
diff --git a/src/backend/metal/codegen/codegen_metal.h 
b/src/backend/metal/codegen/codegen_metal.h
index 7fad4eb398..96b89bb151 100644
--- a/src/backend/metal/codegen/codegen_metal.h
+++ b/src/backend/metal/codegen/codegen_metal.h
@@ -24,6 +24,7 @@
 #ifndef TVM_TARGET_METAL_CODEGEN_METAL_H_
 #define TVM_TARGET_METAL_CODEGEN_METAL_H_
 
+#include <tvm/sym/analyzer.h>
 #include <tvm/target/codegen.h>
 
 #include <string>
@@ -63,6 +64,7 @@ class CodeGenMetal final : public CodeGenC {
   using CodeGenC::PrintType;
 
  private:
+  sym::Analyzer analyzer_;
   std::unordered_map<const VarNode*, std::string> simdgroup_dtype_;
   int thread_index_bits_{32};
   int thread_work_dim_{0};
diff --git a/src/backend/webgpu/codegen/codegen_webgpu.cc 
b/src/backend/webgpu/codegen/codegen_webgpu.cc
index 7d144603bd..3fd0530561 100644
--- a/src/backend/webgpu/codegen/codegen_webgpu.cc
+++ b/src/backend/webgpu/codegen/codegen_webgpu.cc
@@ -27,7 +27,7 @@
 #include <tvm/ffi/reflection/registry.h>
 #include <tvm/ir/prim/builtin.h>
 #include <tvm/support/io.h>
-#include <tvm/sym/analyzer.h>
+#include <tvm/tirx/analysis.h>
 #include <tvm/tirx/builtin.h>
 #include <tvm/tirx/transform.h>
 
@@ -171,6 +171,7 @@ std::string CodeGenWebGPU::Finish() {
 
 void CodeGenWebGPU::InitFuncState(const PrimFunc& f) {
   CodeGenC::InitFuncState(f);
+  analyzer_ = sym::Analyzer();
   workgroup_memory_bytes_ = 0;
   // analyze the data;
   for (Var arg : f->params) {
@@ -652,6 +653,11 @@ void CodeGenWebGPU::Dispatch_(const TensorLoadNode* op, 
std::ostream& os) {  //
 }
 
 void CodeGenWebGPU::Dispatch_(const BindNode* op) {
+  // Stateful reads cannot be substituted after the underlying state changes.
+  if (auto prim_value = op->value.as<PrimExpr>();
+      prim_value && SideEffect(prim_value.value()) <= CallEffectKind::kPure) {
+    analyzer_->Bind(op->var, prim_value.value());
+  }
   // use ssa form.
   if (print_ssa_form_) {
     std::string value = PrintExpr(op->value);
@@ -730,11 +736,10 @@ void CodeGenWebGPU::Dispatch_(const AllocBufferNode* op) {
   TVM_FFI_ICHECK(op->buffer.defined());
   std::string vid = AllocVarID(op->buffer.get());
   size_t constant_size = 1;
-  sym::Analyzer analyzer;
   for (const auto& dim : op->buffer->shape) {
     const auto* dim_imm = dim.as<IntImmNode>();
     int64_t dim_size =
-        dim_imm ? static_cast<int64_t>(dim_imm->value) : 
analyzer->const_int_bound(dim)->max_value;
+        dim_imm ? static_cast<int64_t>(dim_imm->value) : 
analyzer_->const_int_bound(dim)->max_value;
     if (dim_imm == nullptr) {
       const auto* dtype_max = max_value(dim.ty()).as<IntImmNode>();
       // An integer dtype's intrinsic maximum is not a program-derived 
allocation bound.
diff --git a/src/backend/webgpu/codegen/codegen_webgpu.h 
b/src/backend/webgpu/codegen/codegen_webgpu.h
index 489aba3ac8..5c6b57da33 100644
--- a/src/backend/webgpu/codegen/codegen_webgpu.h
+++ b/src/backend/webgpu/codegen/codegen_webgpu.h
@@ -27,6 +27,7 @@
 #ifndef TVM_TARGET_WEBGPU_CODEGEN_WEBGPU_H_
 #define TVM_TARGET_WEBGPU_CODEGEN_WEBGPU_H_
 
+#include <tvm/sym/analyzer.h>
 #include <tvm/target/codegen.h>
 
 #include <cstddef>
@@ -88,6 +89,8 @@ class CodeGenWebGPU final : public CodeGenC {
   void Dispatch_(const ContinueNode* op) final;
 
  private:
+  sym::Analyzer analyzer_;
+
   /*!
    * \brief Enforce value to be U32.
    */
diff --git a/tests/python/codegen/test_target_codegen_metal.py 
b/tests/python/codegen/test_target_codegen_metal.py
index 150d1c40ce..8f6c1a00c3 100644
--- a/tests/python/codegen/test_target_codegen_metal.py
+++ b/tests/python/codegen/test_target_codegen_metal.py
@@ -140,7 +140,7 @@ def test_ramp():
             for i in T.thread_binding(1, thread="threadIdx.x"):
                 with T.sblock("block"):
                     tx = T.axis.spatial(1, i)
-                    r = T.ramp(tx, 3, 2)
+                    r: T.let = T.ramp(tx, 3, 2)
                     A[0, T.ramp(0, 1, 2)] = r
 
     f = tvm.compile(IRModule, target=target)
@@ -410,6 +410,83 @@ def test_bounded_symbolic_stack_allocation():
     assert "thread float scratch[128]" in source
 
 
[email protected]("bounded", [True, False])
+def test_bound_symbolic_stack_allocation(bounded):
+    limit = 64 if bounded else 2147483647
+
+    @I.ir_module
+    class Module:
+        @T.prim_func(s_tir=True)
+        def main(n: T.int32):
+            T.func_attr(
+                {
+                    "calling_conv": 2,
+                    "global_symbol": "main",
+                    "target": T.target("metal"),
+                    "tirx.kernel_launch_params": [],
+                    "tirx.is_global_func": True,
+                }
+            )
+            # Common subexpression elimination can hoist the bounded extent.
+            extent: T.let[T.int32] = T.min(n, limit)
+            elements: T.let[T.int32] = extent * 2
+            scratch = T.alloc_buffer((elements,), "float32", scope="local")
+            T.evaluate(scratch.data)
+
+    if bounded:
+        source = _build_metal(Module).inspect_source()
+        assert "thread float scratch[128]" in source
+    else:
+        with pytest.raises(
+            tvm.error.InternalError,
+            match="Metal allocation extent requires a finite compile-time 
upper bound",
+        ):
+            _build_metal(Module)
+
+
[email protected]("scope", ["local", "shared"])
[email protected]("bounded", [True, False])
+def test_allocation_bound_does_not_substitute_buffer_load(scope, bounded):
+    @I.ir_module
+    class Module:
+        @T.prim_func(s_tir=True)
+        def main():
+            T.func_attr(
+                {
+                    "calling_conv": 2,
+                    "global_symbol": "main",
+                    "target": T.target("metal"),
+                    "tirx.kernel_launch_params": [],
+                    "tirx.is_global_func": True,
+                }
+            )
+            state = T.alloc_buffer((1,), "int32", scope="local")
+            state[0] = 0
+            snapshot: T.let[T.int32] = state[0]
+            state[0] = 32
+            difference: T.let[T.int32] = state[0] - snapshot
+            # The snapshot is immutable, but the buffer it read has changed.
+            # Substituting the load would incorrectly reduce this extent to 1.
+            scratch = T.alloc_buffer(
+                (T.min(T.max(difference, 1), 32 if bounded else 2147483647),),
+                "float32",
+                scope=scope,
+            )
+            scratch[31] = 1.0
+
+    if bounded:
+        source = _build_metal(Module).inspect_source()
+        storage = "threadgroup" if scope == "shared" else "thread"
+        assert f"{storage} float scratch[32]" in source
+        assert "scratch[31] =" in source
+    else:
+        with pytest.raises(
+            tvm.error.InternalError,
+            match="Metal allocation extent requires a finite compile-time 
upper bound",
+        ):
+            _build_metal(Module)
+
+
 def test_bounded_uint64_symbolic_stack_allocation():
     @I.ir_module
     class Module:
diff --git a/tests/python/codegen/test_target_codegen_webgpu.py 
b/tests/python/codegen/test_target_codegen_webgpu.py
index e6dcf1c2d6..42c52144ed 100644
--- a/tests/python/codegen/test_target_codegen_webgpu.py
+++ b/tests/python/codegen/test_target_codegen_webgpu.py
@@ -67,6 +67,116 @@ def test_bounded_symbolic_stack_allocation():
     assert re.search(r"\bvar\s+\w+\s*:\s*array<f32,\s*128>;", source)
 
 
[email protected]("scope", ["local", "shared"])
[email protected]("bounded", [True, False])
+def test_bound_symbolic_allocation(scope, bounded):
+    limit = 64 if bounded else 2147483647
+
+    @I.ir_module
+    class Module:
+        @T.prim_func(s_tir=True)
+        def main(n: T.int32):
+            T.func_attr(
+                {
+                    "calling_conv": 2,
+                    "global_symbol": "main",
+                    "target": T.target("webgpu"),
+                    "tirx.is_global_func": True,
+                }
+            )
+            # Common subexpression elimination can hoist the bounded extent.
+            extent: T.let[T.int32] = T.min(n, limit)
+            first = T.alloc_buffer((extent * 2,), "float32", scope=scope)
+            elements: T.let[T.int32] = extent * 2
+            second = T.alloc_buffer((elements,), "float32", scope=scope)
+            first[0] = 1.0
+            second[0] = first[0]
+
+    if bounded:
+        source = _build_webgpu(Module).inspect_source()
+        declaration = r"var<workgroup>" if scope == "shared" else r"\bvar"
+        assert len(re.findall(declaration + 
r"\s+\w+\s*:\s*array<f32,\s*128>;", source)) == 2
+    else:
+        with pytest.raises(
+            tvm.error.InternalError,
+            match="WebGPU allocation extent requires a finite compile-time 
upper bound",
+        ):
+            _build_webgpu(Module)
+
+
[email protected]("scope", ["local", "shared"])
[email protected]("bounded", [True, False])
+def test_allocation_bound_does_not_substitute_buffer_load(scope, bounded):
+    @I.ir_module
+    class Module:
+        @T.prim_func(s_tir=True)
+        def main():
+            T.func_attr(
+                {
+                    "calling_conv": 2,
+                    "global_symbol": "main",
+                    "target": T.target("webgpu"),
+                    "tirx.is_global_func": True,
+                }
+            )
+            state = T.alloc_buffer((1,), "int32", scope="local")
+            state[0] = 0
+            snapshot: T.let[T.int32] = state[0]
+            state[0] = 32
+            difference: T.let[T.int32] = state[0] - snapshot
+            # The snapshot is immutable, but the buffer it read has changed.
+            # Substituting the load would incorrectly reduce this extent to 1.
+            scratch = T.alloc_buffer(
+                (T.min(T.max(difference, 1), 32 if bounded else 2147483647),),
+                "float32",
+                scope=scope,
+            )
+            scratch[31] = 1.0
+
+    if bounded:
+        source = _build_webgpu(Module).inspect_source()
+        declaration = r"var<workgroup>" if scope == "shared" else r"\bvar"
+        assert re.search(declaration + r"\s+scratch\s*:\s*array<f32,\s*32>;", 
source)
+        assert "scratch[31" in source
+    else:
+        with pytest.raises(
+            tvm.error.InternalError,
+            match="WebGPU allocation extent requires a finite compile-time 
upper bound",
+        ):
+            _build_webgpu(Module)
+
+
[email protected]("target_limit", [512, 496])
+def 
test_bound_symbolic_workgroup_allocation_respects_target_limit(target_limit):
+    @I.ir_module
+    class Module:
+        @T.prim_func(s_tir=True)
+        def main(n: T.int32):
+            T.func_attr(
+                {
+                    "calling_conv": 2,
+                    "global_symbol": "main",
+                    "target": T.target("webgpu"),
+                    "tirx.is_global_func": True,
+                }
+            )
+            extent: T.let[T.int32] = T.min(n, 64)
+            elements: T.let[T.int32] = extent * 2
+            scratch = T.alloc_buffer((elements,), "float32", scope="shared")
+            scratch[0] = 1.0
+
+    target = {"kind": "webgpu", "max_shared_memory_per_block": target_limit}
+    if target_limit == 512:
+        source = _build_webgpu(Module, target).inspect_source()
+        assert re.search(r"var<workgroup>\s+\w+\s*:\s*array<f32,\s*128>;", 
source)
+    else:
+        with pytest.raises(
+            tvm.error.InternalError,
+            match=r"WebGPU workgroup allocations use 512 bytes, .* supports 
only 496 bytes",
+        ):
+            _build_webgpu(Module, target)
+
+
 def test_unbounded_symbolic_stack_allocation_rejected():
     @I.ir_module
     class Module:

Reply via email to