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",
[