Skip to content

[Fix][TOPI][WebGPU] Keep sort merge passes off blockIdx.z - #20410

Merged
tlopex merged 1 commit into
apache:mainfrom
akaashrp:fix/webgpu-sort-merge-grid
Sep 22, 2026
Merged

tlopex merged 1 commit into
apache:mainfrom
akaashrp:fix/webgpu-sort-merge-grid

Conversation

@akaashrp

Copy link
Copy Markdown
Contributor

#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).

apache#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 CodeGenWebGPU now
rejects every sort, argsort, and topk whose size needs a merge pass (larger than one block).
This includes the argsort that MLC-LLM attaches to every WebGPU model library.

On WebGPU, fold the merge-path and section axes into blockIdx.x, whose extent the backend can
already extend, and keep the batch on blockIdx.y. Other targets keep the apache#19900 mapping. The
mapping before apache#19900 folded batch and sections into blockIdx.y, which exceeds WebGPU's 65535
workgroups per dimension for large batches (e.g. a batch of 128 at vocabulary size 151936).

Add WebGPU compile tests for sort, argsort, and topk with static and dynamic sizes that require
merge passes. Forcing the folded mapping on Metal gives exact sort, argsort, and topk results
for batches up to 70 and sizes up to 262144.
@akaashrp
akaashrp requested a review from tlopex September 22, 2026 10:06
@tlopex
tlopex merged commit 827934a into apache:main Sep 22, 2026
8 checks passed
akaashrp added a commit to mlc-ai/mlc-llm that referenced this pull request Sep 29, 2026
…3554)

* [Fix] Match scan hierarchy width to the MLC index policy

Pass an explicit index-width budget to DispatchSortScan and use the same policy for subsequent forced narrowing. Non-CUDA pipelines request 32-bit-compatible scan hierarchies; CUDA retains 64 bits. This avoids out-of-range hierarchy thresholds on Metal without restricting TVM Metal callers globally.

Add pipeline regression coverage for Metal, WebGPU, and CUDA to verify that scan dispatch agrees with the actual narrowing policy. Requires the companion TVM DispatchSortScan(index_bits=...) API change.

* [Compiler] Use TVM structural traversal for symbolic shape rewrites

TVM removed tirx.stmt_functor.substitute and post_order_visit. Rewrite the lifted-buffer and pipeline-parallel shape substitutions with tvm_ffi.structural_map and structural_walk, keeping the post-order Var replacement semantics.

Add coverage that pipeline shapes share one fresh symbol per undefined variable and that lifted buffer shapes resolve to the caller's symbols.

* [Refactor] Adapt to TVM constant and S-TIR node refactors

apache/tvm#20386 removed relax.Constant, relax.StringImm and relax.DataTypeImm in favor of
tvm.ir.GenericConst and the shared tvm.ir.StringImm, and constant payloads are read through
.value. apache/tvm#20378 moved SBlock and SBlockRealize from tvm.tirx to tvm.s_tir.

Update the KV-cache creation arguments, the embedding allocator, memory-usage metadata, RNN
state initial values, FP8/per-tensor scale constants, and the passes that construct or match
root blocks. Data-type arguments use GenericConst(DataType(...), AnyType()), matching TVM's
own replacement for DataTypeImm.

* [Fix] Narrow scheduled TIR with the S-TIR index narrowing pass

After apache/tvm#20378, tirx.transform.ForceNarrowIndexToInt32 expects functions without S-TIR
blocks. MLC narrows scheduled TIR before block lowering in its pipeline and before scheduling
in FuseDequantizeTake, so tensorized blocks kept int64 iterators with int32 domains and every
q4f16_1 Metal build failed with "mismatched types. int64 vs. int32" in CompactBufferAllocation.

Use s_tir.transform.ForceNarrowIndexToInt32, which rewrites block iterators, access regions
and match-buffer regions, at both call sites. The scan index-policy test tracks the S-TIR pass.

* [Refactor] Adopt the S-TIR script namespace and typed buffer parameters

Apache TVM gave S-TIR its own TVMScript namespace and stopped accepting handle parameters bound
by match_buffer, so every S-TIR prim func in the tree needs both changes.

Namespace: `@T.prim_func(s_tir=True)` becomes `@Ts.prim_func`, and the block-level constructs
`sblock`, `sblock_alloc_buffer`, `axis`, `reads`, `writes`, `init`, `where` and `match_buffer`
now come from `tvm.script.s_tir`. The frames behind the MoE top-k cascade were renamed as well,
so `T.If`/`T.Then`/`T.Else` become `T.if_`/`T.then_`/`T.else_`.

Buffer parameters: a buffer that is a prim-func parameter is declared in the signature, and the
symbolic extents its shape needs are declared above the decorator with `T.dynamic` rather than
inside the body. 175 parameters across 16 files move this way. Three spots needed more than the
mechanical change:

- `_attach_take_probs_func` declares int64 relax vars under the same names as the prim func's
  int32 extents, so those declarations now sit after the prim func, which captures its extents
  when it is parsed.
- `get_indptr` has a buffer whose extent is a later scalar parameter. The plain spelling silently
  resolves the extent against the enclosing scope and bakes in a constant, so it uses the form
  upstream covers in test_tir_external_symbol_adopted_by_later_prim_param instead.
- `rnn_state` loaded slot indices into block-local scalars, which leaves the inferred read region
  referring to them and now fails the well-formedness check. The loads are inlined with
  `T.meta_var`.

* [Fix] Use DataTypeImm again for dtype arguments

TVM removed DataTypeImm and then restored it as the node shared across script dialects, so the
GenericConst spelling this branch adopted in the meantime no longer matches what TVMScript
produces for T.dtype(...). The two nodes both carry the dtype, so compilation worked either way,
but a module built by mlc_llm was no longer structurally equal to the same module written in
TVMScript, which is what tests/python/model/test_kv_cache.py checks.

* [Test] Track TVM's explicit dynamic symbols

Symbolic shapes are now declared with T.dynamic or I.dynamic instead of bare strings inside a
shape annotation. Update the scan index-policy test accordingly and regenerate the paged KV cache
expectation, which also picks up R.Any results and the ty_args spelling.

* [Deps] Bump TVM to Apache main with the S-TIR, WebGPU and attention-merge fixes

mlc-ai/relax mlc now tracks Apache TVM main 708fd43ec8 with the four MLC-only commits on top.
The preceding commits in this branch adapt to the APIs that moved in that range. Three fixes it
carries matter here:

- apache/tvm#20409 adds the block-aware S-TIR ForceNarrowIndexToInt32 this branch switches the
  pipeline to, without which every q4f16_1 Metal build fails on an int64/int32 mismatch in
  CompactBufferAllocation.
- apache/tvm#20410 keeps the sort merge passes off blockIdx.z, which WebGPU reserves, so
  argsort_probs and therefore every WebGPU build links again.
- apache/tvm#20449 synchronizes merge_state_inplace before it overwrites the shared LSE, which
  made prefill nondeterministic for any head_dim above 128, Gemma 4 included.

* [Test] Wrap the regenerated KV cache expectation

The printer emits one line per function signature, which exceeds the 100-column limit. Drop the
formatter fence around the expected module so ruff wraps it.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants