Skip to content

[ROCm] Support decode context parallel (DCP) for GLM-5 / DeepSeek-V3.2 - #42618

Draft
EricKing626 wants to merge 9 commits into
sgl-project:mainfrom
EricKing626:amd/dcp-glm-dsa
Draft

EricKing626 wants to merge 9 commits into
sgl-project:mainfrom
EricKing626:amd/dcp-glm-dsa

Conversation

@EricKing626

@EricKing626 EricKing626 commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

[DSA][ROCm] Support decode context parallel (DCP) for GLM-5 / DeepSeek-V3.2

Motivation

Under pure TP, every rank of an MLA model stores a full copy of the KV cache. Decode context parallel (--dcp-size) shards tokens across ranks to remove that duplication. SGLang already supports DCP for dense MLA, but not for DSA models: GLM-5 and DeepSeek-V3.2 run on the dsa backend, which has no DCP path.

DSA picks the top 2048 tokens of the whole sequence before attention. Under DCP this adds three problems:

  • The indexer K cache must be sharded too.
  • No single rank can see the whole sequence to pick the top-k.
  • Each rank may only attend the selected tokens it owns.

With this PR, each rank stores 1/W of both caches, in decode, prefill, MTP speculative decoding and HiCache. The algorithm follows vLLM's DCP for DeepSeek-V3.2 (vllm#58985), including a draft model that inherits the target's DCP layout, and reuses SGLang's existing DCP plumbing.

Modifications

The DSA-specific logic is in a new layers/dcp/dsa.py and a few Triton kernels in kernels/ops/attention/dcp_kernels.py. Other files get small hooks.

  • Writes: MLA KV and indexer K are written only on the owner rank (slot % W == rank).
  • Indexer decode: each rank takes a local top-k over its shard. The ranks all-gather the candidate scores and merge them into the exact global top-k. The merge kernel directly emits the local KV rows this rank owns, so there is no per-layer page-table transform or compaction, and the all-gather carries scores only.
  • Sparse decode: each rank attends only its own top-k slots with the Triton split-K kernel, which also returns the LSE. A custom all-gather merges the partial results.
  • Prefill, chosen per batch:
    • Short prompts (kv_len <= SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD) use MHA one-shot on the all-gathered, dequantized prefix.
    • Otherwise the prefix KV is all-gathered and attended with absorbed MLA. Once the mean key length per query token reaches SGLANG_DCP_DSA_SPLIT_INDEXER_MIN_KV (8192), each rank scores 1/W of the query rows and the int32 top-k is all-gathered.
    • When the cached prefix is long next to the extend (SGLANG_DCP_DSA_OWNED_PREFILL_RATIO, 3.0 × gathered heads × extend tokens), prefill runs like decode: owned top-k plus LSE merge, and the prefix KV is never gathered.
  • MTP / EAGLE (chain, --speculative-eagle-topk 1): the draft shards its KV like the target, so its pool is 1/W the size of a replicated draft. Draft prefill, draft decode and their CUDA graphs reuse the DCP prefill and decode paths. Target verify and draft extend run one query per draft token through the owned top-k plus LSE merge path. MTP IndexShare is off under DCP, because each rank holds only its owned share of the top-k.
  • HiCache (L1/L2): the draft packs into the target's host rows. The DSA indexer host pool pages the widened transfer indices like the MLA host pool.
  • Disabled under DCP: fused top-k and the tilelang fused rope + KV-cache write. Both assume the full KV is local.
  • Tests: test/registered/dcp/test_dcp_dsa_unit.py, test/registered/dcp/test_dcp_layout_unit.py.

Supported scope: ROCm, interleave_size=1, HiCache without an L3 storage backend. Rejected for now: --speculative-eagle-topk > 1, other speculative algorithms, --enable-hisparse and index kpool.

python -m sglang.launch_server --model-path <GLM-5 checkpoint> \
  --tp 4 --dcp-size 4 --kv-cache-dtype fp8_e4m3 \
  --dsa-prefill-backend triton --dsa-decode-backend triton \
  --speculative-algorithm EAGLE --speculative-num-steps 5 \
  --speculative-eagle-topk 1 --speculative-num-draft-tokens 6 \
  --enable-hierarchical-cache --hicache-write-policy write_through

Accuracy Tests

GLM-5.2-MXFP4 on MI355X, fp8 KV cache. MTP uses 5 steps and 6 draft tokens.

Config GSM8K Passkey (9.6k / 38k / 77k tokens)
TP4 0.935 / 0.930 9/9
TP4 + DCP4 0.926 9/9
TP4 + MTP 0.928 –
TP4 + DCP4 + MTP 0.925 9/9
TP4 + DCP4 + MTP + HiCache 0.932 9/9, then 9/9 again from cache
Same, 60k-token device KV pool 0.933 9/9, then 9/9 reloaded from host
  • GSM8K prompts are shorter than the top-k size, so the passkey test covers the cross-rank top-k merge.
  • In the small-pool run, KV usage reaches 0.77 and then drops as prefixes are evicted to host. In the second passkey pass, prefill batches hit up to 153,600 cached tokens while only ~24k tokens are resident on the device, so its KV and indexer K come back from host memory.
  • Accept length under DCP4 + MTP is 4.74, the same as TP4 + MTP (4.74).

Speed Tests and Profiling

Same setup, --max-running-requests 128, random prompts of fixed length, output length 256. TPUT per GPU = output tokens/s ÷ 4. Interactivity = 1000 / TPOT. TTFT, TPOT and ITL are medians.

Config KV pool (tokens)
TP4 3,009,408
TP4 + DCP4 11,081,472 (3.7×)
TP4 + MTP 2,818,624
TP4 + DCP4 + MTP 10,679,552 (3.8×)
input_len concurrency Config Interactivity TPUT per GPU TTFT (ms) TPOT (ms) ITL (ms)
1024 1 TP4 126.0 29.7 125 7.94 7.94
DCP4 81.7 19.5 138 12.24 12.24
TP4 + MTP 486.1 96.7 131 2.06 2.04
DCP4 + MTP 301.2 63.4 147 3.32 3.28
1024 8 TP4 92.6 162.1 356 10.80 10.79
DCP4 62.2 114.5 319 16.08 16.06
TP4 + MTP 212.7 336.2 239 4.70 4.05
DCP4 + MTP 148.9 233.6 339 6.71 6.26
1024 32 TP4 52.7 376.4 669 18.99 17.72
DCP4 36.5 270.4 642 27.42 25.99
TP4 + MTP 102.6 549.1 687 9.74 6.82
DCP4 + MTP 66.4 354.6 734 15.06 12.46
8192 1 TP4 122.9 27.0 328 8.14 8.14
DCP4 71.3 16.4 340 14.03 14.04
TP4 + MTP 476.3 76.4 340 2.10 2.07
DCP4 + MTP 293.6 53.5 354 3.41 3.32
8192 8 TP4 68.2 99.9 1462 14.67 11.12
DCP4 47.6 73.8 1611 21.00 17.99
TP4 + MTP 115.1 119.5 1440 8.69 4.04
DCP4 + MTP 129.1 95.9 2366 7.74 6.37
8192 32 TP4 29.9 147.7 5404 33.40 17.50
DCP4 22.4 121.2 5634 44.57 27.83
TP4 + MTP 41.2 163.6 5110 24.27 6.95
DCP4 + MTP 36.6 132.9 6598 27.34 12.73
32768 1 TP4 121.9 20.8 1237 8.21 8.21
DCP4 71.0 13.9 1225 14.09 14.10
TP4 + MTP 474.3 40.8 1283 2.11 2.08
DCP4 + MTP 296.6 33.7 1273 3.37 3.33
32768 8 TP4 36.3 42.8 4927 27.58 11.06
DCP4 29.5 37.3 5092 33.84 18.00
TP4 + MTP 53.5 49.2 5631 18.68 4.06
DCP4 + MTP 53.9 46.5 6266 18.57 6.34
131072 1 TP4 117.5 9.6 5703 8.51 8.52
DCP4 70.7 8.3 5272 14.14 14.15
TP4 + MTP 459.8 12.4 5850 2.17 2.15
DCP4 + MTP 293.2 12.5 5420 3.41 3.38
131072 4 TP4 31.6 11.6 11795 31.67 9.99
DCP4 22.1 11.5 10636 45.26 15.95
TP4 + MTP 39.2 12.3 14327 25.48 3.10
DCP4 + MTP 34.7 12.9 12547 28.82 4.76

Summary:

  • Capacity: DCP4 holds 3.7× more KV tokens than TP4, and 3.8× with MTP, because the draft pool is sharded too.
  • Prefill: TTFT under DCP4 is on par with TP4, at 0.90–1.11× (0.90× at 1k / c=8, 0.90–0.92× at 128k). Gathering the KV with absorbed MLA for every prompt was 2.5–3× TP4 at 1k input (c=8 / c=32).
  • Decode cost: a DCP4 decode step takes 1.5–1.7× as long as a TP4 step, because of the extra per-layer collectives. Throughput is 0.61–0.99× of TP4 in this sweep. The sweep is not KV-bound, and the gap narrows with longer context (0.99× at 128k / c=4).
  • MTP: DCP4 + MTP delivers 1.10–3.25× the throughput of DCP4 without MTP. Against TP4 + MTP it reaches 0.65–0.94× up to 32k input and 1.01–1.05× at 128k. ITL is 1.53–1.83× TP4 + MTP, close to the 1.47–1.73× without MTP, so MTP adds little DCP-specific overhead.

Long context, long output, 128 concurrent requests (no MTP, --max-running-requests 128, CUDA graphs up to bs=128, 128 prompts at concurrency 128, fixed lengths). Peak running requests and KV usage are sampled from /metrics every 5 s. The TP4 rows come from an earlier session with the same server arguments.

input / output Config Peak running reqs Peak KV usage TPUT per GPU TTFT (s) TPOT (ms) ITL (ms) Duration (s)
64k / 16k TP4 37 1.00 321.9 514.1 23.54 21.03 1629
DCP4 127 0.94 361.9 (1.12×) 159.7 64.78 54.79 1449
128k / 16k TP4 20 0.98 203.2 1200.1 20.04 17.02 2580
DCP4 74 0.99 256.1 (1.26×) 332.8 52.19 44.85 2047
  • TP4 fills its KV cache with 20–37 requests and queues the rest, while DCP4 runs 74–127 at once. Its decode step is 2.6× slower, but with 3.4–3.7× more requests per step its decode throughput (running requests ÷ ITL) is 1.32–1.40× TP4.
  • End to end, DCP4 delivers 1.12–1.26× TP4 throughput and finishes 128k / 16k in 2047 s instead of 2580 s. Less queueing cuts TTFT by 69–72%.
  • Per-request speed is lower: TPOT is 2.6–2.8× TP4. DCP pays off for capacity-bound, throughput-oriented serving, not for latency-bound serving.

Future work

  • The decode gap is on the attention side; MoE and the TP all-reduce cost the same under DCP. A per-step kernel trace shows about 600 tiny torch kernels per step from dcp_localize_write_loc (the owner-rule where on every MLA KV and index-K write). Computing the local write location once per step and fusing it into the write kernels is the cheapest next win.
  • Follow-up PRs: Q-projection replication with fp8 w_kc on gfx950 (--dcp-replicate-q-proj), and projecting V before the LSE merge. In earlier measurements they cut decode ITL by up to 4% and 9% at c=64.

Checklist


CI States

Latest PR Test (Base): Not run yet
Latest PR Test (Extra): ⚠️ Not enabled -- add run-ci-extra label to opt in.
Latest PR Test (AMD ROCm 10): ➖ No AMD PR run found for this commit.

…k-V3.2

The DSA backend had no DCP path, so DSA models could not shard the KV
cache across the DCP group. Reuse the existing DCP plumbing (widened
allocator owner rule, Q all-gather, LSE merge) and add the DSA pieces:

- Writes: MLA fp8 KV and indexer K only land on the owner rank.
- Indexer decode: score the local index-K shard, take a local top-k and
  all-gather (score, global position) candidates to merge the global top-k.
- Indexer prefill: rebuild the full-sequence index K from the shards.
- Sparse decode: keep the owned top-k slots, attend with mla_gluon_decode
  and return the natural-log LSE for the cross-rank merge.
- Sparse prefill: attend the all-gathered dcp_kv_buffer.
- Disable fused top-k, the tilelang fused rope+cache path and MHA one-shot
  prefill under DCP; they assume the full KV is local.

Scope: ROCm, qlen=1 decode; speculative decoding, hisparse and index
kpool are rejected for now.
…ange, custom-AG LSE merge

- Sparse decode under DCP uses the Triton split-K kernel instead of mla_gluon
  (mla_gluon asserts min_kv_seq_len > 784 and DCP shards are short). The
  reduce kernel optionally writes the natural-log LSE (STORE_LSE).
- Owned slots are compacted to the front (dcp_compact_owned_slots) and the
  split kernel bounds its tiles by the per-row owned count.
- dcp_exchange_topk: fused pack / merge Triton kernels; the global merge uses
  the DSA topk_func instead of torch.topk; gather as fp32 so ROCm takes the
  custom all-gather.
- dcp_a2a_lse_reduce: on ROCm use the custom all-gather instead of RCCL a2a.
- Add test/registered/dcp/test_dcp_dsa_unit.py.
…exer, owned top-k + LSE merge

DCP prefill used to gather the whole KV and run absorbed MLA on every rank. Three changes, chosen per batch:

- Short prompts (kv_len <= SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD) take the MHA one-shot path under DCP too; the prefix comes through all_gather_kv_cache_for_mha_extend, which dequantizes the fp8 cache. Absorbed MLA costs ~3.4x the MHA FLOPs per q.k pair, which made short-context TTFT ~3x TP4.
- The gathered-KV indexer splits query rows across ranks once the mean key length per query token reaches SGLANG_DCP_DSA_SPLIT_INDEXER_MIN_KV (default 8192): each rank scores 1/W of the rows against the all-gathered index K and the int32 top-k is all-gathered. Shorter batches score every row on every rank.
- When the cached prefix is long next to the extend (SGLANG_DCP_DSA_OWNED_PREFILL_RATIO, default 3.0 x gathered heads x extend tokens), prefill runs like decode: each rank scores its own index-K shard, the top-k candidates go through the same exchange and merge as decode (global positions, identical on every rank), and attention reads only this rank's owned slots with an LSE merge (forward_batch.dcp_owned_prefill), so the prefix KV is never gathered.

GLM-5.2 MXFP4, TP4 + DCP4 on MI355X: GSM8K 0.933, passkey 9/9 (8K-76K); with every prefill forced onto the owned path (ratio 0, dense threshold 0) GSM8K 0.924, passkey 9/9. Short extend on a cached 16K-128K prefix: 182-308 ms TTFT.
… target; HiCache with MTP

DSA + DCP rejected speculative decoding. Chain EAGLE/NEXTN (--speculative-eagle-topk 1) now runs with the draft inheriting the target's DCP layout, as in vLLM:

- The DSA draft KV pool shards like the target (KVCacheConfigurator.loc_space_scale is 1 for it), so it is 1/W the size of a replicated draft and its writes take the owner rule.
- Draft prefill, draft decode and their CUDA graphs reuse the DCP prefill and decode paths. TARGET_VERIFY and DRAFT_EXTEND_V2 rows are one query per draft token, so the indexer exchanges top-k candidates and forward_extend attends this rank's owned slots and LSE-merges them like decode (is_dcp_mla_decode_phase includes DRAFT_EXTEND_V2).
- The draft runner's extend and multi-step backends are bound to its KVIndexTranslator; the DCP planner reads it through get_attn_backend().
- MTP IndexShare (index_share_for_mtp_iteration) is off under DCP: each rank holds only its owned share of the top-k.
- HiCache + DCP accepts EAGLE/NEXTN on DSA models; the draft packs into the target's host rows. The DSA indexer host pool pages the widened transfer indices (logical_page_size), which previously addressed index-K pages W times too far and faulted once slots passed the per-rank capacity.
@github-actions github-actions Bot added the hicache Hierarchical Caching for SGLang label Oct 8, 2026
The aiter custom all-gather moves W x the bytes the merge needs; it only
wins while the exchange is latency bound. Above
SGLANG_ROCM_DCP_LSE_AG_MAX_BYTES (default 1 MiB per rank) use the byte
all-to-all, which is 2-3x faster at bs 32-64 verify sizes.
Drop SGLANG_ROCM_DCP_LSE_AG_MAX_BYTES and the ROCm custom all-gather
branch; the LSE merge now always uses the byte all-to-all, which moves
1/W of the bytes the all-gather did.
…ens / tables

DcpStep localizes out_cache_loc once per forward (target and draft) and
memoizes the per-rank lens and index block tables, instead of recomputing
them in every layer.
With kv_splits == 1 and return_lse the split partial already is the final
result; write it and the natural-log LSE directly instead of round-tripping
through the split-K workspace and the reduce kernel.
Gate the shared-file changes on DSA so dense MLA + DCP behaves as before:

- is_dcp_mla_decode_phase takes use_dsa; draft extend and owned prefill
  count as the decode phase only for DSA attention.
- The EAGLE v2 draft binds its swapped-in backends to the KV index
  translator only for a DSA draft under DCP.
- dcp_forward_scope opens only for DSA models (target and NextN draft).
@1am9trash 1am9trash added the run-ci CI: run the baseline test suite on this PR label Oct 8, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

amd deepseek hicache Hierarchical Caching for SGLang jit-kernel memory-pool run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants