Add DCP to runtime parallel context (#30478)

This commit is contained in:
Baizhou Zhang
2026-07-10 21:23:41 -07:00
committed by GitHub
parent bbcfcaeefe
commit 07165d5daa
16 changed files with 167 additions and 105 deletions
@@ -5,11 +5,7 @@ import triton
import triton.language as tl
from sglang.jit_kernel.utils import is_arch_support_pdl
from sglang.srt.layers.dcp import (
dcp_enabled,
get_attention_dcp_rank,
get_attention_dcp_world_size,
)
from sglang.srt.runtime_context import get_parallel
@triton.jit
@@ -134,7 +130,7 @@ def set_mla_kv_buffer_triton(
n_loc >= _TMA_BULK_STORE_MIN_LOCS
and is_arch_support_pdl()
and can_use_set_mla_kv_buffer(nope_bytes, rope_bytes)
and not dcp_enabled()
and not get_parallel().dcp_enabled
):
jit_set_mla_kv_buffer(kv_buffer, loc, cache_k_nope, cache_k_rope)
return
@@ -161,8 +157,8 @@ def set_mla_kv_buffer_triton(
nope_dim,
rope_dim,
BLOCK=BLOCK,
DCP_RANK=get_attention_dcp_rank(),
DCP_WORLD_SIZE=get_attention_dcp_world_size(),
DCP_RANK=get_parallel().attn_dcp_rank,
DCP_WORLD_SIZE=get_parallel().attn_dcp_size,
**pdl_kwargs,
)
@@ -25,8 +25,6 @@ from sglang.srt.layers.attention.flashinfer_backend import (
from sglang.srt.layers.attention.utils import assert_buffer_fits
from sglang.srt.layers.dcp import (
DecodeContextParallelMetadata,
dcp_enabled,
get_attention_dcp_world_size,
update_local_kv_lens_for_dcp,
)
from sglang.srt.layers.dcp.planner import plan_dcp_decode_metadata
@@ -647,7 +645,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
k_buffer[:, :, layer.v_head_dim :],
out=o,
# for decode forward_batch, each dcp rank computes total q and partial kv, thus, we need to return_lse for online softmax to get final attn_output
return_lse=forward_batch.forward_mode.is_decode() and dcp_enabled(),
return_lse=(
forward_batch.forward_mode.is_decode() and get_parallel().dcp_enabled
),
)
if isinstance(o, tuple):
out, lse = o
@@ -662,7 +662,7 @@ class FlashInferMLAIndicesUpdaterDecode:
self.num_local_heads = (
model_runner.model_config.num_attention_heads
// get_parallel().attn_tp_size
* get_attention_dcp_world_size()
* get_parallel().attn_dcp_size
)
self.kv_lora_rank = model_runner.model_config.kv_lora_rank
self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim
@@ -733,7 +733,7 @@ class FlashInferMLAIndicesUpdaterDecode:
self.req_to_token.shape[1],
)
if dcp_enabled():
if get_parallel().dcp_enabled:
plan_dcp_decode_metadata(
kv_lens,
kv_indptr,
@@ -17,11 +17,6 @@ from sglang.srt.layers.attention.utils import (
create_flashmla_kv_indices_triton,
get_num_kv_index_blocks_flashmla,
)
from sglang.srt.layers.dcp import (
dcp_enabled,
get_attention_dcp_rank,
get_attention_dcp_world_size,
)
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
@@ -96,8 +91,8 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_num_splits_view = None
# get dcp info
self.dcp_world_size = get_attention_dcp_world_size()
self.dcp_rank = get_attention_dcp_rank()
self.dcp_world_size = get_parallel().attn_dcp_size
self.dcp_rank = get_parallel().attn_dcp_rank
def init_forward_metadata_out_graph(
self,
@@ -390,7 +385,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
# TODO uniform output for forward_decode and forward_extend to
# return tuple instead of single output
# decode context parallel needs lse to correct attn_output via online softmax
if dcp_enabled():
if get_parallel().dcp_enabled:
return o, lse
return o
@@ -14,7 +14,6 @@ from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.distributed.parallel_state import get_dcp_group
from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.dcp import (
@@ -166,8 +165,8 @@ class TritonAttnBackend(AttentionBackend):
and self.topk == 1
)
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.dcp_size = getattr(model_runner, "dcp_size", 1)
self.dcp_rank = getattr(model_runner, "dcp_rank", 0)
self.dcp_size = get_parallel().attn_dcp_size
self.dcp_rank = get_parallel().attn_dcp_rank
self.num_head = (
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) * self.dcp_size
@@ -1387,7 +1386,7 @@ class TritonAttnBackend(AttentionBackend):
"DCP Triton extend does not support sliding window"
)
group = get_dcp_group()
group = get_parallel().dcp_group
q_local = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim).contiguous()
total_tokens, local_heads, _ = q_local.shape
@@ -1712,7 +1711,7 @@ class TritonAttnBackend(AttentionBackend):
attn_logits = self.forward_metadata.swa_attn_logits
if self.dcp_size > 1:
group = get_dcp_group()
group = get_parallel().dcp_group
with use_symmetric_memory(group):
q_for_decode = q.view(
-1, layer.tp_q_head_num, layer.qk_head_dim
+5 -1
View File
@@ -23,7 +23,11 @@ Only the symbols imported by code OUTSIDE this subpackage are re-exported here.
Package-internal helpers (the @triton.jit kernels, ``CPTritonContext``,
``correct_attn_out``, ``create_dcp_kv_indices``, ``update_kv_lens_and_indices``,
``_all_gather_dcp_kv_cache``) stay private to their submodules — import them from
``sglang.srt.layers.dcp.{kernels,comm}`` if ever needed internally."""
``sglang.srt.layers.dcp.{kernels,comm}`` if ever needed internally.
``dcp_enabled`` / ``get_attention_dcp_*`` remain compatibility exports for
out-of-tree callers; in-tree code should use ``get_parallel().dcp_enabled`` and
``get_parallel().attn_dcp_*``."""
from sglang.srt.layers.dcp.comm import (
all_gather_kv_cache_for_dcp,
+34 -28
View File
@@ -20,6 +20,7 @@ PR #25090 vs #14194):
- cp_lse_ag_out_rs_mla: Triton (log2/exp2) correction / reduce-scatter
"""
import warnings
from typing import Optional
import torch
@@ -27,37 +28,39 @@ import torch
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.distributed.parallel_state import (
GroupCoordinator,
get_dcp_group,
get_dcp_group_no_assert,
)
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.layers.dcp.kernels import CPTritonContext, correct_attn_out
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_cuda
def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None:
warnings.warn(
f"{name} is deprecated; use {replacement} instead.",
DeprecationWarning,
stacklevel=2,
)
def dcp_enabled() -> bool:
"""
only checks whether dcp enabled for cuda platform
"""
if get_dcp_group_no_assert() is None:
return False
if not is_cuda():
return False
return get_parallel().dcp_size > 1
"""Deprecated: use ``get_parallel().dcp_enabled``."""
_warn_deprecated_dcp_accessor("dcp_enabled()", "get_parallel().dcp_enabled")
return get_parallel().dcp_enabled
def get_attention_dcp_world_size() -> int:
if not dcp_enabled():
return 1
return get_parallel().dcp_size
"""Deprecated: use ``get_parallel().attn_dcp_size``."""
_warn_deprecated_dcp_accessor(
"get_attention_dcp_world_size()", "get_parallel().attn_dcp_size"
)
return get_parallel().attn_dcp_size
def get_attention_dcp_rank() -> int:
if not dcp_enabled():
return 0
return get_parallel().dcp_rank
"""Deprecated: use ``get_parallel().attn_dcp_rank``."""
_warn_deprecated_dcp_accessor(
"get_attention_dcp_rank()", "get_parallel().attn_dcp_rank"
)
return get_parallel().attn_dcp_rank
def _ag_lse(cp_attn_lse: torch.Tensor, cp_group: GroupCoordinator) -> torch.Tensor:
@@ -132,12 +135,13 @@ def cp_lse_ag_out_rs_mla(
def _all_gather_dcp_kv_cache(kv_a: torch.Tensor):
dcp_world_size = get_parallel().dcp_size
parallel = get_parallel()
dcp_world_size = parallel.dcp_size
# not use symmetric_memory unless torch mem_pool updated, see https://github.com/pytorch/pytorch/issues/178138
gathered_kv_a = kv_a.new_empty(
(kv_a.shape[0] * dcp_world_size, *kv_a.shape[1:]),
)
get_dcp_group().all_gather_into_tensor(gathered_kv_a, kv_a)
parallel.dcp_group.all_gather_into_tensor(gathered_kv_a, kv_a)
gathered_kv_a = (
gathered_kv_a.reshape((dcp_world_size,) + kv_a.shape)
.transpose(0, 1)
@@ -152,7 +156,7 @@ def all_gather_kv_cache_for_mha_chunk_extend(
prefix_kv_lens_cpu: torch.Tensor,
prefix_starts_cpu: torch.Tensor = None,
):
if dcp_enabled():
if get_parallel().dcp_enabled:
kv_a = kv_a.unsqueeze(1)
gathered_kv = all_gather_kv_cache_for_dcp(
kv_a,
@@ -218,10 +222,11 @@ def all_gather_q_for_mla_decode(
q_nope_out: torch.Tensor,
q_pe: torch.Tensor,
):
with use_symmetric_memory(get_dcp_group()):
group = get_parallel().dcp_group
with use_symmetric_memory(group):
# transpose q_pe and q_nope_out from [B, H, L] to [H, B, L]
combined = torch.cat([q_pe.transpose(0, 1), q_nope_out.transpose(0, 1)], dim=-1)
gathered = get_dcp_group().all_gather(combined, dim=0)
gathered = group.all_gather(combined, dim=0)
d_pe = q_pe.size(-1)
d_nope = q_nope_out.size(-1)
q_pe, q_nope_out = gathered.split([d_pe, d_nope], dim=-1)
@@ -278,11 +283,12 @@ def all_gather_kv_cache_for_dcp(
"""
prefix_kv_a and prefix_k_pe should have same shape, expect for last dim
"""
if not dcp_enabled():
parallel = get_parallel()
if not parallel.dcp_enabled:
return torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
# 1. compute max kv_lens for each seq
dcp_world_size = get_parallel().dcp_size
dcp_rank = get_parallel().dcp_rank
dcp_world_size = parallel.dcp_size
dcp_rank = parallel.dcp_rank
if prefix_starts_cpu is None:
prefix_starts_cpu = torch.zeros_like(prefix_kv_lens_cpu)
+7 -8
View File
@@ -17,7 +17,6 @@ the owner-rule local-index filter."""
import torch
from sglang.srt.layers.dcp.comm import dcp_enabled
from sglang.srt.runtime_context import get_parallel
@@ -43,10 +42,11 @@ def get_dcp_lens(
def filter_dcp_local_kv_indices(kv_indices: torch.Tensor):
if dcp_enabled():
parallel = get_parallel()
if parallel.dcp_enabled:
kv_indices = (
kv_indices[kv_indices % get_parallel().dcp_size == get_parallel().dcp_rank]
// get_parallel().dcp_size
kv_indices[kv_indices % parallel.dcp_size == parallel.dcp_rank]
// parallel.dcp_size
)
return kv_indices
@@ -59,8 +59,7 @@ def update_local_kv_lens_for_dcp(kv_len_arr):
in-place mutation because callers (plan_dcp_decode_metadata, the FlashInfer-MLA
cuda-graph replay path) rely on it.
"""
if not dcp_enabled():
parallel = get_parallel()
if not parallel.dcp_enabled:
return
kv_len_arr.copy_(
get_dcp_lens(kv_len_arr, get_parallel().dcp_size, get_parallel().dcp_rank)
)
kv_len_arr.copy_(get_dcp_lens(kv_len_arr, parallel.dcp_size, parallel.dcp_rank))
+8 -7
View File
@@ -20,7 +20,6 @@ from typing import Optional
import torch
from sglang.srt.layers.dcp.comm import dcp_enabled
from sglang.srt.layers.dcp.kernels import (
create_dcp_kv_indices,
update_kv_lens_and_indices,
@@ -43,7 +42,8 @@ def prepare_decode_context_parallel_metadata(
kv_cache_device,
create_chunked_prefix_cache_kv_indices_fn,
) -> Optional[DecodeContextParallelMetadata]:
if not dcp_enabled():
parallel = get_parallel()
if not parallel.dcp_enabled:
return None
# dcp_kv_buffer tokens' layout
# [ rank0_r1.prefix_tokens, rank1_r1.prefix_tokens, ..., rank7_r1.prefix_tokens,
@@ -107,13 +107,13 @@ def prepare_decode_context_parallel_metadata(
extend_cu_prefix_lens,
dcp_kv_indices,
extend_prefix_lens_sum,
get_parallel().dcp_size,
parallel.dcp_size,
)
dcp_local_prefix_kv_indices = (
dcp_prefix_kv_indices[
dcp_prefix_kv_indices % get_parallel().dcp_size == get_parallel().dcp_rank
dcp_prefix_kv_indices % parallel.dcp_size == parallel.dcp_rank
]
// get_parallel().dcp_size
// parallel.dcp_size
)
dcp_kv_buffer = torch.empty(
(
@@ -141,6 +141,7 @@ def plan_dcp_decode_metadata(
fast_decode_kwargs: dict,
bs: int,
):
parallel = get_parallel()
local_kv_lens = kv_lens.clone()
update_local_kv_lens_for_dcp(local_kv_lens)
local_kv_lens.clamp_(min=0)
@@ -178,8 +179,8 @@ def plan_dcp_decode_metadata(
local_kv_lens,
local_kv_lens_cumsum,
local_kv_indices,
dcp_rank=get_parallel().dcp_rank,
dcp_world_size=get_parallel().dcp_size,
dcp_rank=parallel.dcp_rank,
dcp_world_size=parallel.dcp_size,
BLOCK_SIZE=BLOCK_SIZE,
)
kv_indices[:total_local_len] = local_kv_indices[:total_local_len]
@@ -25,11 +25,11 @@ from typing import TYPE_CHECKING
from sglang.srt.configs.model_config import ModelImpl
from sglang.srt.environ import envs
from sglang.srt.layers.dcp import dcp_enabled
from sglang.srt.managers.mm_utils import init_mm_embedding_cache
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.registry import TreeCacheBuildContext, create_tree_cache
from sglang.srt.model_loader.utils import get_resolved_model_impl
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING:
@@ -204,7 +204,9 @@ def build_kv_cache(
# TreeCache.page_size should keep the same as allocator.page_size to
# avoid kv page eviction conflicts.
page_size=(
page_size if not dcp_enabled() else token_to_kv_pool_allocator.page_size
page_size
if not get_parallel().dcp_enabled
else token_to_kv_pool_allocator.page_size
),
is_eagle=spec_algorithm.is_eagle(),
tp_cache_group=(
+4 -9
View File
@@ -50,11 +50,6 @@ from sglang.srt.layers.attention.dsa.quant_k_cache import (
quantize_k_cache_separate,
)
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
from sglang.srt.layers.dcp import (
dcp_enabled,
get_attention_dcp_rank,
get_attention_dcp_world_size,
)
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
@@ -73,6 +68,7 @@ from sglang.srt.mem_cache.utils import (
set_mla_kv_scale_buffer_triton,
)
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import (
cpu_has_amx_support,
is_cpu,
@@ -2834,10 +2830,9 @@ class MLATokenToKVPool(KVCache):
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
layer_id = layer.layer_id
assert not self.dsa_kv_cache_store_fp8
if dcp_enabled():
valid_mask = (
loc % get_attention_dcp_world_size() == get_attention_dcp_rank()
)
parallel = get_parallel()
if parallel.dcp_enabled:
valid_mask = loc % parallel.attn_dcp_size == parallel.attn_dcp_rank
if not valid_mask.all():
loc = loc[valid_mask]
cache_k = cache_k[valid_mask]
@@ -12,7 +12,6 @@ from sglang.srt.layers.communicator import get_attn_tp_context
from sglang.srt.layers.dcp import (
all_gather_kv_cache_for_mha_chunk_extend,
all_gather_kv_cache_for_mha_extend,
dcp_enabled,
filter_dcp_local_kv_indices,
)
from sglang.srt.layers.quantization.fp8_utils import (
@@ -31,7 +30,7 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95,
_use_aiter_gfx95,
)
from sglang.srt.runtime_context import get_server_args
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2
_use_fp8_prefill_attn = (
@@ -287,7 +286,7 @@ class DeepseekMHAForwardMixin:
kv_a, k_pe = self._get_mla_kv_buffer_from_fp8_for_dsa(forward_batch)
else:
# BF16/FP16 path: directly fetch from cache
if dcp_enabled():
if get_parallel().dcp_enabled:
kv_a, k_pe = all_gather_kv_cache_for_mha_extend(
get_token_to_kv_pool(),
self.attn_mha,
@@ -7,7 +7,6 @@ from typing import TYPE_CHECKING, Optional
import torch
from sglang.srt.compilation.compilation_config import register_split_op
from sglang.srt.distributed.parallel_state import get_dcp_group
from sglang.srt.environ import envs
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.attention.dsa.utils import (
@@ -19,8 +18,6 @@ from sglang.srt.layers.dcp import (
all_gather_kv_cache_for_mla_extend,
all_gather_q_for_mla_decode,
cp_lse_ag_out_rs_mla,
dcp_enabled,
get_attention_dcp_world_size,
)
from sglang.srt.layers.quantization.fp8_kernel import (
fp8_dtype,
@@ -67,7 +64,7 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95,
_use_aiter_gfx95,
)
from sglang.srt.runtime_context import get_server_args
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.state_capturer.indexer_topk import (
maybe_capture_indexer_topk,
)
@@ -572,7 +569,7 @@ class DeepseekMLAForwardMixin:
)
# all_gather q_pe, q_nope_out,take tp8 as an example q_pe [B, H, ROPE_DIM], q_nope_out [B, H, NOPE_DIM] gathered to [B, H * dcp_world_size, ROPE_DIM] [B, H * dcp_world_size, NOPE_DIM] for decode batch, and all gather k_pe, k_nope for extend batch.
if dcp_enabled():
if get_parallel().dcp_enabled:
if forward_batch.forward_mode.is_decode():
# if forward_batch.forward_mode is decode, gather q
q_nope_out, q_pe = all_gather_q_for_mla_decode(
@@ -726,7 +723,10 @@ class DeepseekMLAForwardMixin:
topk_indices=topk_indices,
)
attn_output = fusion_plan.attn_output_buf
elif forward_batch.forward_mode.is_decode() and dcp_enabled():
elif (
forward_batch.forward_mode.is_decode()
and get_parallel().dcp_enabled
):
# set return_lse=True to correct attn_output
attn_output, lse = self.attn_mqa_for_dcp_decode(
q_nope_out,
@@ -800,13 +800,15 @@ class DeepseekMLAForwardMixin:
)
# correct attn_output with respect to lse from other ranks
if forward_batch.forward_mode.is_decode() and dcp_enabled():
if forward_batch.forward_mode.is_decode() and get_parallel().dcp_enabled:
attn_output = attn_output.view(
-1,
self.num_local_heads * get_attention_dcp_world_size(),
self.num_local_heads * get_parallel().attn_dcp_size,
self.kv_lora_rank,
)
attn_output = cp_lse_ag_out_rs_mla(attn_output, lse, get_dcp_group())
attn_output = cp_lse_ag_out_rs_mla(
attn_output, lse, get_parallel().dcp_group
)
attn_output = attn_output.transpose(0, 1)
attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank)
+2 -3
View File
@@ -74,7 +74,6 @@ from sglang.srt.layers.communicator_dsa_cp import (
DSACPLayerCommunicator,
maybe_prefetch_next_full_attention_kv,
)
from sglang.srt.layers.dcp import dcp_enabled, get_attention_dcp_world_size
from sglang.srt.layers.dcp.planner import (
prepare_decode_context_parallel_metadata,
)
@@ -1724,9 +1723,9 @@ class DeepseekV2AttentionMLA(
prefix=add_prefix("attn_mqa", prefix),
)
# use num_local_heads * dcp_world_size because q_nope, q_rope is all gathered from dcp ranks
if dcp_enabled():
if get_parallel().dcp_enabled:
self.attn_mqa_for_dcp_decode = RadixAttention(
self.num_local_heads * get_attention_dcp_world_size(),
self.num_local_heads * get_parallel().attn_dcp_size,
self.kv_lora_rank + self.qk_rope_head_dim,
self.scaling,
num_kv_heads=1,
+27 -3
View File
@@ -13,7 +13,7 @@
# ==============================================================================
"""A single structured accessor for process-static runtime state.
``get_parallel()`` returns a ``ParallelContext`` whose attributes tp / pp /
``get_parallel()`` returns a ``ParallelContext`` whose attributes tp / dcp / pp /
moe / attn size and rank, plus the process-group handles each delegate live to
the canonical getter in ``distributed.parallel_state`` / ``layers.dp_attention``.
Returned values are exactly what those getters return; this is a read-through
@@ -79,10 +79,13 @@ _PARALLEL_FIELDS = frozenset(
"attn_tp_rank",
"attn_cp_size",
"attn_cp_rank",
"attn_dp_size",
"attn_dp_rank",
"dcp_enabled",
"dcp_size",
"dcp_rank",
"attn_dcp_size",
"attn_dcp_rank",
"attn_dp_size",
"attn_dp_rank",
"world_group",
"tp_group",
"pp_group",
@@ -194,6 +197,27 @@ class ParallelContext:
def dcp_rank(self) -> int:
return self._v("dcp_rank", _ps().get_dcp_rank)
@property
def dcp_enabled(self) -> bool:
def getter():
if _ps().get_dcp_group_no_assert() is None:
return False
return self.dcp_size > 1
return self._v("dcp_enabled", getter)
@property
def attn_dcp_size(self) -> int:
return self._v(
"attn_dcp_size", lambda: self.dcp_size if self.dcp_enabled else 1
)
@property
def attn_dcp_rank(self) -> int:
return self._v(
"attn_dcp_rank", lambda: self.dcp_rank if self.dcp_enabled else 0
)
@property
def attn_dp_size(self) -> int:
return self._v("attn_dp_size", _dp().get_attention_dp_size)