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."""