Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
116 commits
Select commit Hold shift + click to select a range
fea7120
[AMD] HIP siblings of the DeepSeek V4 JIT kernels and the sorted AOT …
kevin-mii Sep 12, 2026
4d8a93d
[AMD] Triton hosts for DeepSeek-V4.1 attention, mHC and decode glue o…
kevin-mii Sep 12, 2026
75895f6
[AMD] MoE, quantization and GEMM kernels for gfx950
kevin-mii Sep 12, 2026
833b98c
[AMD] DeepSeek V4 attention on HIP: radix backend, low-ratio indexer,…
kevin-mii Sep 12, 2026
853cc5c
[AMD] MoE and dense quantization wiring for aiter on ROCm
kevin-mii Sep 12, 2026
e26fdce
[AMD] DeepSeek V4.1 model wiring for HIP: fused mHC boundary, gfx950 …
kevin-mii Sep 12, 2026
a672c34
Pick the DSA decode graph variant from the DP-group max seq len
kevin-mii Sep 12, 2026
7bd847b
[AMD] ROCm image: aiter pin, FlyDSL stage-1 LDS-DMA drain patch, tune…
kevin-mii Sep 12, 2026
da47a34
[AMD] Tests: numerics guards on the served shapes and regression tests
kevin-mii Sep 12, 2026
2b88220
Fix ROCm 10 source pointer types in mHC LDS prefetch
kevin-mii Sep 15, 2026
d6b26c0
[AMD] Integrate dsv4.1 metadata and port KV stores and low-ratio prep…
kevin-mii Sep 15, 2026
5d49163
Remove deployment tuning from the AMD runtime change
kevin-mii Sep 15, 2026
d0e18f1
Keep the AMD environment reference update outside this runtime PR
kevin-mii Sep 15, 2026
9886bdd
Keep the AITER patch compatible with existing whitespace hooks
kevin-mii Sep 15, 2026
cc769af
Place MI35x regression tests under the AMD test directory
kevin-mii Sep 15, 2026
b707873
Restore DSV4.1 EP4 FMoE tuning while upstreaming the configuration
kevin-mii Sep 15, 2026
cab6f15
[AMD] Fuse query RoPE into padded attention output
kevin-mii Sep 16, 2026
9127dc6
[AMD] Overlap shared Engram lookup and MXFP8 projection
kevin-mii Sep 16, 2026
b547826
[AMD] Use paged FP4 indexing for context-parallel prefill
kevin-mii Sep 16, 2026
4ce3b54
[AMD] Honor fixed context limits in prefill graphs
kevin-mii Sep 16, 2026
b8a8270
[AMD] Fix DSV4 pool sizing runtime-context lookup
kevin-mii Sep 16, 2026
dbecb94
[AMD] Support encoder SWA bounded replay on HIP
kevin-mii Sep 16, 2026
bc3c595
[AMD] Preserve CP token ownership during decoder tail replay
kevin-mii Sep 16, 2026
d4165c1
Allow decoder tail replay with single-rank DP attention plumbing
kevin-mii Sep 16, 2026
d416f3c
[AMD] Allow validated V4.1 interleave context parallelism
kevin-mii Sep 16, 2026
920729a
[AMD] Read compact V4.1 FP8 and FP4 KV pages on gfx950
kevin-mii Sep 16, 2026
c10a618
[AMD] Prefetch compact KV tiles during attention compute
kevin-mii Sep 16, 2026
1337da0
[AMD] Port native 16-head swap-AB attention to gfx950
kevin-mii Sep 16, 2026
5ee301e
[AMD] Validate DSpark graphs with production communicator warmup
kevin-mii Sep 16, 2026
aeea303
[AMD] Accelerate mHC prefill with compensated BF16 projections
kevin-mii Sep 16, 2026
64a7220
[AMD] Preserve attention streams so HIP overlap is reachable
kevin-mii Sep 16, 2026
521ecdc
[AMD] Fuse TP4 attention all-reduce with native mHC post
kevin-mii Sep 16, 2026
9d6f13f
[AMD] Fuse TP4 MoE all-reduce with mHC post
kevin-mii Sep 16, 2026
4ce0c36
[AMD] Tune standalone mHC post for medium prefill batches
kevin-mii Sep 16, 2026
c15034d
[AMD] Fix default BF16 SiLU MXFP4 MoE dispatch and clamp
kevin-mii Sep 16, 2026
3d86a94
[AMD] Avoid BF16 WO-A prefill output layout copies
kevin-mii Sep 16, 2026
a5d3e67
[AMD] Port BF16 WO-A decode and verify projection kernels
kevin-mii Sep 16, 2026
10b32cf
[AMD] Construct stream test layers on the target device
kevin-mii Sep 16, 2026
0948a79
[AMD] Use FP32 weights in the AITER mHC numerical test
kevin-mii Sep 16, 2026
acfa898
[AMD] Avoid DSpark sequence-length CPU copies in HIP metadata
kevin-mii Sep 16, 2026
62db90c
[AMD] Build EAGLE draft metadata without CPU length mirrors
kevin-mii Sep 16, 2026
ab67c29
[AMD] Bound the SWA-only draft page-table placeholder
kevin-mii Sep 16, 2026
dd6e9cd
[AMD] Preserve BF16 bits when reconstructing sharded Engram rows
kevin-mii Sep 16, 2026
39729c5
[AMD] Skip Blackwell norm inspection for pending HIP mHC state
kevin-mii Sep 16, 2026
28e76d2
[AMD] Test indexer chunking within oversized requests
jhinpan Sep 16, 2026
abcc4f2
Keep HIP BF16x3 prefill opt-in after TP4 throughput regression
kevin-mii Sep 16, 2026
ae222ce
[AMD] Enable gfx950 V4.1 KV dequantization with exact NaN bits
kevin-mii Sep 16, 2026
f0a587d
[AMD] Verify C2 padding leaves output and pair state untouched
kevin-mii Sep 16, 2026
1ac9be0
[AMD] Allow Engram prefetch for vision-capable V4.1 checkpoints
kevin-mii Sep 16, 2026
5e817a1
Describe ROCm support in encoder SWA replay help
kevin-mii Sep 16, 2026
72c9a9a
[AMD] Fuse WO-A verify reduction with native MXFP8 quantization
kevin-mii Sep 16, 2026
707d341
Describe both HIP WO-A operand types in the return contract
kevin-mii Sep 16, 2026
ee8f8e7
style: Format the relocated HIP KV-store test decorator
kevin-mii Sep 16, 2026
18f1877
test: Preserve compressed KV quantization regressions under AMD
kevin-mii Sep 16, 2026
6f0992e
Keep HIP WO-A quantization fusion opt-in after serving comparison
kevin-mii Sep 16, 2026
fc4e5df
[AMD] Preserve FP32 scores in the V4 router GEMM fallback
kevin-mii Sep 16, 2026
0a39eaa
[AMD] Reconcile main prefill graphs and updated V4.1 compressor APIs
kevin-mii Sep 16, 2026
f64d7d8
test: Keep mixed CUDA and AMD regressions in the vendor taxonomy
kevin-mii Sep 16, 2026
9c2bc53
style: Format the main rebase compatibility updates
kevin-mii Sep 16, 2026
03602a0
test: Update remaining compressor callers and share the KV oracle
kevin-mii Sep 16, 2026
255ccb2
style: Sort updated quantization test imports
kevin-mii Sep 16, 2026
d79b8f8
[AMD] Pass per-token request mapping to raw verify metadata
kevin-mii Sep 16, 2026
fb20ec9
Trim comments to numerical and graph-lifetime constraints
kevin-mii Sep 16, 2026
14cd8a8
Shorten contracts for routing, graph metadata and kernel dispatch
kevin-mii Sep 16, 2026
7418dd2
Keep prerequisite-only comments aligned with upstream
kevin-mii Sep 17, 2026
dda0018
[AMD] Follow upstream FP4 and mHC module moves
kevin-mii Sep 17, 2026
381cfc9
[AMD] Include the renamed FP4 indexer header
kevin-mii Sep 17, 2026
d08a1da
test: Import the candidate reference from its new module
kevin-mii Sep 17, 2026
825e4c0
style: Format the moved candidate-reference import
kevin-mii Sep 17, 2026
c7faf8d
test: Trim experimental suites and duplicate AMD regressions
kevin-mii Sep 17, 2026
1d4531c
style: Format the reduced AMD test suites
kevin-mii Sep 17, 2026
ac55177
test: Keep focused AMD coverage and reuse existing kernel suites
kevin-mii Sep 17, 2026
83ac51d
style: Format the focused AMD regression suites
kevin-mii Sep 17, 2026
6047d94
test: Limit added compressor and inverse-RoPE cases to AMD
kevin-mii Sep 17, 2026
69c924d
test: Place AMD operator suites in kernel subsystems
kevin-mii Sep 17, 2026
700f4a4
test: Consolidate AMD coverage and remove inadmissible cases
kevin-mii Sep 17, 2026
6430728
ci: Discover AMD kernel suites in existing MI350 jobs
kevin-mii Sep 17, 2026
f3de711
style: Format consolidated AMD tests
kevin-mii Sep 17, 2026
2d30666
Reconcile AMD DeepSeek-V4.1 tests and kernels with main's module layout
kevin-mii Sep 22, 2026
049d2a7
Keep the ROCm wo_a fast paths independent of fast_path; update AMD te…
kevin-mii Sep 22, 2026
2010607
[AMD] Serve 9-128 verify rows through the direct-output wo_a path
kevin-mii Sep 22, 2026
ca87fa8
[AMD] Skip the collapsed-input pre-computation when the fused boundar…
kevin-mii Sep 22, 2026
1233032
[AMD] Accept trailing aiter moe_sorting parameters in the fused sorti…
kevin-mii Sep 22, 2026
2721eb6
[AMD] Default to adaptive split-KV and native TP4 attention outside d…
kevin-mii Sep 23, 2026
24e0e05
docs: list the HIP DeepSeek-V4.1 environment variables
kevin-mii Sep 23, 2026
6099f12
test: consolidate the AMD DeepSeek-V4.1 suites per the admission rule
kevin-mii Sep 23, 2026
25f3179
Warn when the gfx95 MXFP8 tile table cannot be read
kevin-mii Sep 23, 2026
9c7ba17
Rename engram's Triton-kernel gate now that HIP takes it
kevin-mii Sep 23, 2026
5a3ee1f
Trim the comments around the ROCm mHC fusion and wo_a gates
kevin-mii Sep 23, 2026
8b011b5
test: follow main's parallel-topology validation in the AMD DeepSeek-…
kevin-mii Sep 23, 2026
321832b
Drop the unused replay flag from the request-window layout
kevin-mii Sep 23, 2026
a4385cb
test: cut and merge the AMD DeepSeek-V4.1 cases per the admission rule
kevin-mii Sep 23, 2026
fd2a348
test: share one TP4 torchrun across the DeepSeek-V4.1 collectives
kevin-mii Sep 23, 2026
557a5a9
test: cover the ROCm sorting fold-in hand-off on CPU
kevin-mii Sep 23, 2026
1d0194c
test: run the shared DeepSeek-V4.1 kernels on the CUDA lanes too
kevin-mii Sep 23, 2026
75f4550
Own the ROCm-resolved layer attributes in the model classes
kevin-mii Sep 23, 2026
ae21fff
Read the aiter override arguments without signature.bind per launch
kevin-mii Sep 23, 2026
87ba395
Drop the unreachable return after the non-low-ratio capture path
kevin-mii Sep 23, 2026
d0ee583
Point the vision top-k gate at its platform-wide branch
kevin-mii Sep 23, 2026
4fd22d8
test: estimate the AMD DeepSeek-V4.1 kernel suites from measured runs
kevin-mii Sep 23, 2026
802c941
Drop the backticks and Sphinx roles from the branch's Python docstrings
kevin-mii Sep 23, 2026
9556408
Apply the code-style review to the AMD DeepSeek-V4.1 port
kevin-mii Sep 23, 2026
24de2ee
Fold the ROCm query rope into the K norm-rope kernel behind a templat…
kevin-mii Sep 23, 2026
13521ff
Keep one regime in the gfx950 MXFP8 gemv: the waves split K
kevin-mii Sep 23, 2026
cbbaef9
Unwrap the gfx950 dense operand once, then dispatch once
kevin-mii Sep 23, 2026
c42ab36
Move the ROCm MoE reduction glue out of deepseek_v2
kevin-mii Sep 23, 2026
8bea71a
Test the trailing aiter sorting parameters by identity
kevin-mii Sep 23, 2026
638b290
Keep main's backticks in docstrings the branch only reflowed
kevin-mii Sep 23, 2026
c3fded8
Drop the ROCm env switches and keep the paths they default to
kevin-mii Sep 23, 2026
7dcde05
Read HIP attributes directly and hoist imports that have no cycle
kevin-mii Sep 23, 2026
e3c0fbf
Restore main's replay floor and drop a call to a missing compressor m…
kevin-mii Sep 23, 2026
d75cc18
Move the attention-DP graph-variant sync to its own branch
kevin-mii Sep 23, 2026
d891abf
Keep main's comments, docstrings and names as main has them
kevin-mii Sep 23, 2026
7f75d9e
Drop the aiter overrides and keep aiter's own sorting and reduction
kevin-mii Sep 24, 2026
86ab3ad
Drop the Docker aiter patches and the DSV4.1 FMoE tuning CSV
kevin-mii Sep 24, 2026
e2e824d
Merge branch 'main' into dsv41-amd-main
kkHuang-amd Sep 24, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions .github/workflows/pr-test-amd.yml
Original file line number Diff line number Diff line change
Expand Up @@ -656,7 +656,7 @@ jobs:
- name: Run test
timeout-minutes: 30
run: |
bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-b-test-1-gpu-small-amd-mi35x ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }}
bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-b-test-1-gpu-small-amd-mi35x,stage-b-kernel-test-1-gpu-amd-mi35x ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }}

stage-b-test-1-gpu-large-amd:
name: ${{ format('stage-b-test-1-gpu-large-amd ({0}, linux-{1}-1gpu-sglang, {2})', matrix.rocm_version, inputs.runner_arch || 'mi300', matrix.part) }}
Expand Down Expand Up @@ -1203,7 +1203,7 @@ jobs:
- name: Run test
timeout-minutes: 60
run: |
bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-c-test-large-8-gpu-amd-mi35x --auto-partition-id ${{ matrix.part }} --auto-partition-size 3 --timeout-per-file 3600 ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }}
bash scripts/ci/amd/amd_ci_exec.sh -w "/sglang-checkout/test" python3 run_suite.py --hw amd --suite stage-c-test-large-8-gpu-amd-mi35x,stage-c-kernel-test-4-gpu-amd-mi35x --auto-partition-id ${{ matrix.part }} --auto-partition-size 3 --timeout-per-file 3600 ${{ needs.check-changes.outputs.continue_on_error == 'true' && '--continue-on-error' || '' }}

# =============================================== DeepSeek-V4 (MI35x, 8-GPU) ====================================================
# GSM8K accuracy on the nightly dsv4 suites, ~20min each. Gated as stage-C so
Expand Down
60 changes: 60 additions & 0 deletions docs/docs/references/environment_variables.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -1867,6 +1867,66 @@ SGLang supports various environment variables that can be used to configure its
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Disable linear-layer quantization on ROCm.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_FUSE_SINKHORN_INTO_NORM</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1 on gfx950: host the mHC Sinkhorn reduce in the next RMSNorm launch instead of a separate kernel.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_MHC_BF16X3_PREFILL</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1 on gfx950: compensated BF16x3 mHC projection for prefill rows. Opt-in; it can regress full-model decode throughput.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_MHC_POST_SPLIT_H</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1 on gfx950: split-H standalone mHC post kernel for medium prefill batches.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_WO_A_BF16_DECODE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1 on gfx950 TP4: native BF16 WO-A GEMV / split-K for decode and verify rows, direct token-major output for large verify batches.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_WO_A_MXFP8_EPILOGUE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1 on gfx950 TP4 verify: quantize the WO-A partial sums directly for the native MXFP8 WO-B.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_ALL_REDUCE_MHC</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1 on gfx950 TP4: fuse the attention and MoE all-reduce with the mHC post for 1-8 row batches.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_FUSED_MOE_REDUCE_ADD</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>aiter MoE on ROCm: add the shared expert inside the FlyDSL top-k reduction launch.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_FUSED_DECODE_GLUE</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4 HIP radix backend: Triton launches for the decode metadata glue (page table, index widening, image select) instead of torch ops.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_HACK_FLASHMLA_BACKEND</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4 HIP decode attention kernel: <code>auto</code> takes aiter's gluon sparse kernel on gfx950 and the tilelang partial + combine elsewhere; <code>aiter_sparse</code>, <code>tilelang</code> or <code>unified_kv_triton</code> force one.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>auto</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_USE_AITER_BATCHED_GEMM</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4 wo_a decode GEMM through aiter's tuned batched GEMM (gfx950); off it stays on the rocBLAS batched GEMM.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true on ROCm</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_OPT_HIP_ATTN_KV_SPLITS</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>HIP sparse decode split-KV count. Unset: adaptive splits and the native TP4 16-head attention, or a fixed 4 under deterministic inference.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>unset</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_ENABLE_DSV41_ENGRAM_KV_PREFETCH</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DeepSeek-V4.1: overlap the engram layer's shared-host lookup and WKV projection with the earlier layers at batch size 1.</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_ROCM_K3_FUSE_KDA_INPROJ</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Kimi-K3 on ROCm: fold the KDA <code>[f_a|b]</code> tail into the wide <code>[q,k,v,g]</code> projection so the whole input projection is one GEMM. Applies to unquantized weights only; falls back to the split projection otherwise.</td>
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/kernels/aot/csrc/common_extension_rocm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ TORCH_LIBRARY_EXPAND(sgl_kernel, m) {

m.def(
"deepseek_v4_topk_transform_512(Tensor scores, Tensor seq_lens, Tensor page_table, Tensor! "
"page_indices, int page_size, Tensor!? raw_indices) -> ()");
"page_indices, int page_size, Tensor!? raw_indices, bool sort_output=False) -> ()");
m.impl("deepseek_v4_topk_transform_512", torch::kCUDA, &deepseek_v4_topk_transform_512);

m.def(
Expand Down
192 changes: 186 additions & 6 deletions python/sglang/kernels/aot/csrc/elementwise/deepseek_v4_topk.cu
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@ constexpr size_t kSMEM = 48 * 1024; // bytes
#endif
static_assert(kSMEM % (2 * sizeof(int32_t)) == 0, "kSMEM must be a multiple of 8 bytes.");

// seq_lens[b] must not exceed scores.size(1) or page_table.size(1) << page_bits: a row reads up to its length
struct TopKParams {
const float* __restrict__ scores;
const int32_t* __restrict__ seq_lens;
Expand All @@ -57,6 +58,8 @@ struct TopKParams {
uint32_t page_bits;
uint32_t topk;
int64_t output_stride;
// Emit each row's picks in ascending order (see bitonic_sort_u32).
bool sort_output;
};

__device__ __forceinline__ uint8_t convert_to_uint8(float x) {
Expand Down Expand Up @@ -251,6 +254,135 @@ radix_topk(const float* __restrict__ input, int32_t* __restrict__ output, uint32
}
}

// Bitonic sort of n (a power of two, 64 <= n <= kMaxTopK) 32-bit keys, one per thread: strides
// below the wavefront width exchange through lane shuffles, the wider ones through LDS.

// lane ^ J's value in registers: DPP for J <= 8, gfx950 permlane swaps for J = 16, 32; __shfl_xor
// is an LDS round trip on every stage's dependent chain.
template <uint32_t J>
__device__ __forceinline__ uint32_t lane_xor(uint32_t v) {
#if defined(__HIP_PLATFORM_AMD__) && (defined(__gfx90a__) || defined(__gfx942__) || defined(__gfx950__))
if constexpr (J == 1) { // quad_perm [1, 0, 3, 2]
return static_cast<uint32_t>(__builtin_amdgcn_update_dpp(0, static_cast<int>(v), 0xB1, 0xF, 0xF, true));
}
if constexpr (J == 2) { // quad_perm [2, 3, 0, 1]
return static_cast<uint32_t>(__builtin_amdgcn_update_dpp(0, static_cast<int>(v), 0x4E, 0xF, 0xF, true));
}
if constexpr (J == 4 || J == 8) {
// within a 16-lane row: bit clear reads J lanes up (row_shl), bit set J lanes down (row_shr)
const uint32_t shl =
static_cast<uint32_t>(__builtin_amdgcn_update_dpp(0, static_cast<int>(v), 0x100 | J, 0xF, 0xF, true));
const uint32_t shr =
static_cast<uint32_t>(__builtin_amdgcn_update_dpp(0, static_cast<int>(v), 0x110 | J, 0xF, 0xF, true));
return (__lane_id() & J) ? shr : shl;
}
#endif
#if defined(__HIP_PLATFORM_AMD__) && defined(__gfx950__)
if constexpr (J == 16) {
// rows 1 and 3 of the first operand swap with rows 0 and 2 of the second
const auto pair = __builtin_amdgcn_permlane16_swap(v, v, false, false);
return (__lane_id() & 16) ? pair[0] : pair[1];
}
if constexpr (J == 32) {
const auto pair = __builtin_amdgcn_permlane32_swap(v, v, false, false);
return (__lane_id() & 32) ? pair[0] : pair[1];
}
#endif
return static_cast<uint32_t>(__shfl_xor(static_cast<int>(v), static_cast<int>(J), 64));
}

// One bitonic stage (merge size K, stride J); templated so the network unrolls with no stride switch.
template <uint32_t K, uint32_t J>
__device__ __forceinline__ uint32_t bitonic_stage(uint32_t v, uint32_t* __restrict__ s_vals, uint32_t tx, uint32_t n) {
const bool up = (tx & K) == 0;
const bool lower = (tx & J) == 0;
uint32_t w;
if constexpr (J >= 64) {
__syncthreads();
if (tx < n) s_vals[tx] = v;
__syncthreads();
w = tx < n ? s_vals[tx ^ J] : ~0u;
} else {
w = lane_xor<J>(v);
}
return (lower == up) ? min(v, w) : max(v, w);
}

template <uint32_t N, uint32_t K, uint32_t J>
__device__ __forceinline__ uint32_t bitonic_network(uint32_t v, uint32_t* __restrict__ s_vals, uint32_t tx) {
v = bitonic_stage<K, J>(v, s_vals, tx, N);
if constexpr (J > 1) {
return bitonic_network<N, K, J / 2>(v, s_vals, tx);
} else if constexpr (K < N) {
return bitonic_network<N, K * 2, K>(v, s_vals, tx);
} else {
return v;
}
}

// Sort the N (a power of two, 64 <= N <= kBlockSize) values in s_vals ascending, one per thread.
template <uint32_t N>
__device__ void bitonic_sort_fixed(uint32_t* __restrict__ s_vals) {
static_assert(N >= 64 && N <= kBlockSize && (N & (N - 1)) == 0, "one value per thread");
const uint32_t tx = threadIdx.x;
uint32_t v = tx < N ? s_vals[tx] : ~0u;
v = bitonic_network<N, 2, 1>(v, s_vals, tx);
__syncthreads();
if (tx < N) s_vals[tx] = v;
__syncthreads();
}

__device__ void bitonic_sort_u32(uint32_t* __restrict__ s_vals, uint32_t n) {
const uint32_t tx = threadIdx.x;
if (n <= kBlockSize) {
switch (n) {
case 64:
bitonic_sort_fixed<64>(s_vals);
return;
case 128:
bitonic_sort_fixed<128>(s_vals);
return;
case 256:
bitonic_sort_fixed<256>(s_vals);
return;
case 512:
bitonic_sort_fixed<512>(s_vals);
return;
default:
if constexpr (kBlockSize >= 1024) {
bitonic_sort_fixed<1024>(s_vals);
return;
}
break;
}
}
// more values than threads (a smaller block than kMaxTopK): every stage through LDS
for (uint32_t k = 2; k <= n; k <<= 1) {
for (uint32_t j = k >> 1; j > 0; j >>= 1) {
for (uint32_t i = tx; i < n; i += kBlockSize) {
const uint32_t partner = i ^ j;
if (partner > i) {
const uint32_t a = s_vals[i];
const uint32_t b = s_vals[partner];
const bool up = (i & k) == 0;
if ((a > b) == up) {
s_vals[i] = b;
s_vals[partner] = a;
}
}
}
__syncthreads();
}
}
}

__device__ __forceinline__ uint32_t next_pow2_at_least_64(uint32_t x) {
uint32_t n = 64;
while (n < x)
n <<= 1;
return n;
}

__global__ __launch_bounds__(kBlockSize) void deepseek_v4_topk_transform_kernel(const TopKParams params) {
const auto bid = blockIdx.x;
const auto seq_len = params.seq_lens[bid];
Expand All @@ -261,15 +393,61 @@ __global__ __launch_bounds__(kBlockSize) void deepseek_v4_topk_transform_kernel(
const auto raw_indices_ptr =
params.raw_indices != nullptr ? params.raw_indices + bid * params.output_stride : nullptr;

__shared__ int32_t s_topk_indices[kMaxTopK];
__shared__ uint32_t s_sort_vals[kMaxTopK];

// key: the position when the row has raw indices, else the slot (sort_selection_rows' order)
const bool key_is_position = raw_indices_ptr != nullptr;
uint32_t count = topk;
if (seq_len <= static_cast<int32_t>(topk)) {
naive_paged_transform(seq_len, topk, params.page_bits, page_ptr, indices_ptr, raw_indices_ptr);
return;
if (!params.sort_output || key_is_position) {
// ascending positions with the -1 padding last: already the sorted row
naive_paged_transform(seq_len, topk, params.page_bits, page_ptr, indices_ptr, raw_indices_ptr);
return;
}
// every position is a pick; the slot order still has to be established
count = static_cast<uint32_t>(seq_len);
for (uint32_t i = threadIdx.x; i < count; i += kBlockSize) {
s_topk_indices[i] = static_cast<int32_t>(i);
}
} else {
radix_topk(score_ptr, s_topk_indices, static_cast<uint32_t>(seq_len), topk);
}
__syncthreads();

__shared__ int32_t s_topk_indices[kMaxTopK];
radix_topk(score_ptr, s_topk_indices, static_cast<uint32_t>(seq_len), topk);
if (params.sort_output) {
const uint32_t n = next_pow2_at_least_64(count);
for (uint32_t i = threadIdx.x; i < n; i += kBlockSize) {
uint32_t key = ~0u;
if (i < count) {
const int32_t raw = s_topk_indices[i];
key = static_cast<uint32_t>(
key_is_position ? raw : page_to_slot(page_ptr, static_cast<uint32_t>(raw), params.page_bits));
}
s_sort_vals[i] = key;
}
__syncthreads();
bitonic_sort_u32(s_sort_vals, n);
for (uint32_t i = threadIdx.x; i < topk; i += kBlockSize) {
int32_t slot = -1;
int32_t raw = -1;
if (i < count) {
const int32_t key = static_cast<int32_t>(s_sort_vals[i]);
if (key_is_position) {
raw = key;
slot = page_to_slot(page_ptr, static_cast<uint32_t>(raw), params.page_bits);
} else {
slot = key;
}
}
indices_ptr[i] = slot;
if (raw_indices_ptr != nullptr) {
raw_indices_ptr[i] = raw;
}
}
return;
}

__syncthreads();
for (uint32_t i = threadIdx.x; i < topk; i += kBlockSize) {
const auto raw = s_topk_indices[i];
indices_ptr[i] = page_to_slot(page_ptr, static_cast<uint32_t>(raw), params.page_bits);
Expand Down Expand Up @@ -304,7 +482,8 @@ void deepseek_v4_topk_transform_512(
const at::Tensor& page_table,
at::Tensor& page_indices,
int64_t page_size,
std::optional<at::Tensor> raw_indices_opt) {
std::optional<at::Tensor> raw_indices_opt,
bool sort_output) {
CHECK_CUDA(scores);
CHECK_CUDA(seq_lens);
CHECK_CUDA(page_table);
Expand Down Expand Up @@ -367,6 +546,7 @@ void deepseek_v4_topk_transform_512(
.page_bits = page_bits,
.topk = static_cast<uint32_t>(topk),
.output_stride = topk,
.sort_output = sort_output,
};

const auto stream = at::cuda::getCurrentCUDAStream().stream();
Expand Down
3 changes: 2 additions & 1 deletion python/sglang/kernels/aot/include/sgl_kernel_ops.h
Original file line number Diff line number Diff line change
Expand Up @@ -170,7 +170,8 @@ void deepseek_v4_topk_transform_512(
const at::Tensor& page_table,
at::Tensor& page_indices,
int64_t page_size,
std::optional<at::Tensor> raw_indices_opt = std::nullopt);
std::optional<at::Tensor> raw_indices_opt = std::nullopt,
bool sort_output = false);
#endif

/*
Expand Down
9 changes: 7 additions & 2 deletions python/sglang/kernels/aot/python/sgl_kernel/top_k.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,7 @@ def deepseek_v4_topk_transform_512(
page_indices: torch.Tensor,
page_size: int,
raw_indices: Optional[torch.Tensor] = None,
sort_output: bool = False,
) -> None:
"""
Performs the DeepSeek-V4 indexer top-k selection and writes the paged
Expand All @@ -96,19 +97,23 @@ def deepseek_v4_topk_transform_512(

Args:
scores: float32 ``[B, max_seq_len]`` indexer logits, contiguous on dim 1.
seq_lens: int32 ``[B]``, true KV length per batch row.
seq_lens: int32 [B], true KV length per batch row; each at most
max_seq_len and at most num_pages * page_size (a row reads
its scores and page table up to its length).
page_table: int32 ``[B, num_pages]``, logical->physical page table,
contiguous on dim 1.
page_indices: int32 ``[B, topk]``, output buffer, contiguous. Filled
with paged physical slots; -1 for padding entries.
page_size: power-of-2 page size.
raw_indices: optional int32 ``[B, topk]``, contiguous. If provided,
filled with raw token positions within each row.
sort_output: order every row ascending by position (-1 padding last)
in the kernel epilogue, so the consumer needs no sort launch.
"""
if raw_indices is not None:
assert raw_indices.dim() == 2
torch.ops.sgl_kernel.deepseek_v4_topk_transform_512(
scores, seq_lens, page_table, page_indices, page_size, raw_indices
scores, seq_lens, page_table, page_indices, page_size, raw_indices, sort_output
)


Expand Down
Loading
Loading