[feature] implement dcp for deepseek_v2 (#14194)

This commit is contained in:
Augusto Yao
2026-06-25 15:15:04 -07:00
committed by GitHub
parent e4696ed62d
commit ea8f4e9f3f
20 changed files with 1770 additions and 30 deletions
@@ -34,8 +34,16 @@ from sglang.srt.layers.pooler import EmbeddingPoolerOutput
from sglang.srt.model_executor.cuda_graph_buffer_registry import (
build_eager_registry,
)
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
create_chunked_prefix_cache_kv_indices,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.forward_context import (
ForwardContext,
forward_context,
get_req_to_token_pool,
get_token_to_kv_pool,
)
from sglang.srt.model_executor.runner.base_runner import BaseRunner
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
enable_tc_piecewise_cuda_graph,
@@ -256,6 +264,23 @@ class EagerRunner(BaseRunner):
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
if forward_batch.needs_forward_metadata_init():
if hasattr(model_runner.model, "prepare_context_parallel_metadata_for_dcp"):
# prepare kv cache buffer for dcp to gather kv cache
forward_batch.attn_dcp_metadata = (
model_runner.model.prepare_context_parallel_metadata_for_dcp(
forward_batch.seq_lens,
forward_batch.extend_prefix_lens,
forward_batch.extend_prefix_lens_cpu,
forward_batch.extend_seq_lens,
forward_batch.req_pool_indices,
get_req_to_token_pool().req_to_token,
forward_batch.seq_lens_sum,
get_token_to_kv_pool().get_key_buffer(0).shape,
model_runner.kv_cache_dtype,
model_runner.device,
create_chunked_prefix_cache_kv_indices,
)
)
if hasattr(model_runner.model, "prepare_forward_batch"):
# Prepare model-specific attention metadata before planning,
# e.g. Moss-VL's prefill cross-attention custom mask.