From b0ce16d0c577e127b3f7b82ada1739277d988687 Mon Sep 17 00:00:00 2001 From: Khoa Pham Date: Sat, 23 May 2026 03:07:09 -0700 Subject: [PATCH] [CP] 1/N: Support MLA Prefill Context Parallel (#23292) Co-authored-by: Claude Opus 4.7 (1M context) --- .../attention/flashattention_backend.py | 184 +++++++++++----- python/sglang/srt/layers/communicator.py | 14 +- .../sglang/srt/layers/communicator_dsa_cp.py | 5 +- python/sglang/srt/layers/utils/cp_utils.py | 55 +++-- .../srt/model_executor/cuda_graph_runner.py | 17 +- .../sglang/srt/model_executor/model_runner.py | 7 + .../attention_backend_handler.py | 7 + .../attention_forward_methods/forward_mla.py | 3 +- python/sglang/srt/models/deepseek_nextn.py | 39 +++- python/sglang/srt/models/deepseek_v2.py | 87 ++++++-- python/sglang/srt/models/deepseek_v4.py | 1 + python/sglang/srt/models/deepseek_v4_nextn.py | 1 + .../srt/models/mistral_large_3_eagle.py | 7 + python/sglang/srt/models/qwen3_moe.py | 1 + python/sglang/srt/server_args.py | 60 +++++- .../cp/test_deepseek_v3_cp_single_node.py | 89 ++++++++ .../cp/test_deepseek_v4_flash_fp4_b200_cp.py | 85 ++++++++ .../{4-gpu-models => cp}/test_qwen3_30b.py | 2 +- .../kernels/test_cp_prefix_len_fa3_parity.py | 143 ++++++++++++ .../kernels/test_mla_cp_fa3_parity.py | 203 ++++++++++++++++++ .../test_deepseek_v4_flash_fp4_b200.py | 51 +---- 21 files changed, 900 insertions(+), 161 deletions(-) create mode 100644 test/registered/cp/test_deepseek_v3_cp_single_node.py create mode 100644 test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py rename test/registered/{4-gpu-models => cp}/test_qwen3_30b.py (98%) create mode 100644 test/registered/kernels/test_cp_prefix_len_fa3_parity.py create mode 100644 test/registered/kernels/test_mla_cp_fa3_parity.py diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index aa8eb009e..2853cdd69 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -508,6 +508,25 @@ class FlashAttentionBackend(AttentionBackend): metadata.cu_seqlens_k = torch.nn.functional.pad( torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) ) + + # MLA/MHA CP: prepare_mlp_sync_batch pads extend tokens up to + # lcm(attn_tp_size, attn_cp_size), so cache_seqlens_cp can exceed + # seq_lens_cpu.max(). Widen page_table by the pad delta to keep + # FA3's causal reads in-bounds; widened columns index KV slot 0 + # (req_to_token is zero-init) and outputs for padding queries are + # discarded downstream. + if ( + self.attn_cp_size > 1 + and forward_batch.global_num_tokens_cpu is not None + and forward_batch.extend_num_tokens is not None + and forward_batch.extend_seq_lens_cpu is not None + ): + padded_extend = int(forward_batch.extend_num_tokens) + real_extend = int(sum(forward_batch.extend_seq_lens_cpu)) + pad_delta = padded_extend - real_extend + if pad_delta > 0: + metadata.max_seq_len_k += pad_delta + metadata.page_table = self.req_to_token_pool.req_to_token[ forward_batch.req_pool_indices, : metadata.max_seq_len_k ] @@ -627,36 +646,43 @@ class FlashAttentionBackend(AttentionBackend): k_rope: Optional[torch.Tensor] = None, sinks: Optional[torch.Tensor] = None, ): + is_cp_mode = ( + forward_batch.forward_mode.is_context_parallel_extend() + and forward_batch.attn_cp_metadata is not None + and self.attn_cp_size > 1 + ) + if k is not None: assert v is not None - is_cp_mode = ( - forward_batch.forward_mode.is_context_parallel_extend() - and forward_batch.attn_cp_metadata is not None - and self.attn_cp_size > 1 - ) - - if save_kv_cache and not is_cp_mode and not self.fa_skip_kv_cache: + if save_kv_cache and not self.fa_skip_kv_cache: cache_loc = ( forward_batch.out_cache_loc if not layer.is_cross_attention else forward_batch.encoder_out_cache_loc ) - if not self.use_mla: - self.token_to_kv_pool.set_kv_buffer( - layer, cache_loc, k, v, layer.k_scale, layer.v_scale - ) - else: + if self.use_mla: + # MLA: under CP, k and k_rope arrive full-sequence + # (rebuild_cp_kv_cache ran upstream in + # forward_absorb_prepare); rank-local otherwise. + # out_cache_loc is never zigzag-split, so the write + # lands in the right slots on every rank in either case. self.token_to_kv_pool.set_mla_kv_buffer( layer, cache_loc, k, k_rope, ) - if is_cp_mode: - cp_allgather_and_save_kv_cache( - forward_batch, layer, k, v, self.attn_cp_size - ) + elif is_cp_mode: + # Dense-MHA CP: k, v are still rank-local; backend + # all-gathers and writes to the per-rank pool. + cp_allgather_and_save_kv_cache( + forward_batch, layer, k, v, self.attn_cp_size + ) + else: + self.token_to_kv_pool.set_kv_buffer( + layer, cache_loc, k, v, layer.k_scale, layer.v_scale + ) # Use precomputed metadata across all layers metadata = self.forward_metadata @@ -974,57 +1000,103 @@ class FlashAttentionBackend(AttentionBackend): q_nope = q_all[:, :, : layer.v_head_dim] q_rope = q_all[:, :, layer.v_head_dim :] - result = flash_attn_with_kvcache( - q=q_rope, - k_cache=k_rope_cache, - v_cache=c_kv_cache, - qv=q_nope, - page_table=page_table, - cache_seqlens=cache_seqlens, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k_new=cu_seqlens_k if not use_local_attn else None, - max_seqlen_q=max_seqlen_q, - softmax_scale=layer.scaling, - causal=False if use_cascade_attn else causal, - softcap=layer.logit_cap, - k_descale=k_descale, - v_descale=v_descale, - return_softmax_lse=use_cascade_attn, - num_splits=self.num_splits, - ver=self.fa_impl_ver, - ) - if use_cascade_attn: - o, softmax_lse, *rest = result - o_expand, softmax_lse_expand, *rest_expand = ( - flash_attn_with_kvcache( - q=q_rope, + if is_cp_mode: + # MLA CP: q is rank-local zigzag-split; run the + # absorbed-MLA kernel twice (prev/next halves) against + # the full latent KV pool (which rebuild_cp_kv_cache + # populated upstream) via cp_attn_forward_extend. + # Concat q_nope + q_rope along dim=-1 so the wrapper's + # chunk(2, dim=0) keeps their alignment; split back + # inside the closure. + assert ( + not use_cascade_attn + ), "Cascade attention under MLA CP is not supported in v1." + q_fused = torch.cat([q_nope, q_rope], dim=-1) + + def _mla_cp_attn( + q_chunk, + cu_seqlens_q_cp, + cache_seqlens_cp, + max_seqlen_q_cp, + ): + q_nope_chunk = q_chunk[..., : layer.v_head_dim] + q_rope_chunk = q_chunk[..., layer.v_head_dim :] + return flash_attn_with_kvcache( + q=q_rope_chunk, + qv=q_nope_chunk, k_cache=k_rope_cache, v_cache=c_kv_cache, - qv=q_nope, - page_table=self.forward_metadata_spec_decode_expand.page_table, - cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32, - cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q, - cu_seqlens_k_new=self.forward_metadata_spec_decode_expand.cu_seqlens_k, - max_seqlen_q=self.forward_metadata_spec_decode_expand.max_seq_len_q, + page_table=page_table, + cache_seqlens=cache_seqlens_cp, + cu_seqlens_q=cu_seqlens_q_cp, + cu_seqlens_k_new=( + cu_seqlens_k if not use_local_attn else None + ), + max_seqlen_q=max_seqlen_q_cp, softmax_scale=layer.scaling, - causal=False, - window_size=window_size, + causal=causal, softcap=layer.logit_cap, k_descale=k_descale, v_descale=v_descale, - return_softmax_lse=True, num_splits=self.num_splits, ver=self.fa_impl_ver, ) - ) - o, _ = merge_state_v2_wrapper( - o, - softmax_lse.T.contiguous(), - o_expand, - softmax_lse_expand.T.contiguous(), + + o = cp_attn_forward_extend( + forward_batch, q_fused, self.device, _mla_cp_attn ) else: - o = result + result = flash_attn_with_kvcache( + q=q_rope, + k_cache=k_rope_cache, + v_cache=c_kv_cache, + qv=q_nope, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k_new=cu_seqlens_k if not use_local_attn else None, + max_seqlen_q=max_seqlen_q, + softmax_scale=layer.scaling, + causal=False if use_cascade_attn else causal, + softcap=layer.logit_cap, + k_descale=k_descale, + v_descale=v_descale, + return_softmax_lse=use_cascade_attn, + num_splits=self.num_splits, + ver=self.fa_impl_ver, + ) + if use_cascade_attn: + o, softmax_lse, *rest = result + o_expand, softmax_lse_expand, *rest_expand = ( + flash_attn_with_kvcache( + q=q_rope, + k_cache=k_rope_cache, + v_cache=c_kv_cache, + qv=q_nope, + page_table=self.forward_metadata_spec_decode_expand.page_table, + cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32, + cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q, + cu_seqlens_k_new=self.forward_metadata_spec_decode_expand.cu_seqlens_k, + max_seqlen_q=self.forward_metadata_spec_decode_expand.max_seq_len_q, + softmax_scale=layer.scaling, + causal=False, + window_size=window_size, + softcap=layer.logit_cap, + k_descale=k_descale, + v_descale=v_descale, + return_softmax_lse=True, + num_splits=self.num_splits, + ver=self.fa_impl_ver, + ) + ) + o, _ = merge_state_v2_wrapper( + o, + softmax_lse.T.contiguous(), + o_expand, + softmax_lse_expand.T.contiguous(), + ) + else: + o = result return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 09f38265f..2efc1e775 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -65,6 +65,10 @@ from sglang.srt.layers.moe import ( should_use_dp_reduce_scatterv, should_use_flashinfer_cutlass_moe_fp4_allgather, ) +from sglang.srt.layers.utils.cp_utils import ( + is_mla_prefill_cp_enabled, + mla_use_prefill_cp, +) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -202,7 +206,7 @@ class ScatterMode(Enum): @staticmethod def model_input_output(): """The scatter mode for model forward pass input and output data""" - if is_dsa_enable_prefill_cp(): + if is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled(): return ScatterMode.SCATTERED return ScatterMode.TP_ATTN_FULL @@ -379,8 +383,10 @@ class LayerScatterModes: or should_use_flashinfer_cutlass_moe_fp4_allgather() ): return ScatterMode.SCATTERED - # DSA CP doesn't support MOE_FULL yet; fall back to FULL - if is_enable_moe_cp_allgather() and not is_dsa_enable_prefill_cp(): + # DSA CP and MLA CP both don't support MOE_FULL yet; fall back to FULL. + if is_enable_moe_cp_allgather() and not ( + is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled() + ): return ScatterMode.MOE_FULL return ScatterMode.FULL else: @@ -709,7 +715,7 @@ class LayerCommunicator: return True if forward_batch.dp_padding_mode.is_max_len(): return True - if dsa_use_prefill_cp(forward_batch): + if dsa_use_prefill_cp(forward_batch) or mla_use_prefill_cp(forward_batch): return True if get_attn_tp_context().input_scattered and not self.is_last_layer: return True diff --git a/python/sglang/srt/layers/communicator_dsa_cp.py b/python/sglang/srt/layers/communicator_dsa_cp.py index 20b220278..46346ca3b 100644 --- a/python/sglang/srt/layers/communicator_dsa_cp.py +++ b/python/sglang/srt/layers/communicator_dsa_cp.py @@ -37,6 +37,7 @@ from sglang.srt.layers.dp_attention import ( get_attention_cp_group, get_local_dp_buffer, ) +from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -152,7 +153,7 @@ class DSACPCommunicateWithAllReduceAndLayerNormFn( hidden_states, residual = layernorm(hidden_states, residual) # for prefill: attn tp scattered -> full # for decode: attn tp full -> full - if dsa_use_prefill_cp(forward_batch): + if dsa_use_prefill_cp(forward_batch) or mla_use_prefill_cp(forward_batch): assert context.attn_dp_size == 1 hidden_states, local_hidden_states = ( get_local_dp_buffer(get_attention_cp_group()), @@ -205,7 +206,7 @@ class DSACPCommunicateSummableTensorPairFn(CommunicateSummableTensorPairFn): ): # for prefill: full -> attn tp scattered # for decode: full -> attn tp full - if dsa_use_prefill_cp(forward_batch): + if dsa_use_prefill_cp(forward_batch) or mla_use_prefill_cp(forward_batch): assert context.attn_dp_size == 1 input_hidden_states = hidden_states hidden_states = hidden_states.tensor_split(context.attn_cp_size)[ diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index e998dfc67..b87225dc7 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -51,19 +51,41 @@ def is_prefill_cp_in_seq_split(): ) +def is_mla_prefill_cp_enabled() -> bool: + sa = get_global_server_args() + return sa.enable_prefill_context_parallel and sa.use_mla_backend + + +def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None): + if mla_enable_prefill_cp is None: + mla_enable_prefill_cp = is_mla_prefill_cp_enabled() + return ( + forward_batch.attn_cp_metadata is not None + and mla_enable_prefill_cp + and forward_batch.forward_mode.is_context_parallel_extend() + ) + + def can_cp_split(seq_len: int, cp_size: int, forward_batch): - # CP metadata (zigzag split) only supports batch=1 for now. + from sglang.srt.model_executor.forward_batch_info import ForwardMode + + # TODO current just support prefill batch=1 and len(input_ids) > self.cp_size * 2 + # Note: (self.cp_size * 2) To achieve load balancing for seq computation, + # the seq data needs to be divided and recombined at twice the size of cp_size. cur_cp_seq_len = seq_len // (cp_size * 2) - if ( + return ( cur_cp_seq_len != 0 and cp_size > 1 + # prepare_context_parallel_metadata hard-codes bs_per_cp_group = 1; + # guard explicitly to avoid silent mis-partitioning under continuous batching. + # TODO: remove this guard once we support multi-batch-cp-split + and forward_batch.batch_size == 1 and forward_batch.forward_mode.is_context_parallel_extend() + # is_context_parallel_extend() returns True for MIXED (prefill+decode + # in one step), but the zigzag split only makes sense on pure extend. + and forward_batch.forward_mode != ForwardMode.MIXED and is_prefill_context_parallel_enabled() - and forward_batch.seq_lens_cpu.shape[0] == 1 - ): - return True - else: - return False + ) def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor): @@ -395,6 +417,7 @@ def prepare_context_parallel_metadata( cp_rank, cp_size, seqs_len, + extend_lens, ): from sglang.srt.layers.attention.dsa.utils import ( is_dsa_prefill_cp_round_robin_split, @@ -449,18 +472,12 @@ def prepare_context_parallel_metadata( bs_per_cp_group = 1 kv_len_origin = kv_len - # Derive prefix offset from the full sequence length on CPU. - # NOTE: forward_batch.seq_lens_cpu includes cached prefix + extend tokens. - # In CP we only split the extend tokens, but cache_seqlens passed to FA must - # include the cached prefix. - prefix_len = 0 - try: - if seqs_len is not None and len(seqs_len) == 1: - prefix_len = int(seqs_len[0]) - int(kv_len_origin.item()) - if prefix_len < 0: - prefix_len = 0 - except Exception: - prefix_len = 0 + # Derive prefix offset from unpadded CPU tensors. Both `seqs_len` and `extend_lens` are unpadded by the caller + # Using the padded `kv_len` here would undercount `prefix_len` by the padding amount and shift the FA causal horizon. + assert ( + len(seqs_len) == 1 and len(extend_lens) == 1 + ), "Prefill Context Parallel only supports batch_size == 1 for now" + prefix_len = max(0, int(seqs_len[0]) - int(extend_lens[0])) # get zigzag index cp_segment_num = cp_size * 2 seq_per_batch = kv_len // cp_segment_num # seq_len for each batch and segment diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 65e38c2d0..63aff21f8 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -56,6 +56,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer from sglang.srt.layers.moe.utils import get_deepep_mode, get_moe_a2a_backend from sglang.srt.layers.utils import MultiPlatformOp +from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, ForwardBatch, @@ -567,7 +568,15 @@ class CudaGraphRunner: self.attn_tp_size = get_attention_tp_size() self.attn_tp_rank = get_attention_tp_rank() - self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + # True if a DSACPLayerCommunicator-style prefill-CP flavor is active + # (DSA or MLA). These flavors feed a zigzag-split rank-local layout + # into the runner; MHA-arch prefill CP (Qwen3/Qwen2 MoE via PR + # #18233) uses the plain LayerCommunicator with an attn_tp-replicated + # layout and is intentionally excluded so the attn_tp-local + # num_token_non_padded adjustment still runs for it. + self.enable_prefill_cp = ( + is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled() + ) self.deepep_adapter = DeepEPCudaGraphRunnerAdapter() @@ -928,7 +937,7 @@ class CudaGraphRunner: if ( enable_num_token_non_padded() and self.require_gathered_buffer - and not self.dsa_enable_prefill_cp + and not self.enable_prefill_cp ): local = compute_local_num_token_non_padded( global_num_token_non_padded=buffers.num_token_non_padded, @@ -1191,7 +1200,9 @@ class CudaGraphRunner: seq_len_fill_value=self.seq_len_fill_value, require_gathered_buffer=self.require_gathered_buffer, num_tokens_per_bs=self.num_tokens_per_bs, - dsa_enable_prefill_cp=self.dsa_enable_prefill_cp, + # Parameter name retained for API stability; semantically this is + # "any prefill-CP flavor enabled" (DSA CP or MLA CP). + dsa_enable_prefill_cp=self.enable_prefill_cp, enable_num_token_non_padded_flag=enable_num_token_non_padded(), pp_proxy_tensors=pp_proxy_tensors, ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 3d0423d9b..eb0a633a3 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -126,6 +126,7 @@ from sglang.srt.layers.pooler import EmbeddingPoolerOutput from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype from sglang.srt.layers.sampler import create_sampler from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model +from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled from sglang.srt.lora.lora_manager import LoRAManager from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value @@ -3285,11 +3286,17 @@ class ModelRunner(ModelRunnerKVCacheMixin): forward_batch.prepare_attn_tp_scatter_input(self) # Normalize num_token_non_padded to be local to this attention TP rank if needed. + # The skip is scoped to DSACPLayerCommunicator-style CP (DSA, MLA): those + # flavors already feed a zigzag-split rank-local layout whose token count + # should not be further divided by attn_tp_size. MHA-arch prefill CP + # (Qwen3/Qwen2 MoE) keeps the attn_tp-replicated layout and wants the + # adjustment to run — see docs/design/prefill-cp-mla.md §Phase 5. if ( forward_batch.num_token_non_padded is not None and forward_batch.global_num_tokens_gpu is not None and require_gathered_buffer(self.server_args) and not is_dsa_enable_prefill_cp() + and not is_mla_prefill_cp_enabled() ): forward_batch.adjust_num_token_non_padded_for_attn_tp( server_args=self.server_args, diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index e4cac55f9..6dcf9bc79 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -1,5 +1,6 @@ from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph from sglang.srt.layers.attention.tbo_backend import TboAttnBackend +from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import ( AttnForwardMethod, @@ -74,6 +75,12 @@ def _handle_attention_backend(attn, forward_batch, backend_name): if is_in_piecewise_cuda_graph(): return AttnForwardMethod.MLA + # MLA prefill CP forces absorbed MLA regardless of prefix length: the + # CP path gathers latent KV via rebuild_cp_kv_cache and feeds the + # backend's absorbed-MLA kernel. + if mla_use_prefill_cp(forward_batch): + return _dispatch_mla_subtype(attn, forward_batch) + sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) disable_ragged = ( backend_name in ["flashinfer", "flashmla"] diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 904f0ee03..364f065b4 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -13,6 +13,7 @@ from sglang.srt.layers.quantization.fp8_kernel import ( per_tensor_quant_mla_fp8, per_token_group_quant_mla_deep_gemm_masked_fp8, ) +from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp from sglang.srt.lora.deepseek_mla_correction import ( apply_q_correction as apply_kv_b_lora_q_correction, ) @@ -380,7 +381,7 @@ class DeepseekMLAForwardMixin: ): q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) - if dsa_use_prefill_cp(forward_batch): + if dsa_use_prefill_cp(forward_batch) or mla_use_prefill_cp(forward_batch): # support allgather+rerrange k_nope, k_pe = self.rebuild_cp_kv_cache( latent_cache, forward_batch, k_nope, k_pe diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 59d26282e..1e1697252 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -43,9 +43,12 @@ from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization import Fp8Config from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.utils.cp_utils import ( + can_cp_split, cp_all_gather_rerange_output, cp_split_and_rebuild_data, cp_split_and_rebuild_position, + is_mla_prefill_cp_enabled, + mla_use_prefill_cp, prepare_context_parallel_metadata, ) from sglang.srt.layers.vocab_parallel_embedding import ( @@ -136,6 +139,14 @@ class DeepseekModelNextN(nn.Module): layer_name = "layers." + str(config.num_hidden_layers) self.quant_config = quant_config + self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.mla_enable_prefill_cp = ( + is_mla_prefill_cp_enabled() and not is_deepseek_dsa(config) + ) + if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: + self.cp_size = get_attention_cp_size() + else: + self.cp_size = None self.decoder = DeepseekV2DecoderLayer( config, 0, @@ -144,15 +155,12 @@ class DeepseekModelNextN(nn.Module): is_nextn=True, prefix=add_prefix(layer_name, prefix), alt_stream=self.alt_stream, + dsa_enable_prefill_cp=self.dsa_enable_prefill_cp, + mla_enable_prefill_cp=self.mla_enable_prefill_cp, ) self.shared_head = nn.Module() self.shared_head.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() - if self.dsa_enable_prefill_cp: - self.cp_size = get_attention_cp_size() - else: - self.cp_size = None def forward( self, @@ -193,7 +201,9 @@ class DeepseekModelNextN(nn.Module): else: hidden_states = self.eh_proj(eh_input) - if dsa_use_prefill_cp(forward_batch, self.dsa_enable_prefill_cp): + if dsa_use_prefill_cp( + forward_batch, self.dsa_enable_prefill_cp + ) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp): hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) residual = None @@ -212,7 +222,9 @@ class DeepseekModelNextN(nn.Module): else: hidden_states = self.shared_head.norm(hidden_states) - if dsa_use_prefill_cp(forward_batch, self.dsa_enable_prefill_cp): + if dsa_use_prefill_cp( + forward_batch, self.dsa_enable_prefill_cp + ) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp): # allgather + rerrange hidden_states = cp_all_gather_rerange_output( hidden_states, @@ -250,7 +262,8 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): self.determine_num_fused_shared_experts("DeepseekV3ForCausalLMNextN") self.use_dsa = is_deepseek_dsa(config) self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() - if self.dsa_enable_prefill_cp: + self.mla_enable_prefill_cp = is_mla_prefill_cp_enabled() and not self.use_dsa + if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: self.cp_rank = get_attention_cp_rank() self.cp_size = get_attention_cp_size() else: @@ -298,6 +311,16 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): self.cp_rank, self.cp_size, forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, + ) + elif self.mla_enable_prefill_cp: + if can_cp_split(len(input_ids), self.cp_size, forward_batch): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len(input_ids), + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, ) hidden_states = self.model(input_ids, positions, forward_batch) return self.logits_processor( diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 955edb980..71e233f40 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -123,9 +123,12 @@ from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.rotary_embedding import get_rope_wrapper from sglang.srt.layers.utils import PPMissingLayer from sglang.srt.layers.utils.cp_utils import ( + can_cp_split, cp_all_gather_rerange_output, cp_split_and_rebuild_data, cp_split_and_rebuild_position, + is_prefill_context_parallel_enabled, + mla_use_prefill_cp, prepare_context_parallel_metadata, ) from sglang.srt.layers.vocab_parallel_embedding import ( @@ -339,6 +342,8 @@ class MoEGate(nn.Module): is_nextn: bool = False, is_hash_moe: bool = False, is_deepseek_v4: bool = False, + dsa_enable_prefill_cp: bool = False, + mla_enable_prefill_cp: bool = False, ): super().__init__() self.is_nextn = is_nextn @@ -368,7 +373,9 @@ class MoEGate(nn.Module): self.e_score_correction_bias = None if _is_cpu and _is_cpu_amx_available: self.quant_method = PackWeightMethod(weight_names=["weight"]) - self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.use_dsa = is_deepseek_dsa(config) + self.dsa_enable_prefill_cp = dsa_enable_prefill_cp + self.mla_enable_prefill_cp = mla_enable_prefill_cp def forward( self, @@ -390,7 +397,10 @@ class MoEGate(nn.Module): if ( not self.is_deepseek_v4 and forward_batch is not None - and dsa_use_prefill_cp(forward_batch) + and ( + dsa_use_prefill_cp(forward_batch, self.dsa_enable_prefill_cp) + or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp) + ) ): logits = F.linear(hidden_states, self.weight, None) else: @@ -442,6 +452,8 @@ class DeepseekV2MoE(nn.Module): alt_stream: Optional[torch.cuda.Stream] = None, is_nextn: bool = False, is_deepseek_v4: bool = False, + dsa_enable_prefill_cp: bool = False, + mla_enable_prefill_cp: bool = False, ): super().__init__() self.tp_size = get_tensor_model_parallel_world_size() @@ -503,6 +515,8 @@ class DeepseekV2MoE(nn.Module): is_nextn=is_nextn, is_hash_moe=self.is_hash, is_deepseek_v4=is_deepseek_v4, + dsa_enable_prefill_cp=dsa_enable_prefill_cp, + mla_enable_prefill_cp=mla_enable_prefill_cp, ) # scaling factor for fused shared experts on AMD-platform. @@ -1340,6 +1354,8 @@ class DeepseekV2AttentionMLA( alt_stream: Optional[torch.cuda.Stream] = None, skip_rope: bool = False, is_nextn: bool = False, + dsa_enable_prefill_cp: bool = False, + mla_enable_prefill_cp: bool = False, ) -> None: super().__init__() self.layer_id = layer_id @@ -1354,11 +1370,14 @@ class DeepseekV2AttentionMLA( attn_tp_rank = get_attention_tp_rank() attn_tp_size = get_attention_tp_size() self.use_dsa = is_deepseek_dsa(config) - self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.dsa_enable_prefill_cp = dsa_enable_prefill_cp + self.mla_enable_prefill_cp = mla_enable_prefill_cp if self.dsa_enable_prefill_cp: assert self.use_dsa, "CP currently only supports deepseek v3.2 model" - # cp reuse the attn_tp comm group but need to duplicate the weights - if self.dsa_enable_prefill_cp and self.use_dsa: + # cp reuses the attn_tp comm group but needs to duplicate the weights; + # store cp_size whenever either CP flavor is active so rebuild_cp_kv_cache + # and the FA3 MLA wrapper can reach it on the dense MLA path too. + if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: self.cp_size = get_attention_cp_size() self.num_heads = num_heads assert num_heads % attn_tp_size == 0 @@ -1793,6 +1812,8 @@ class DeepseekV2DecoderLayer(nn.Module): is_nextn: bool = False, prefix: str = "", alt_stream: Optional[torch.cuda.Stream] = None, + dsa_enable_prefill_cp: bool = False, + mla_enable_prefill_cp: bool = False, ) -> None: super().__init__() self.hidden_size = config.hidden_size @@ -1809,7 +1830,8 @@ class DeepseekV2DecoderLayer(nn.Module): self.speculative_algorithm = SpeculativeAlgorithm.from_string( get_global_server_args().speculative_algorithm ) - self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.dsa_enable_prefill_cp = dsa_enable_prefill_cp + self.mla_enable_prefill_cp = mla_enable_prefill_cp self.layer_id = layer_id self.is_nextn = is_nextn self.self_attn = DeepseekV2AttentionMLA( @@ -1832,6 +1854,8 @@ class DeepseekV2DecoderLayer(nn.Module): prefix=add_prefix("self_attn", prefix), alt_stream=alt_stream, is_nextn=is_nextn, + dsa_enable_prefill_cp=dsa_enable_prefill_cp, + mla_enable_prefill_cp=mla_enable_prefill_cp, ) if not hasattr(config, "q_lora_rank") and envs.SGLANG_USE_AG_AFTER_QLORA.get(): raise ValueError( @@ -1858,6 +1882,8 @@ class DeepseekV2DecoderLayer(nn.Module): layer_id=self.layer_id, alt_stream=alt_stream, is_nextn=is_nextn, + dsa_enable_prefill_cp=dsa_enable_prefill_cp, + mla_enable_prefill_cp=mla_enable_prefill_cp, ) else: if enable_moe_dense_fully_dp(): @@ -1882,7 +1908,10 @@ class DeepseekV2DecoderLayer(nn.Module): self._gfx95_quant_format = self._detect_gfx95_quant_format() - if self.dsa_enable_prefill_cp: + if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: + # DSACPLayerCommunicator is flavor-agnostic; its internal gates + # read both dsa_use_prefill_cp and mla_use_prefill_cp. The rename + # to CPLayerCommunicator is deferred to a cleanup PR. self.layer_communicator = DSACPLayerCommunicator( layer_scatter_modes=self.layer_scatter_modes, input_layernorm=self.input_layernorm, @@ -1998,7 +2027,10 @@ class DeepseekV2DecoderLayer(nn.Module): gemm_output_zero_allocator, ) - if not self.dsa_enable_prefill_cp and should_allreduce_fusion: + if ( + not (self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp) + and should_allreduce_fusion + ): hidden_states._sglang_needs_allreduce_fusion = True if not should_allreduce_fusion: @@ -2096,7 +2128,10 @@ class DeepseekV2Model(nn.Module): self.first_k_dense_replace = config.first_k_dense_replace self.pp_group = get_pp_group() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() - if self.dsa_enable_prefill_cp: + self.mla_enable_prefill_cp = ( + is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) + ) + if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: self.cp_size = get_attention_cp_size() else: self.cp_size = None @@ -2129,6 +2164,8 @@ class DeepseekV2Model(nn.Module): quant_config=quant_config, prefix=prefix, alt_stream=self.alt_stream, + dsa_enable_prefill_cp=self.dsa_enable_prefill_cp, + mla_enable_prefill_cp=self.mla_enable_prefill_cp, ), pp_rank=self.pp_group.rank_in_group, pp_size=self.pp_group.world_size, @@ -2255,7 +2292,9 @@ class DeepseekV2Model(nn.Module): else None ) - if dsa_use_prefill_cp(forward_batch): + if dsa_use_prefill_cp( + forward_batch, self.dsa_enable_prefill_cp + ) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp): if self.pp_group.is_first_rank: hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) @@ -2340,7 +2379,10 @@ class DeepseekV2Model(nn.Module): else: hidden_states, _ = self.norm(hidden_states, residual) - if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch): + if self.pp_group.is_last_rank and ( + dsa_use_prefill_cp(forward_batch, self.dsa_enable_prefill_cp) + or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp) + ): # allgather + rerrange hidden_states = cp_all_gather_rerange_output( hidden_states, @@ -2418,7 +2460,10 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self.capture_aux_hidden_states = False self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() - if self.dsa_enable_prefill_cp: + self.mla_enable_prefill_cp = ( + is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) + ) + if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: self.cp_rank = get_attention_cp_rank() self.cp_size = get_attention_cp_size() else: @@ -2510,15 +2555,29 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): input_embeds: torch.Tensor = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> torch.Tensor: + # Minor fix for multi-modal model: input_ids is None + len_input_ids = ( + input_ids.shape[0] if input_ids is not None else input_embeds.shape[0] + ) if self.dsa_enable_prefill_cp: if can_dsa_cp_split( - len(input_ids), self.cp_size, self.use_dsa, forward_batch + len_input_ids, self.cp_size, self.use_dsa, forward_batch ): forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( - len(input_ids), + len_input_ids, self.cp_rank, self.cp_size, forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, + ) + elif self.mla_enable_prefill_cp: + if can_cp_split(len_input_ids, self.cp_size, forward_batch): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len_input_ids, + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, ) with get_attn_tp_context().maybe_input_scattered(forward_batch): diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 4b062a119..ecdfaabf2 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -1286,6 +1286,7 @@ class DeepseekV4ForCausalLM(nn.Module): self.cp_rank, self.cp_size, forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, ) if is_dsa_prefill_cp_round_robin_split(): attn_backend = get_attn_backend() diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index 1260320a8..6b5c89e50 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -249,6 +249,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM): self.cp_rank, self.cp_size, forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, ) if is_dsa_prefill_cp_round_robin_split(): attn_backend = get_attn_backend() diff --git a/python/sglang/srt/models/mistral_large_3_eagle.py b/python/sglang/srt/models/mistral_large_3_eagle.py index 65ffbd820..ae487c52c 100644 --- a/python/sglang/srt/models/mistral_large_3_eagle.py +++ b/python/sglang/srt/models/mistral_large_3_eagle.py @@ -7,11 +7,13 @@ import torch from torch import nn from transformers import PretrainedConfig +from sglang.srt.configs.model_config import is_deepseek_dsa from sglang.srt.distributed import get_pp_group from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import RowParallelLinear from sglang.srt.layers.quantization.base_config import QuantizationConfig +from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV2Model @@ -36,6 +38,9 @@ class MistralLarge3EagleModel(DeepseekV2Model): assert get_pp_group().world_size == 1 self.pp_group = get_pp_group() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.mla_enable_prefill_cp = ( + is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) + ) self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -50,6 +55,8 @@ class MistralLarge3EagleModel(DeepseekV2Model): prefix=add_prefix(prefix, f"layers.{i}"), quant_config=quant_config, layer_id=i, + dsa_enable_prefill_cp=self.dsa_enable_prefill_cp, + mla_enable_prefill_cp=self.mla_enable_prefill_cp, ) for i in range(self.config.num_hidden_layers) ] diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 0887bfc3f..1e31930fb 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -1005,6 +1005,7 @@ class Qwen3MoeForCausalLM(nn.Module): self.attn_cp_rank, self.attn_cp_size, forward_batch.seq_lens_cpu.tolist(), + extend_lens=forward_batch.extend_seq_lens_cpu, ) hidden_states = self.model( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index e5cd400be..2b65605d1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1825,10 +1825,20 @@ class ServerArgs: assert ( self.tp_size <= 8 ), "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues." + # Note(kpham-sgl): Keep attn_tp_size == 1 under DSA CP. + # DSACPLayerCommunicator does not all-reduce attention-TP + # partial o_proj outputs before replicated dense FFNs. self.attn_cp_size = self.tp_size // self.dp_size - + self.disable_piecewise_cuda_graph = True logger.warning( - f"Enable Context Parallel opt for deeeseekv3.2-DSA, Setting dp_size == {self.dp_size} and moe_dense_tp_size == {self.moe_dense_tp_size}, ep_size == {self.ep_size}, tp_size == {self.tp_size}, kv_cache_dtype == {self.kv_cache_dtype}, moe_a2a_backend {self.moe_a2a_backend} " + f"Enable DSA Context Parallel opt, " + f"Setting dp_size == {self.dp_size} and " + f"moe_dense_tp_size == {self.moe_dense_tp_size}, " + f"ep_size == {self.ep_size}, " + f"tp_size == {self.tp_size}, " + f"kv_cache_dtype == {self.kv_cache_dtype}, " + f"moe_a2a_backend {self.moe_a2a_backend}, " + f"disable_piecewise_cuda_graph=True" ) else: # Pure TP and partial DP Attention mode is active for DSA, logging a warning @@ -1870,7 +1880,7 @@ class ServerArgs: ), "CP is only supported for prefill when PD disaggregation, please remove --enable-dsa-prefill-context-parallel." else: - # DeepSeek V3/R1/V3.1 + # DeepSeek V3/R1/V3.1 and Kimi K2.5 if not self.disable_piecewise_cuda_graph: logger.info("Piecewise CUDA graph is enabled, use MLA for prefill.") @@ -1885,6 +1895,37 @@ class ServerArgs: "Use trtllm_mla as attention backend on sm100 for DeepseekV3ForCausalLM" ) + # MLA prefill CP auto-config. Mirrors the NSA CP block above + # (minus the in-seq/round-robin mode split, which MLA CP does not support) + if self.enable_prefill_context_parallel and self.use_mla_backend(): + logger.warning( + "MLA prefill context parallel is still experimental. " + "Verified on Hopper with the fa3 backend." + ) + self.enable_dp_attention = True + # TODO(kpham-sgl) Supports moe_dense_tp_size != 1. + self.moe_dense_tp_size = 1 + self.moe_a2a_backend = "deepep" + self.ep_size = self.tp_size + logger.warning( + "For MLA CP, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, batch_size == 1" + ) + # FIXME(kpham-sgl): Keep attn_tp_size == 1 under MLA CP. + # DSACPLayerCommunicator does not all-reduce attention-TP + # partial o_proj outputs before replicated dense FFNs. + self.attn_cp_size = self.tp_size // self.dp_size + self.disable_piecewise_cuda_graph = True + logger.warning( + f"Enable Context Parallel opt for MLA, " + f"Setting dp_size == {self.dp_size} and " + f"attn_cp_size == {self.attn_cp_size}, " + f"moe_dense_tp_size == {self.moe_dense_tp_size}, " + f"ep_size == {self.ep_size}, " + f"tp_size == {self.tp_size}, " + f"moe_a2a_backend {self.moe_a2a_backend}, " + f"disable_piecewise_cuda_graph=True" + ) + # Set moe backend for DeepSeek if is_sm100_supported(): quant_method = get_quantization_config(hf_config) @@ -3039,6 +3080,19 @@ class ServerArgs: ) def _handle_context_parallelism(self): + if ( + self.enable_prefill_context_parallel + and self.enable_dsa_prefill_context_parallel + ): + raise ValueError( + "--enable-prefill-context-parallel and " + "--enable-nsa-prefill-context-parallel are mutually " + "exclusive. Use --enable-nsa-prefill-context-parallel for " + "DeepSeek V3.2 (NSA) models and " + "--enable-prefill-context-parallel for MLA-based models " + "(DeepSeek V3/R1, Kimi K2.5) or MHA/GQA-based models." + ) + if self.attn_cp_size > 1: # The tp_size is the world size, not the real tensor parallel size assert ( diff --git a/test/registered/cp/test_deepseek_v3_cp_single_node.py b/test/registered/cp/test_deepseek_v3_cp_single_node.py new file mode 100644 index 000000000..74e0258d5 --- /dev/null +++ b/test/registered/cp/test_deepseek_v3_cp_single_node.py @@ -0,0 +1,89 @@ +import unittest +from types import SimpleNamespace + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.run_eval import run_eval +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + is_in_ci, + popen_launch_server, + write_github_step_summary, +) + +register_cuda_ci(est_time=500, stage="extra-b", runner_config="deepep-8-gpu-h200") + +DEEPSEEK_V3_MODEL_PATH = "deepseek-ai/DeepSeek-V3-0324" + +# Matches the non-CP DSv3 production baseline in +# ``test_deepseek_v3_basic.py`` / ``test_deepseek_v3_mtp.py``. Pinning +# MLA CP to the same threshold makes this test double as a regression +# gate against the known production accuracy. +GSM8K_ACCURACY_THRESHOLD = 0.935 + + +class TestDeepseekV3CPInSeqSplit(CustomTestCase): + """tp=8, dp=2, attn-cp=4 — DP attention + DeepEP MoE + MLA CP.""" + + @classmethod + def setUpClass(cls): + cls.model = DEEPSEEK_V3_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + other_args = [ + "--trust-remote-code", + "--tp", + "8", + "--dp", + "2", + "--enable-prefill-context-parallel", + "--attention-backend", + "fa3", + "--mem-frac", + "0.7", + "--cuda-graph-max-bs", + "32", + "--max-running-requests", + "32", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + ] + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5, + other_args=other_args, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + # "test_a_" prefix pins alphabetical first-run ordering so this + # warms up the server before any follow-up sibling test methods. + def test_a_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=500, + num_threads=32, + num_shots=20, + ) + metrics = run_eval(args) + print(f"{metrics=}") + + if is_in_ci(): + write_github_step_summary( + f"### test_a_gsm8k (deepseek-v3-mla-cp-in-seq-split)\n" + f'{metrics["score"]=:.3f}\n' + ) + self.assertGreater(metrics["score"], GSM8K_ACCURACY_THRESHOLD) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py new file mode 100644 index 000000000..3cb9b4e1a --- /dev/null +++ b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py @@ -0,0 +1,85 @@ +"""B200 extra CI: DeepSeek-V4-Flash FP4 with attn-CP (DSA prefill CP). + +Balanced recipe (TP=4, DeepEP, EAGLE) plus --attn-cp-size=4 with the +DSA prefill-CP round-robin-split mode. Split out of +models_e2e/test_deepseek_v4_flash_fp4_b200.py so the `cp` group covers +all context-parallel tests. + +Registry: extra-b-test-4-gpu-b200 (label-gated extra CI, 4x B200) +""" + +import unittest + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.basic_decode_correctness_kit import BasicDecodeCorrectnessMixin +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.test_utils import ( + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, + try_cached_model, +) + +register_cuda_ci(est_time=235, stage="extra-b", runner_config="4-gpu-b200") + +MODEL = "deepseek-ai/DeepSeek-V4-Flash" +SERVER_LAUNCH_TIMEOUT = 3600 +DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}' + +_DEEPEP_ENV = { + "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", +} + + +class TestDSV4FlashFP4B200Balanced_CP( + BasicDecodeCorrectnessMixin, + GSM8KMixin, + CustomTestCase, +): + """Balanced recipe: TP=4, DP=4, DeepEP, EAGLE (1-step spec).""" + + gsm8k_accuracy_thres = 0.93 + + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(MODEL) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp", + "4", + "--attn-cp-size", + "4", + "--enable-dp-attention", + "--moe-a2a-backend", + "deepep", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "1", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "2", + "--enable-dsa-prefill-context-parallel", + "--dsa-prefill-cp-mode", + "round-robin-split", + "--deepep-config", + DEEPEP_CONFIG, + ], + env=_DEEPEP_ENV, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/4-gpu-models/test_qwen3_30b.py b/test/registered/cp/test_qwen3_30b.py similarity index 98% rename from test/registered/4-gpu-models/test_qwen3_30b.py rename to test/registered/cp/test_qwen3_30b.py index b81534268..e21dafb5c 100644 --- a/test/registered/4-gpu-models/test_qwen3_30b.py +++ b/test/registered/cp/test_qwen3_30b.py @@ -11,7 +11,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=261, stage="base-c", runner_config="4-gpu-h100") +register_cuda_ci(est_time=261, stage="extra-b", runner_config="4-gpu-h100") QWEN3_30B_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8" diff --git a/test/registered/kernels/test_cp_prefix_len_fa3_parity.py b/test/registered/kernels/test_cp_prefix_len_fa3_parity.py new file mode 100644 index 000000000..11d7acbe5 --- /dev/null +++ b/test/registered/kernels/test_cp_prefix_len_fa3_parity.py @@ -0,0 +1,143 @@ +""" +FA3 parity test for `prepare_context_parallel_metadata`. + +Drives the real function and feeds its `kv_len_prev/next_tensor` into FA3 +via `flash_attn_with_kvcache`. Compares per-rank CP output against a +full-sequence FA3 reference computed over the unpadded `(prefix + extend)` +KV. Any discrepancy indicates the metadata function emitted wrong +`cache_seqlens` for at least one rank. +""" + +import unittest +from unittest.mock import patch + +import torch + +from sglang.srt.layers.utils.cp_utils import prepare_context_parallel_metadata +from sglang.srt.utils.common import ceil_align +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=5, stage="extra-a", runner_config="1-gpu-large") + +_DSA_UTILS = "sglang.srt.layers.attention.dsa.utils" +_DEVICE = "cuda" +_DTYPE = torch.bfloat16 +_HEAD_NUM = 8 +_HEAD_DIM = 128 +_SCALE = _HEAD_DIM**-0.5 + + +class TestCPPrefixLenFA3Parity(CustomTestCase): + """Per-rank FA3 output under CP must match a full-sequence reference.""" + + def _run_parity(self, prefix_len: int, extend_len: int, cp_size: int): + from sgl_kernel.flash_attn import flash_attn_with_kvcache + + torch.manual_seed(extend_len * 1_000_003 + prefix_len * 101 + cp_size) + + padded_extend = ceil_align(extend_len, cp_size) + pad = padded_extend - extend_len + self.assertGreaterEqual( + padded_extend, + 2 * cp_size, + "runtime `can_cp_split` would skip this case; pick a larger extend", + ) + + # Reference: one full-sequence FA3 call over the unpadded KV. + q_full = torch.randn( + extend_len, _HEAD_NUM, _HEAD_DIM, device=_DEVICE, dtype=_DTYPE + ) + k_full = torch.randn( + prefix_len + extend_len, _HEAD_NUM, _HEAD_DIM, device=_DEVICE, dtype=_DTYPE + ) + v_full = torch.randn( + prefix_len + extend_len, _HEAD_NUM, _HEAD_DIM, device=_DEVICE, dtype=_DTYPE + ) + ref = flash_attn_with_kvcache( + q=q_full.unsqueeze(0), + k_cache=k_full.unsqueeze(0), + v_cache=v_full.unsqueeze(0), + cache_seqlens=torch.tensor( + [k_full.shape[0]], dtype=torch.int32, device=_DEVICE + ), + softmax_scale=_SCALE, + causal=True, + ).squeeze(0) + + # CP path sees tensors padded to `ceil_align(extend, cp_size)`, + # matching what `prepare_mlp_sync_batch` does in production. + zeros = torch.zeros(pad, _HEAD_NUM, _HEAD_DIM, device=_DEVICE, dtype=_DTYPE) + q_padded = torch.cat([q_full, zeros], dim=0) + k_padded = torch.cat([k_full, zeros], dim=0) + v_padded = torch.cat([v_full, zeros], dim=0) + + seqs_len = [prefix_len + extend_len] + extend_lens = [extend_len] + + def _call_meta(rank: int): + return prepare_context_parallel_metadata( + padded_extend, rank, cp_size, seqs_len, extend_lens=extend_lens + ) + + # Exercise the non-DSA branch; the DSA branch uses a separate + # `prefix_len` pathway re-added by `_get_topk_ragged_with_cp`. + with ( + patch(f"{_DSA_UTILS}.is_dsa_enable_prefill_cp", return_value=False), + patch( + f"{_DSA_UTILS}.is_dsa_prefill_cp_round_robin_split", + return_value=False, + ), + ): + meta0 = _call_meta(0) + cp_segment_num = 2 * cp_size + blocks_q = list(torch.split(q_padded, meta0.split_list, dim=0)) + outs = [None] * cp_segment_num + + for rank in range(cp_size): + meta = meta0 if rank == 0 else _call_meta(rank) + for idx, cs_tensor in ( + (rank, meta.kv_len_prev_tensor), + (cp_size * 2 - rank - 1, meta.kv_len_next_tensor), + ): + if meta0.split_list[idx] == 0: + outs[idx] = torch.empty( + 0, _HEAD_NUM, _HEAD_DIM, device=_DEVICE, dtype=_DTYPE + ) + continue + outs[idx] = flash_attn_with_kvcache( + q=blocks_q[idx].unsqueeze(0), + k_cache=k_padded.unsqueeze(0), + v_cache=v_padded.unsqueeze(0), + cache_seqlens=cs_tensor, + softmax_scale=_SCALE, + causal=True, + ).squeeze(0) + + cp_out = torch.cat(outs, dim=0) + err = (cp_out[:extend_len].float() - ref.float()).abs().max().item() + + self.assertLess( + err, + 1e-2, + f"CP output diverges from full-sequence FA3 reference by " + f"max_err={err:.5f} " + f"(prefix_len={prefix_len}, extend_len={extend_len}, " + f"cp_size={cp_size}, pad={pad})", + ) + + def test_cp2_prefix1_extend3(self): + """cp_size=2, prefix_len=1, extend_len=3 (pad=1).""" + self._run_parity(prefix_len=1, extend_len=3, cp_size=2) + + def test_cp4_prefix1_extend7(self): + """cp_size=4, prefix_len=1, extend_len=7 (pad=1).""" + self._run_parity(prefix_len=1, extend_len=7, cp_size=4) + + def test_cp8_prefix1_extend17(self): + """cp_size=8, prefix_len=1, extend_len=17 (pad=7).""" + self._run_parity(prefix_len=1, extend_len=17, cp_size=8) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kernels/test_mla_cp_fa3_parity.py b/test/registered/kernels/test_mla_cp_fa3_parity.py new file mode 100644 index 000000000..9e2a4564a --- /dev/null +++ b/test/registered/kernels/test_mla_cp_fa3_parity.py @@ -0,0 +1,203 @@ +"""FA3 numerical parity for MLA prefill CP. + +Verifies the rank-local zigzag-split FA3 path (``_mla_cp_attn`` + +``cp_attn_forward_extend`` in ``flashattention_backend.py``) matches a +single non-CP ``flash_attn_with_kvcache`` over the full sequence. + +Single-process, single-layer, pre-populated paged KV cache. Requires +FA3 ver=3 (Hopper+). +""" + +import math +import sys +from types import SimpleNamespace + +import pytest +import torch + +from sglang.srt.layers.utils.cp_utils import ( + ContextParallelMetadata, + cp_attn_forward_extend, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, stage="extra-a", runner_config="1-gpu-large") + +if not torch.cuda.is_available(): + pytest.skip(reason="CUDA required for FA3", allow_module_level=True) + +_cap = torch.cuda.get_device_capability(0) +if _cap[0] < 9: + pytest.skip( + reason=f"FA3 ver=3 requires Hopper (sm90+); got sm{_cap[0]}{_cap[1]}", + allow_module_level=True, + ) + +try: + from sgl_kernel.flash_attn import flash_attn_with_kvcache +except ImportError as e: + pytest.skip( + reason=f"sgl_kernel.flash_attn unavailable: {e}", + allow_module_level=True, + ) + +DEVICE = torch.device("cuda") +DTYPE = torch.bfloat16 + +# Default shape is DeepSeek V3/R1 TP=8 MLA: 16 heads, v=512, rope=64. +NUM_HEADS = 16 +V_HEAD_DIM = 512 +QK_ROPE_HEAD_DIM = 64 +PAGE_SIZE = 1 + + +def _build_cache_and_q(seq_len): + """Pre-populated paged KV cache + full-sequence q. + + Pre-population mirrors upstream ``rebuild_cp_kv_cache``, which all-gathers + rank-local KV into the global pool before the attention call, so each + rank's FA3 invocation sees the same fully-populated cache. + """ + num_pages = (seq_len + PAGE_SIZE - 1) // PAGE_SIZE + c_kv_cache = torch.randn( + num_pages, PAGE_SIZE, 1, V_HEAD_DIM, dtype=DTYPE, device=DEVICE + ) + k_rope_cache = torch.randn( + num_pages, PAGE_SIZE, 1, QK_ROPE_HEAD_DIM, dtype=DTYPE, device=DEVICE + ) + q_nope = torch.randn(seq_len, NUM_HEADS, V_HEAD_DIM, dtype=DTYPE, device=DEVICE) + q_rope = torch.randn( + seq_len, NUM_HEADS, QK_ROPE_HEAD_DIM, dtype=DTYPE, device=DEVICE + ) + page_table = torch.arange(num_pages, dtype=torch.int32, device=DEVICE).unsqueeze(0) + return c_kv_cache, k_rope_cache, q_nope, q_rope, page_table + + +def _full_seq_attn( + seq_len, q_nope, q_rope, c_kv_cache, k_rope_cache, page_table, softmax_scale +): + """Non-CP reference: single flash_attn_with_kvcache over the full seq.""" + return flash_attn_with_kvcache( + q=q_rope, + qv=q_nope, + k_cache=k_rope_cache, + v_cache=c_kv_cache, + page_table=page_table, + cache_seqlens=torch.tensor([seq_len], dtype=torch.int32, device=DEVICE), + cu_seqlens_q=torch.tensor([0, seq_len], dtype=torch.int32, device=DEVICE), + cu_seqlens_k_new=None, + max_seqlen_q=seq_len, + softmax_scale=softmax_scale, + causal=True, + ver=3, + ) + + +def _cp_attn_for_rank( + rank, + cp_size, + block_size, + q_nope, + q_rope, + c_kv_cache, + k_rope_cache, + page_table, + softmax_scale, +): + """Run the rank-local CP closure from ``flashattention_backend.py``. + + Zigzag layout: rank r gets blocks [r, num_blocks - 1 - r] where + num_blocks = cp_size * 2. kv_len for each half is the cumulative KV + extent through the end of that block. + """ + num_blocks = cp_size * 2 + b_prev, b_next = rank, num_blocks - 1 - rank + prev_slice = slice(b_prev * block_size, (b_prev + 1) * block_size) + next_slice = slice(b_next * block_size, (b_next + 1) * block_size) + + q_nope_local = torch.cat([q_nope[prev_slice], q_nope[next_slice]], dim=0) + q_rope_local = torch.cat([q_rope[prev_slice], q_rope[next_slice]], dim=0) + q_fused = torch.cat([q_nope_local, q_rope_local], dim=-1) + + cp_meta = ContextParallelMetadata( + kv_len_prev_tensor=torch.tensor( + [(b_prev + 1) * block_size], dtype=torch.int32, device=DEVICE + ), + kv_len_next_tensor=torch.tensor( + [(b_next + 1) * block_size], dtype=torch.int32, device=DEVICE + ), + actual_seq_q_prev=block_size, + actual_seq_q_next=block_size, + ) + fb = SimpleNamespace(attn_cp_metadata=cp_meta) + + def _mla_cp_attn(q_chunk, cu_seqlens_q_cp, cache_seqlens_cp, max_seqlen_q_cp): + q_nope_chunk = q_chunk[..., :V_HEAD_DIM] + q_rope_chunk = q_chunk[..., V_HEAD_DIM:] + return flash_attn_with_kvcache( + q=q_rope_chunk, + qv=q_nope_chunk, + k_cache=k_rope_cache, + v_cache=c_kv_cache, + page_table=page_table, + cache_seqlens=cache_seqlens_cp, + cu_seqlens_q=cu_seqlens_q_cp, + cu_seqlens_k_new=None, + max_seqlen_q=max_seqlen_q_cp, + softmax_scale=softmax_scale, + causal=True, + ver=3, + ) + + local_out = cp_attn_forward_extend(fb, q_fused, DEVICE, _mla_cp_attn) + return local_out, prev_slice, next_slice + + +@pytest.mark.parametrize( + "cp_size, block_size", + [ + (2, 64), # DSv3 TP=8 baseline + (2, 128), # longer per-block seq + (4, 32), # multi-rank zigzag: rank r gets blocks [r, 7-r] + ], +) +def test_cp_parity(cp_size, block_size): + torch.manual_seed(0) + seq_len = block_size * cp_size * 2 + softmax_scale = 1.0 / math.sqrt(V_HEAD_DIM + QK_ROPE_HEAD_DIM) + + c_kv_cache, k_rope_cache, q_nope, q_rope, page_table = _build_cache_and_q(seq_len) + ref_out = _full_seq_attn( + seq_len, q_nope, q_rope, c_kv_cache, k_rope_cache, page_table, softmax_scale + ) + + for rank in range(cp_size): + local_out, prev_slice, next_slice = _cp_attn_for_rank( + rank, + cp_size, + block_size, + q_nope, + q_rope, + c_kv_cache, + k_rope_cache, + page_table, + softmax_scale, + ) + torch.testing.assert_close( + local_out[:block_size], + ref_out[prev_slice], + rtol=1e-3, + atol=5e-3, + msg=f"rank={rank} prev-half mismatch", + ) + torch.testing.assert_close( + local_out[block_size:], + ref_out[next_slice], + rtol=1e-3, + atol=5e-3, + msg=f"rank={rank} next-half mismatch", + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200.py b/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200.py index 5e659af08..f8536767d 100644 --- a/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200.py +++ b/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200.py @@ -20,7 +20,7 @@ from sglang.test.test_utils import ( try_cached_model, ) -register_cuda_ci(est_time=700, stage="base-c", runner_config="dsv4-4-gpu-b200") +register_cuda_ci(est_time=465, stage="base-c", runner_config="dsv4-4-gpu-b200") MODEL = "deepseek-ai/DeepSeek-V4-Flash" SERVER_LAUNCH_TIMEOUT = 3600 @@ -156,54 +156,5 @@ class TestDSV4FlashFP4NonMTPB200( kill_process_tree(cls.process.pid) -class TestDSV4FlashFP4B200Balanced_CP( - BasicDecodeCorrectnessMixin, - GSM8KMixin, - CustomTestCase, -): - """Balanced recipe: TP=4, DP=4, DeepEP, EAGLE (1-step spec).""" - - gsm8k_accuracy_thres = 0.93 - - @classmethod - def setUpClass(cls): - cls.model = try_cached_model(MODEL) - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=SERVER_LAUNCH_TIMEOUT, - other_args=[ - "--trust-remote-code", - "--tp", - "4", - "--attn-cp-size", - "4", - "--enable-dp-attention", - "--moe-a2a-backend", - "deepep", - "--speculative-algorithm", - "EAGLE", - "--speculative-num-steps", - "1", - "--speculative-eagle-topk", - "1", - "--speculative-num-draft-tokens", - "2", - "--enable-dsa-prefill-context-parallel", - "--dsa-prefill-cp-mode", - "round-robin-split", - "--deepep-config", - DEEPEP_CONFIG, - ], - env=_DEEPEP_ENV, - ) - - @classmethod - def tearDownClass(cls): - if hasattr(cls, "process") and cls.process: - kill_process_tree(cls.process.pid) - - if __name__ == "__main__": unittest.main()