[Fix][TOPI][WebGPU] Keep sort merge passes off blockIdx.z - #20410
Merged
Merged
Conversation
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.
tlopex
approved these changes
Sep 22, 2026
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
#19900 moved the merge path axis of the GPU sort merge passes to
blockIdx.zon every target. The WebGPU backend reservesblockIdx.zto extendblockIdx.xbeyond 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 intoblockIdx.xand keep the batch onblockIdx.y; other targets keep the #19900 mapping. This also avoids theblockIdx.yoverflow of the mapping before #19900, which produced wrong results for large batches (for example, a batch of 128 at size 151936).