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 827934accc [Fix][TOPI][WebGPU] Keep sort merge passes off blockIdx.z 
(#20410)
827934accc is described below

commit 827934accc4dd6928c41e04b4a3013ce2c12ba5e
Author: Akaash Parthasarathy <[email protected]>
AuthorDate: Tue Sep 22 11:18:58 2026 -0700

    [Fix][TOPI][WebGPU] Keep sort merge passes off blockIdx.z (#20410)
    
    #19900 moved the merge path axis of the GPU sort merge passes to
    `blockIdx.z` on every target. The WebGPU backend reserves `blockIdx.z`
    to extend `blockIdx.x` beyond 65535, so sort, argsort, and topk fail to
    compile for WebGPU whenever the input needs a merge pass. On WebGPU,
    fold the merge path and section axes into `blockIdx.x` and keep the
    batch on `blockIdx.y`; other targets keep the #19900 mapping. This also
    avoids the `blockIdx.y` overflow of the mapping before #19900, which
    produced wrong results for large batches (for example, a batch of 128 at
    size 151936).
---
 python/tvm/topi/gpu/sort.py                        | 31 ++++++++++++++--------
 .../relax/test_backend_dispatch_sort_scan.py       | 23 ++++++++++++++++
 2 files changed, 43 insertions(+), 11 deletions(-)

diff --git a/python/tvm/topi/gpu/sort.py b/python/tvm/topi/gpu/sort.py
index 0971bcccf2..9b9871e6c3 100644
--- a/python/tvm/topi/gpu/sort.py
+++ b/python/tvm/topi/gpu/sort.py
@@ -507,23 +507,32 @@ def _sort_common(
             nbx = cast(ceil_div(width, max_threads * thread_work), "int32")
             nbz = cast(ceil_div(size, width), "int32")
 
-        is_cuda = target.kind.name == "cuda"
         tx = te.thread_axis("threadIdx.x")
-        bx = te.thread_axis("blockIdx.z")  # nbx
-        if is_cuda:
-            by = te.thread_axis("blockIdx.x")  # batch
-            bz = te.thread_axis("blockIdx.y")  # nbz
-        else:
+        if target.kind.name == "webgpu":
+            # WebGPU reserves blockIdx.z to extend blockIdx.x beyond 65535, so 
fold the
+            # merge-path (nbx) and section (nbz) axes into blockIdx.x.
+            bxz = te.thread_axis("blockIdx.x")
             by = te.thread_axis("blockIdx.y")  # batch
-            bz = te.thread_axis("blockIdx.x")  # nbz
-        with T.frame_scope(
-            [
-                T.attr(tx, "thread_extent", ntx),
+            grid_extents = [
+                T.attr(bxz, "thread_extent", nbx * nbz),
+                T.attr(by, "thread_extent", nthread_by),
+            ]
+            bx = tvm.tirx.indexmod(bxz, nbx)
+            bz = tvm.tirx.indexdiv(bxz, nbx)
+        else:
+            bx = te.thread_axis("blockIdx.z")  # nbx
+            if target.kind.name == "cuda":
+                by = te.thread_axis("blockIdx.x")  # batch
+                bz = te.thread_axis("blockIdx.y")  # nbz
+            else:
+                by = te.thread_axis("blockIdx.y")  # batch
+                bz = te.thread_axis("blockIdx.x")  # nbz
+            grid_extents = [
                 T.attr(bx, "thread_extent", nbx),
                 T.attr(by, "thread_extent", nthread_by),
                 T.attr(bz, "thread_extent", nbz),
             ]
-        ):
+        with T.frame_scope([T.attr(tx, "thread_extent", ntx), *grid_extents]):
             base_idx = by * size
 
             # calculate the start, mid, and end points of this section
diff --git a/tests/python/relax/test_backend_dispatch_sort_scan.py 
b/tests/python/relax/test_backend_dispatch_sort_scan.py
index c9f4fda65e..f616a08b61 100644
--- a/tests/python/relax/test_backend_dispatch_sort_scan.py
+++ b/tests/python/relax/test_backend_dispatch_sort_scan.py
@@ -478,6 +478,29 @@ def test_dispatch_topk_cuda_large_batch():
     tvm.testing.run_with_gpu_lock(run_and_check)
 
 
[email protected](not tvm.testing.env.has_llvm(), reason="need llvm for the 
host module")
[email protected]("op", ["sort", "argsort", "topk"])
[email protected]("size", [300, None])
+def test_dispatch_sort_webgpu_merge_grid(op, size):
+    """Merge passes must not use blockIdx.z, which WebGPU reserves to extend 
blockIdx.x."""
+    n = size if size is not None else tirx.Var("n", "int64")
+    x = relax.Var("x", relax.TensorType((tirx.Var("m", "int64"), n), 
"float32"))
+    bb = relax.BlockBuilder()
+    with bb.function("main", [x]):
+        if op == "sort":
+            out = relax.op.sort(x, axis=-1, descending=False)
+        elif op == "argsort":
+            out = relax.op.argsort(x, axis=-1, descending=True, dtype="int32")
+        else:
+            out = relax.op.topk(x, k=4, axis=-1, ret_type="values", 
largest=True)
+        bb.emit_func_output(bb.emit(out))
+
+    target = tvm.target.Target("webgpu", host="llvm")
+    with target:
+        mod = DispatchSortScan()(bb.get())
+    tvm.compile(mod, target)
+
+
 @pytest.mark.parametrize(
     "target",
     [

Reply via email to