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 902a2a3e48 [Fix][Relax] Allow callers to specify scan index-width 
budget (#20336)
902a2a3e48 is described below

commit 902a2a3e485ce2896c6372cdc636105ae27c8eed
Author: Akaash Parthasarathy <[email protected]>
AuthorDate: Sun Sep 20 22:13:54 2026 -0700

    [Fix][Relax] Allow callers to specify scan index-width budget (#20336)
    
    Add an optional `index_bits` argument to `DispatchSortScan` so pipelines
    that subsequently force indices to `int32` can generate compatible
    continuous GPU cumsum hierarchies. The hierarchy calculation can
    otherwise emit 64-bit thresholds that fail forced int32 narrowing, as
    seen in MLC's Metal compilation path. Existing callers retain the
    upstream defaults: 32 bits for WebGPU and 64 bits for other targets.
---
 python/tvm/relax/backend/dispatch_sort_scan.py     | 31 +++++++++++++--
 .../relax/test_backend_dispatch_sort_scan.py       | 46 +++++++++++++++++++++-
 2 files changed, 72 insertions(+), 5 deletions(-)

diff --git a/python/tvm/relax/backend/dispatch_sort_scan.py 
b/python/tvm/relax/backend/dispatch_sort_scan.py
index df1543761d..7918c92e03 100644
--- a/python/tvm/relax/backend/dispatch_sort_scan.py
+++ b/python/tvm/relax/backend/dispatch_sort_scan.py
@@ -37,9 +37,10 @@ class SortScanDispatcher(BackendDispatcher):
 
     calls_to_update: dict[GlobalVar, Target]
 
-    def __init__(self, mod):
+    def __init__(self, mod, index_bits: int | None = None):
         super().__init__(mod)
         self.calls_to_update = {}
+        self.index_bits = index_bits
 
     def apply_dlight_gpu_fallback(
         self,
@@ -172,10 +173,15 @@ class SortScanDispatcher(BackendDispatcher):
                 if normalized_axis == len(shape) - 1:
                     outer = reduce(mul, shape_values[:-1], 1)
                     kernel_shape = relax.ShapeExpr([outer, shape[-1]])
+                    index_bits = self.index_bits
+                    if index_bits is None:
+                        index_bits = 32 if tgt.kind.name == "webgpu" else 64
+                    if tgt.kind.name == "webgpu" and index_bits != 32:
+                        raise ValueError("WebGPU scan kernels require 
index_bits=32")
                     kernel = gpu_2d_continuous_cumsum(
                         in_dtype=in_dtype,
                         out_dtype=out_dtype,
-                        index_bits=32 if tgt.kind.name == "webgpu" else 64,
+                        index_bits=index_bits,
                     )
                     kernel_name = "gpu_2d_continuous_cumsum"
                 else:
@@ -263,10 +269,29 @@ class SortScanDispatcher(BackendDispatcher):
 class DispatchSortScan:
     """
     Pass to dispatch scan and sort operators to platform dependent 
implementation.
+
+    Parameters
+    ----------
+    index_bits : Optional[int]
+        Signed index-width budget for the generated continuous GPU cumsum 
hierarchy.
+        Must be 32 or 64. By default, use 32 for WebGPU and 64 for other 
targets.
+        WebGPU does not support an explicit 64-bit budget.
+
+        Pipelines that subsequently force indices to int32 should request 32 to
+        avoid generating hierarchy thresholds outside the signed int32 range.
+        The caller must ensure runtime indices fit the requested width; this
+        option does not insert runtime bounds checks.
+        This option does not narrow the generated TIR, change tensor dtypes, or
+        affect other sort/scan implementations.
     """
 
+    def __init__(self, index_bits: int | None = None):
+        if index_bits not in (None, 32, 64):
+            raise ValueError("index_bits must be either 32 or 64")
+        self.index_bits = index_bits
+
     def transform_module(self, mod: IRModule, ctx: PassContext) -> IRModule:
-        sort_scan_dispater = SortScanDispatcher(mod)
+        sort_scan_dispater = SortScanDispatcher(mod, self.index_bits)
         for gv, func in mod.functions_items():
             if isinstance(func, relax.Function):
                 func = sort_scan_dispater.visit_expr(func)
diff --git a/tests/python/relax/test_backend_dispatch_sort_scan.py 
b/tests/python/relax/test_backend_dispatch_sort_scan.py
index 795a56537e..c9f4fda65e 100644
--- a/tests/python/relax/test_backend_dispatch_sort_scan.py
+++ b/tests/python/relax/test_backend_dispatch_sort_scan.py
@@ -482,10 +482,12 @@ def test_dispatch_topk_cuda_large_batch():
     "target",
     [
         pytest.param("cuda", marks=pytest.mark.gpu),
+        pytest.param("metal", marks=pytest.mark.gpu),
         pytest.param({"kind": "vulkan", "supports_int64": True}, 
marks=pytest.mark.gpu),
     ],
 )
-def test_dispatch_cumsum_gpu(target):
[email protected]("index_bits", [None, 32, 64])
+def test_dispatch_cumsum_gpu(target, index_bits):
     """Test cumsum kernel dispatch and numerical correctness"""
     if not tvm.testing.device_enabled(target):
         pytest.skip(f"{target} not enabled")
@@ -503,7 +505,9 @@ def test_dispatch_cumsum_gpu(target):
     np_data = np.random.randint(0, 10, size).astype("int32")
     np_cumsum = np.cumsum(np_data, axis=-1)
     with tvm.target.Target(target):
-        mod = DispatchSortScan()(Module)
+        mod = DispatchSortScan(index_bits=index_bits)(Module)
+        if index_bits == 32:
+            mod = tirx.transform.ForceNarrowIndexToInt32()(mod)
         ex = tvm.compile(mod, target)
 
     def run_and_check():
@@ -516,6 +520,44 @@ def test_dispatch_cumsum_gpu(target):
     tvm.testing.run_with_gpu_lock(run_and_check)
 
 
[email protected]("target_kind", ["metal", "webgpu", "cuda"])
[email protected]("index_bits", [None, 32, 64])
+def test_dispatch_cumsum_index_width(target_kind, index_bits):
+    """Respect the caller's index budget without restricting Metal's 
default."""
+    from tvm.relax.backend.gpu_generic import gpu_2d_continuous_cumsum
+
+    @I.ir_module
+    class Module:
+        @R.function
+        def main(x: R.Tensor(("m", "n"), "float32")):
+            gv = R.cumsum(x, axis=-1)
+            return gv
+
+    with tvm.target.Target(target_kind, host="llvm"):
+        if target_kind == "webgpu" and index_bits == 64:
+            with pytest.raises(ValueError, match="WebGPU scan kernels require 
index_bits=32"):
+                DispatchSortScan(index_bits=index_bits)(Module)
+            return
+        mod = DispatchSortScan(index_bits=index_bits)(Module)
+
+    expected_bits = (
+        index_bits if index_bits is not None else (32 if target_kind == 
"webgpu" else 64)
+    )
+    expected = gpu_2d_continuous_cumsum(
+        in_dtype="float32", out_dtype="float32", index_bits=expected_bits
+    )
+    assert_structural_equal(mod["gpu_2d_continuous_cumsum"], expected)
+    if expected_bits == 32:
+        # This previously failed on Metal with a 2**35 IntImm.
+        tirx.transform.ForceNarrowIndexToInt32()(mod)
+
+
[email protected]("index_bits", [0, 16, 128])
+def test_dispatch_cumsum_invalid_index_width(index_bits):
+    with pytest.raises(ValueError, match="index_bits must be either 32 or 64"):
+        DispatchSortScan(index_bits=index_bits)
+
+
 @pytest.mark.gpu
 def test_dispatch_cumprod_cuda_large_batch():
     """Test that GPU scan supports more batches than CUDA's grid-y limit."""

Reply via email to