diff --git a/python/sglang/jit_kernel/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh b/python/sglang/jit_kernel/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh index 7d5f97824..9b5ce3178 100644 --- a/python/sglang/jit_kernel/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh +++ b/python/sglang/jit_kernel/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh @@ -155,8 +155,10 @@ struct MegaMoEPreDispatchKernel { .with_dtype() .with_device(device) .verify(topk_weights); + // DeepGEMM versions expose this fp8 dispatch buffer either as raw int8 + // storage or as torch.float8_e4m3fn; the kernel writes fp8 bytes in both. TensorMatcher({P, H}) // buf.x - .with_dtype() + .with_dtype() .with_device(device) .verify(buf_x); // buf.x_sf is the contiguous row-major int32 view from DeepGEMM's mega diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index d0d36d7fe..d8c843485 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -381,6 +381,14 @@ class TboDPAttentionPreparer: self.enable_two_batch_overlap = enable_two_batch_overlap + # Short-circuit when TBO is off: prepare_mlp_sync_batch_raw invokes + # this preparer unconditionally for the forward_mode all-gather, but + # compute_split_seq_index is TBO-only and undefined for some modes + # (e.g. MIXED from enable_mixed_chunk). + if not enable_two_batch_overlap: + self.local_tbo_split_seq_index = None + return False, self._compute_local_forward_mode(local_batch) + if local_batch is not None: token_num_per_seq = get_token_num_per_seq( forward_mode=local_batch.forward_mode, spec_info=local_batch.spec_info @@ -692,6 +700,7 @@ class TboForwardBatchPreparer: "all_extend_in_batch", "return_logprob", "can_run_dp_cuda_graph", + "can_run_dp_breakable_cuda_graph", "dp_padding_mode", "global_forward_mode", "is_prefill_only", diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index 49cf92ba7..8694299b8 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -85,10 +85,40 @@ class AttentionBackend(ABC): # Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum. needs_cpu_seq_lens: bool = True + # Most attention backends can rebuild and replace forward metadata before + # every forward. BCG capture is different: some backends expose metadata + # tensors to kernels across graph breaks, so the captured graph depends on + # those tensor addresses. Such backends opt in here, create the metadata + # object during capture, and refresh its dynamic fields before each replay. + use_captured_forward_metadata_for_breakable_cuda_graph: bool = False + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): """Init the global shared states for cuda graph.""" raise NotImplementedError() + def init_forward_metadata_for_breakable_cuda_graph_capture( + self, + forward_batch: ForwardBatch, + ): + """Create forward metadata whose tensor addresses will be graph-captured.""" + raise NotImplementedError() + + def prepare_forward_metadata_for_breakable_cuda_graph_replay( + self, + capture_metadata, + forward_batch: ForwardBatch, + *, + static_forward_batch: Optional[ForwardBatch] = None, + ) -> None: + """Refresh captured metadata for the current batch before BCG replay. + + Implementations should update ``capture_metadata`` in place where graph + address stability is required, assign any safe per-replay objects, and + make the backend's active ``forward_metadata`` point to the captured + metadata object. + """ + raise NotImplementedError() + def get_cuda_graph_seq_len_fill_value(self): """Get the fill value for padded seq lens. Typically, it is 0 or 1.""" raise NotImplementedError() diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index aa80dfefd..f57486e68 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -184,6 +184,47 @@ class DSV4AttnMetadata: ], ) + def refresh_for_breakable_cuda_graph_replay_(self, other: DSV4AttnMetadata) -> None: + assert self.c4_sparse_topk == other.c4_sparse_topk + assert self.page_size == other.page_size + assert self.cuda_int32_kwargs == other.cuda_int32_kwargs + + tensor_copy_fields = [ + "raw_out_loc", + "seq_lens_casual", + "positions_casual", + "c4_out_loc", + "c128_out_loc", + "c4_topk_lengths_raw", + "c4_topk_lengths_clamp1", + "c4_sparse_topk_lengths", + ] + reference_assign_fields = [ + "page_table", + "swa_page_indices", + "swa_topk_lengths", + "c128_page_indices", + "c128_topk_lengths_clamp1", + "c1_flashmla_metadata", + "c4_flashmla_metadata", + "c128_flashmla_metadata", + ] + # Keep graph-captured tensor objects alive for fields that captured + # kernels read by address; overwrite only their contents. + for field_name in tensor_copy_fields: + src_val = getattr(other, field_name) + dst_val = getattr(self, field_name) + if src_val is None and dst_val is None: + continue + assert dst_val is not None, f"{field_name=} {src_val=} {dst_val=}" + dst_val.copy_(src_val) + + # These fields are safe to replace because captured kernels only need + # the current per-replay objects, or the field is produced inside the + # captured graph before the attention graph break consumes it. + for field_name in reference_assign_fields: + setattr(self, field_name, getattr(other, field_name)) + def init_compression_metadata(self): assert self.page_table.dim() == 2 assert ( @@ -312,6 +353,24 @@ class DSV4Metadata: ) self.sparse_prefill_cache = None + def refresh_for_breakable_cuda_graph_replay_(self, static_metadata: DSV4Metadata): + self.core_attn_metadata.refresh_for_breakable_cuda_graph_replay_( + static_metadata.core_attn_metadata + ) + maybe_copy_inplace(self.indexer_metadata, src=static_metadata.indexer_metadata) + maybe_copy_inplace( + self.c4_compress_metadata, src=static_metadata.c4_compress_metadata + ) + if envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get(): + # Online c128 prefill metadata may carry Python-side planner state, + # so assign the freshly built per-replay object. + self.c128_compress_metadata = static_metadata.c128_compress_metadata + else: + maybe_copy_inplace( + self.c128_compress_metadata, + src=static_metadata.c128_compress_metadata, + ) + @dataclass class DSV4RawVerifyMetadata: @@ -360,6 +419,8 @@ class _GraphBucket(enum.Enum): class DeepseekV4AttnBackend( AttentionBackend, C4IndexerBackendMixin, CompressorBackendMixin ): + use_captured_forward_metadata_for_breakable_cuda_graph: bool = True + def __init__( self, model_runner: ModelRunner, @@ -477,6 +538,7 @@ class DeepseekV4AttnBackend( num_tokens: int, extend_seq_lens: torch.Tensor, extend_seq_lens_cpu: List[int], + extend_start_loc: Optional[torch.Tensor] = None, need_compress: bool = True, use_prefill_cuda_graph: bool = False, ) -> DSV4Metadata: @@ -486,6 +548,9 @@ class DeepseekV4AttnBackend( extend_seq_lens=extend_seq_lens_cpu, req_pool_indices=req_pool_indices, padded_num_tokens=out_cache_loc.shape[0], + seq_lens_tensor=seq_lens, + extend_seq_lens_tensor=extend_seq_lens, + extend_start_loc=extend_start_loc, ) core_attn_metadata = self.make_core_attn_metadata( req_to_token=self.req_to_token, @@ -504,23 +569,48 @@ class DeepseekV4AttnBackend( if not need_compress: create = _create_dummy_paged_compress_data else: - create = functools.partial( - create_paged_compressor_data, - is_prefill=True, - token_to_kv_pool=self.token_to_kv_pool, - req_to_token=self.req_to_token, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=seq_lens_cpu, - extend_lens=extend_seq_lens, - extend_lens_cpu=extend_seq_lens_cpu, - use_prefill_cuda_graph=use_prefill_cuda_graph, - ) + + def create(compress_ratio: Literal[4, 128]): + # Online c128 uses a different planner that cannot be created in + # prefill cuda-graph mode. Keep c4 graph-friendly while matching + # c128's existing online path. + use_graph_plan = use_prefill_cuda_graph and not ( + compress_ratio == 128 and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get() + ) + if use_graph_plan: + return create_paged_compressor_data( + compress_ratio=compress_ratio, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=None, + extend_lens=extend_seq_lens, + extend_lens_cpu=None, + use_prefill_cuda_graph=True, + num_q_tokens=out_cache_loc.shape[0], + ) + return create_paged_compressor_data( + compress_ratio=compress_ratio, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + extend_lens=extend_seq_lens, + extend_lens_cpu=extend_seq_lens_cpu, + use_prefill_cuda_graph=use_graph_plan, + ) + + c4_compress_metadata = create(compress_ratio=4) + c128_compress_metadata = create(compress_ratio=128) return DSV4Metadata( core_attn_metadata, indexer_metadata, - c4_compress_metadata=create(compress_ratio=4), - c128_compress_metadata=create(compress_ratio=128), + c4_compress_metadata=c4_compress_metadata, + c128_compress_metadata=c128_compress_metadata, ) def init_forward_metadata_target_verify( @@ -582,6 +672,7 @@ class DeepseekV4AttnBackend( num_tokens=num_tokens, extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, + extend_start_loc=None, need_compress=True, use_prefill_cuda_graph=use_prefill_cuda_graph, ) @@ -689,6 +780,7 @@ class DeepseekV4AttnBackend( num_tokens=num_tokens, extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, + extend_start_loc=None, need_compress=False, use_prefill_cuda_graph=use_prefill_cuda_graph, ) @@ -867,6 +959,16 @@ class DeepseekV4AttnBackend( if self.mtp_enabled and forward_batch.forward_mode.is_idle(): return + self.forward_metadata = self._build_forward_metadata(forward_batch) + self.init_forward_metadata_in_graph(forward_batch) + + def _build_forward_metadata( + self, + forward_batch: ForwardBatch, + *, + max_seq_len_override: Optional[int] = None, + use_prefill_cuda_graph: bool = False, + ): req_pool_indices = forward_batch.req_pool_indices seq_lens = forward_batch.seq_lens.to(torch.int32) seq_lens_cpu = forward_batch.seq_lens_cpu @@ -874,7 +976,11 @@ class DeepseekV4AttnBackend( assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0 assert seq_lens_cpu is not None - max_seq_len = int(seq_lens_cpu.max().item()) + max_seq_len = ( + int(seq_lens_cpu.max().item()) + if max_seq_len_override is None + else max_seq_len_override + ) if forward_batch.forward_mode.is_decode_or_idle(): # DSv4 bakes this step's KV write target (c4/c128) into metadata, @@ -919,13 +1025,43 @@ class DeepseekV4AttnBackend( num_tokens=sum(extend_seq_lens_cpu), extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, + extend_start_loc=forward_batch.extend_start_loc, need_compress=not is_draft, + use_prefill_cuda_graph=use_prefill_cuda_graph, ) else: raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") - self.forward_metadata = metadata - self.init_forward_metadata_in_graph(forward_batch) + return metadata + + def init_forward_metadata_for_breakable_cuda_graph_capture( + self, forward_batch: ForwardBatch + ): + self.forward_metadata = self._build_forward_metadata( + forward_batch, + max_seq_len_override=self.MAX_SEQ_LEN_FOR_CAPTURE, + use_prefill_cuda_graph=True, + ) + return self.forward_metadata + + def prepare_forward_metadata_for_breakable_cuda_graph_replay( + self, + capture_metadata, + forward_batch: ForwardBatch, + *, + static_forward_batch: Optional[ForwardBatch] = None, + ) -> None: + # Build graph-compatible metadata against the padded static batch. The + # batch still carries live seq/extend lens, so the online c128 prefill + # plan remains batch-specific without constructing a second metadata set. + static_metadata = self._build_forward_metadata( + static_forward_batch if static_forward_batch is not None else forward_batch, + max_seq_len_override=self.MAX_SEQ_LEN_FOR_CAPTURE, + use_prefill_cuda_graph=True, + ) + assert isinstance(capture_metadata, DSV4Metadata) + capture_metadata.refresh_for_breakable_cuda_graph_replay_(static_metadata) + self.forward_metadata = capture_metadata def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None: self.cuda_graph_metadata_of_bucket_and_bs: Dict[ @@ -1087,16 +1223,17 @@ class DeepseekV4AttnBackend( swa_page_indices = core_attn_metadata.swa_page_indices swa_topk_lengths = core_attn_metadata.swa_topk_lengths - if self.mtp_enabled: - if swa_page_indices.shape[0] != q.shape[0]: - swa_page_indices = _pad_tensor_to_size( - swa_page_indices, q.shape[0], value=0 - ) + def match_num_queries(x, value): + if x is None or x.shape[0] == q.shape[0]: + return x + if x.shape[0] > q.shape[0]: + return x[: q.shape[0]] + return _pad_tensor_to_size(x, q.shape[0], value=value) - if swa_topk_lengths.shape[0] != q.shape[0]: - swa_topk_lengths = _pad_tensor_to_size( - swa_topk_lengths, q.shape[0], value=1 - ) + swa_page_indices = match_num_queries(swa_page_indices, value=0) + swa_topk_lengths = match_num_queries(swa_topk_lengths, value=1) + extra_indices = match_num_queries(extra_indices, value=-1) + extra_topk_lengths = match_num_queries(extra_topk_lengths, value=1) if q.ndim == 3: q = q.unsqueeze(1) @@ -1281,7 +1418,24 @@ class DeepseekV4AttnBackend( extend_seq_lens: List[int], req_pool_indices: torch.Tensor, padded_num_tokens: Optional[int], + seq_lens_tensor: Optional[torch.Tensor] = None, + extend_seq_lens_tensor: Optional[torch.Tensor] = None, + extend_start_loc: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: + if ( + seq_lens_tensor is not None + and extend_seq_lens_tensor is not None + and extend_start_loc is not None + ): + return self._expand_prefill_casually_vectorized( + num_tokens=num_tokens, + seq_lens=seq_lens_tensor, + extend_seq_lens=extend_seq_lens_tensor, + extend_start_loc=extend_start_loc, + req_pool_indices=req_pool_indices, + padded_num_tokens=padded_num_tokens, + ) + seq_lens_casual = torch.empty(num_tokens, **self.cuda_int32_kwargs) idx_to_req_repeated = torch.empty(num_tokens, **self.cuda_int32_kwargs) offset = 0 @@ -1309,6 +1463,48 @@ class DeepseekV4AttnBackend( return seq_lens_casual, req_pool_indices_repeated + def _expand_prefill_casually_vectorized( + self, + num_tokens: int, + seq_lens: torch.Tensor, + extend_seq_lens: torch.Tensor, + extend_start_loc: torch.Tensor, + req_pool_indices: torch.Tensor, + padded_num_tokens: Optional[int], + ) -> Tuple[torch.Tensor, torch.Tensor]: + repeats = extend_seq_lens.to(torch.int64) + req_pool_indices_repeated = torch.repeat_interleave( + req_pool_indices, repeats, output_size=num_tokens + ) + + start_positions = seq_lens.to(torch.int32) - extend_seq_lens.to(torch.int32) + 1 + start_positions_repeated = torch.repeat_interleave( + start_positions, repeats, output_size=num_tokens + ) + start_locs_repeated = torch.repeat_interleave( + extend_start_loc.to(torch.int32), repeats, output_size=num_tokens + ) + token_offsets = ( + torch.arange(num_tokens, **self.cuda_int32_kwargs) - start_locs_repeated + ) + seq_lens_casual = start_positions_repeated + token_offsets + + if padded_num_tokens is not None and padded_num_tokens > num_tokens: + pad_size = padded_num_tokens - num_tokens + seq_lens_casual = torch.nn.functional.pad( + seq_lens_casual, + (0, pad_size), + value=1, + ) + req_pool_indices_repeated = torch.cat( + ( + req_pool_indices_repeated, + req_pool_indices_repeated[-1:].expand(pad_size), + ) + ) + + return seq_lens_casual, req_pool_indices_repeated + def expand_extend_with_same_length( self, bs: int, @@ -1467,6 +1663,35 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend): for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_forward_metadata(forward_batch) + def init_forward_metadata_for_breakable_cuda_graph_capture( + self, forward_batch: ForwardBatch + ): + ret = [] + for i in range(self.speculative_num_steps - 1): + ret.append( + self.attn_backends[ + i + ].init_forward_metadata_for_breakable_cuda_graph_capture(forward_batch) + ) + return ret + + def prepare_forward_metadata_for_breakable_cuda_graph_replay( + self, + capture_metadata, + forward_batch: ForwardBatch, + *, + static_forward_batch: Optional[ForwardBatch] = None, + ) -> None: + assert len(capture_metadata) == self.speculative_num_steps - 1 + for i in range(self.speculative_num_steps - 1): + self.attn_backends[ + i + ].prepare_forward_metadata_for_breakable_cuda_graph_replay( + capture_metadata[i], + forward_batch, + static_forward_batch=static_forward_batch, + ) + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): for i in range(self.speculative_num_steps): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index 48172f06f..bc8b2111b 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -179,6 +179,8 @@ class CompressorBackendMixin: if compressor.ratio == 4 else core_metadata.c128_out_loc ) + if out_loc.shape[0] > new_compressed_kv.shape[0]: + out_loc = out_loc[: new_compressed_kv.shape[0]] if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get(): token_to_kv_pool.set_extra_key_buffer_fused( layer_id=layer_id, @@ -202,16 +204,19 @@ class CompressorBackendMixin: assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) new_compressed_kv = compressor(x, forward_batch, attn_backend=self) + out_loc = self.forward_metadata.core_metadata.c4_out_loc + if out_loc.shape[0] > new_compressed_kv.shape[0]: + out_loc = out_loc[: new_compressed_kv.shape[0]] if self.enable_deepseek_v4_fp4_indexer: token_to_kv_pool.set_index_k_fp4( layer_id=layer_id, - loc=self.forward_metadata.core_metadata.c4_out_loc, + loc=out_loc, cache_k=new_compressed_kv, ) elif envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get(): token_to_kv_pool.set_index_k_fused( layer_id=layer_id, - loc=self.forward_metadata.core_metadata.c4_out_loc, + loc=out_loc, cache_k=new_compressed_kv, ) else: @@ -220,7 +225,7 @@ class CompressorBackendMixin: ) token_to_kv_pool.set_index_k_scale_buffer( layer_id=layer_id, - loc=self.forward_metadata.core_metadata.c4_out_loc, + loc=out_loc, index_k=new_compressed_kv_fp8, index_k_scale=new_compressed_kv_scale, ) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index e14bdc648..eceba4e18 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -455,13 +455,22 @@ class C4IndexerBackendMixin: assert isinstance(indexer_metadata, PagedIndexerMetadata) + positions = core_metadata.positions + num_queries = min(x.shape[0], q_lora.shape[0], positions.shape[0]) + if x.shape[0] != num_queries: + x = x[:num_queries] + if q_lora.shape[0] != num_queries: + q_lora = q_lora[:num_queries] + if positions.shape[0] != num_queries: + positions = positions[:num_queries] + if enable_multi_stream: q_indexer, weights, c4_indexer_kv_cache = ( self._forward_prepare_multi_stream( x=x, q_lora=q_lora, c4_indexer=c4_indexer, - positions=core_metadata.positions, + positions=positions, forward_batch=forward_batch, token_to_kv_pool=token_to_kv_pool, alt_streams=alt_streams, @@ -474,7 +483,7 @@ class C4IndexerBackendMixin: x=x, q_lora=q_lora, c4_indexer=c4_indexer, - positions=core_metadata.positions, + positions=positions, forward_batch=forward_batch, token_to_kv_pool=token_to_kv_pool, skip_compressor=skip_compressor, @@ -519,7 +528,22 @@ class C4IndexerBackendMixin: else: from deep_gemm import fp8_paged_mqa_logits as fn - _c4sl = indexer_metadata.c4_seq_lens + query_rows = q_indexer[0].shape[0] if use_fp4_indexer else q_indexer.shape[0] + + def match_num_queries(tensor: torch.Tensor, value: int) -> torch.Tensor: + if tensor.shape[0] == query_rows: + return tensor + if tensor.shape[0] > query_rows: + return tensor[:query_rows] + pad = (0, 0) * (tensor.dim() - 1) + (0, query_rows - tensor.shape[0]) + return F.pad(tensor, pad, value=value) + + c4_seq_lens = match_num_queries(indexer_metadata.c4_seq_lens, value=1) + _c4sl = c4_seq_lens + page_table = match_num_queries(indexer_metadata.page_table, value=0) + c4_sparse_page_indices = match_num_queries( + core_metadata.c4_sparse_page_indices, value=-1 + ) _use_tilelang = ( envs.SGLANG_OPT_USE_TILELANG_INDEXER.get() and not use_fp4_indexer ) @@ -531,7 +555,7 @@ class C4IndexerBackendMixin: c4_indexer_kv_cache, weights, _c4sl, - indexer_metadata.page_table, + page_table, indexer_metadata.deep_gemm_metadata, indexer_metadata.max_c4_seq_len, False, @@ -551,10 +575,10 @@ class C4IndexerBackendMixin: raw_indices = None if capture_enabled: - raw_indices = torch.empty_like(core_metadata.c4_sparse_page_indices) + raw_indices = torch.empty_like(c4_sparse_page_indices) elif hisparse_decode: raw_indices = hisparse_coordinator.raw_indices_buffer[ - : core_metadata.c4_sparse_page_indices.size(0) + : c4_sparse_page_indices.size(0) ] elif core_metadata.c4_sparse_raw_indices is not None: raw_indices = core_metadata.c4_sparse_raw_indices @@ -562,27 +586,27 @@ class C4IndexerBackendMixin: if envs.SGLANG_TOPK_TRANSFORM_512_TORCH.get(): topk_transform_512_pytorch_vectorized( logits, - indexer_metadata.c4_seq_lens, - core_metadata.page_table, - core_metadata.c4_sparse_page_indices, + c4_seq_lens, + page_table, + c4_sparse_page_indices, indexer_metadata.c4_page_size, raw_indices, ) elif envs.SGLANG_OPT_USE_TOPK_V2.get() and raw_indices is None: topk_transform_512_v2( logits, - indexer_metadata.c4_seq_lens, - core_metadata.page_table, - core_metadata.c4_sparse_page_indices, + c4_seq_lens, + page_table, + c4_sparse_page_indices, indexer_metadata.c4_page_size, indexer_metadata.topk_metadata, ) else: topk_transform_512( logits, - indexer_metadata.c4_seq_lens, - core_metadata.page_table, - core_metadata.c4_sparse_page_indices, + c4_seq_lens, + page_table, + c4_sparse_page_indices, indexer_metadata.c4_page_size, raw_indices, ) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index db7ad15ff..42ea84310 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1667,6 +1667,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): is_extend_in_batch: bool = False all_extend_in_batch: bool = False # plumbing for downstream forks (PR #19639) can_run_dp_cuda_graph: bool = False + can_run_dp_breakable_cuda_graph: bool = False tbo_split_seq_index: Optional[int] = None # For processing logprobs @@ -2749,6 +2750,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): global_num_tokens=self.global_num_tokens, global_num_tokens_for_logprob=self.global_num_tokens_for_logprob, can_run_dp_cuda_graph=self.can_run_dp_cuda_graph, + can_run_dp_breakable_cuda_graph=self.can_run_dp_breakable_cuda_graph, is_extend_in_batch=self.is_extend_in_batch, all_extend_in_batch=self.all_extend_in_batch, is_prefill_only=self.is_prefill_only, diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index 518b1dd18..6e8c1d4c9 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -39,6 +39,7 @@ class MLPSyncBatchInfo: is_extend_in_batch: bool local_can_run_tbo: bool local_forward_mode: int + can_run_breakable_cuda_graph: bool # some gathered elements tp0_info: torch.Tensor = None @@ -57,6 +58,7 @@ class MLPSyncBatchInfo: int(self.is_extend_in_batch), int(self.local_can_run_tbo), self.local_forward_mode, + int(self.can_run_breakable_cuda_graph), ], device=device, dtype=dtype, @@ -71,6 +73,7 @@ class MLPSyncBatchInfo: 0, # is_extend_in_batch 1, # local_can_run_tbo ForwardMode.IDLE.value, # local_forward_mode + 0, # can_run_breakable_cuda_graph ], device=device, dtype=dtype, @@ -79,7 +82,7 @@ class MLPSyncBatchInfo: def all_gather(self, device, group: torch.distributed.ProcessGroup): local_info_tensor = self._get_local_tensor(device=device) global_info_tensor = torch.empty( - (self.dp_size, self.tp_size * self.cp_size, 6), + (self.dp_size, self.tp_size * self.cp_size, 7), dtype=torch.int64, device=device, ) @@ -95,7 +98,7 @@ class MLPSyncBatchInfo: tp_active_ranks = get_tp_group().active_ranks # Set fallback values for inactive ranks - tp_info = global_info_tensor.view(self.dp_size * self.tp_size * self.cp_size, 6) + tp_info = global_info_tensor.view(self.dp_size * self.tp_size * self.cp_size, 7) tp_info[tp_active_ranks == 0] = self._get_fallback_tensor(device=device) tp0_info = global_info_tensor[:, 0, :] @@ -106,6 +109,7 @@ class MLPSyncBatchInfo: self.global_num_tokens_for_logprob = cpu_data[:, 1].tolist() self.can_cuda_graph = bool(tp0_info[:, 2].min().item()) self.is_extend_in_batch = bool(tp0_info[:, 3].max().item()) + self.can_run_breakable_cuda_graph = bool(tp0_info[:, 6].min().item()) if _ENABLE_METRICS_DP_ATTENTION: self.dp_cooperation_info = DPCooperationInfo.create(tp0_info[:, 5].tolist()) @@ -132,6 +136,7 @@ def _update_gather_batch( # Check forward mode for cuda graph batch.can_run_dp_cuda_graph = mlp_sync_info.can_cuda_graph + batch.can_run_dp_breakable_cuda_graph = mlp_sync_info.can_run_breakable_cuda_graph def prepare_mlp_sync_batch_raw( @@ -178,6 +183,11 @@ def prepare_mlp_sync_batch_raw( or local_batch.forward_mode.is_decode_or_idle() or local_batch.forward_mode.is_prebuilt() ) and not disable_cuda_graph + can_run_breakable_cuda_graph = ( + local_batch is not None + and local_batch.forward_mode in (ForwardMode.EXTEND, ForwardMode.MIXED) + and not disable_cuda_graph + ) is_extend_in_batch = local_batch.forward_mode.is_extend() if local_batch else False if local_batch is not None: @@ -206,6 +216,7 @@ def prepare_mlp_sync_batch_raw( is_extend_in_batch=is_extend_in_batch, local_can_run_tbo=local_can_run_tbo, local_forward_mode=local_forward_mode, + can_run_breakable_cuda_graph=can_run_breakable_cuda_graph, ) if not skip_all_gather: diff --git a/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py b/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py index e78746a5f..d5efcd421 100644 --- a/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py @@ -36,10 +36,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( set_graph_pool_id, ) from sglang.srt.distributed.parallel_state import graph_capture -from sglang.srt.layers.dp_attention import ( - set_dp_buffer_len, - set_is_extend_in_batch, -) +from sglang.srt.layers.dp_attention import set_dp_buffer_len, set_is_extend_in_batch from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.pooler import EmbeddingPoolerOutput from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import ( @@ -123,6 +120,10 @@ class BreakableCudaGraphRunner: self.attention_layers = model_runner.attention_layers self.moe_layers = model_runner.moe_layers self.moe_fusions = model_runner.moe_fusions + self.use_captured_attn_metadata = ( + model_runner.attn_backend.use_captured_forward_metadata_for_breakable_cuda_graph + ) + self.attn_metadata_buffers = {} if self.use_captured_attn_metadata else None # Resolve the inner transformer-stack module (the same boundary PCG draws # via patch_model). At replay we monkey-patch this module's forward with @@ -167,6 +168,18 @@ class BreakableCudaGraphRunner: self.raw_num_tokens = 0 + def _has_inactive_dp_rank(self, forward_batch: "ForwardBatch") -> bool: + global_num_tokens = forward_batch.global_num_tokens_cpu + if global_num_tokens is None: + return False + + # DSV4 DP attention / DeepEP collectives need every DP rank to enter + # the same replay path. Sparse-DP batches fall back to eager to avoid + # hanging ranks that have zero local tokens. + return len(global_num_tokens) > 1 and any( + int(num_tokens) == 0 for num_tokens in global_num_tokens + ) + def _init_buffers(self, model_runner): """Initialize input buffers.""" from sglang.srt.model_executor.cuda_graph_buffer_registry import ( @@ -317,12 +330,47 @@ class BreakableCudaGraphRunner: """Warmup the model with a forward pass.""" num_tokens = self.capture_num_tokens[0] forward_batch = self._build_capture_forward_batch(num_tokens) - with forward_context( - ForwardContext(attn_backend=self.model_runner.attn_backend) + with ( + forward_context( + ForwardContext(attn_backend=self.model_runner.attn_backend) + ), + set_forward_context( + forward_batch, + self.attention_layers, + self.quant_config, + self.moe_layers, + self.moe_fusions, + ), ): - self.model_runner.attn_backend.init_forward_metadata(forward_batch) + self._init_forward_metadata_for_capture(forward_batch, num_tokens) self._run_forward(forward_batch, num_tokens) + def _init_forward_metadata_for_capture(self, forward_batch, num_tokens): + attn_backend = self.model_runner.attn_backend + if not self.use_captured_attn_metadata: + attn_backend.init_forward_metadata(forward_batch) + return + metadata = attn_backend.init_forward_metadata_for_breakable_cuda_graph_capture( + forward_batch + ) + assert self.attn_metadata_buffers is not None + self.attn_metadata_buffers[num_tokens] = metadata + + def _prepare_forward_metadata_for_replay( + self, forward_batch, static_forward_batch, num_tokens + ): + attn_backend = self.model_runner.attn_backend + if not self.use_captured_attn_metadata: + attn_backend.init_forward_metadata(forward_batch) + return + assert self.attn_metadata_buffers is not None + metadata = self.attn_metadata_buffers[num_tokens] + attn_backend.prepare_forward_metadata_for_breakable_cuda_graph_replay( + metadata, + forward_batch, + static_forward_batch=static_forward_batch, + ) + def _capture_all(self): """Capture breakable CUDA graphs for all token sizes.""" with ( @@ -364,6 +412,13 @@ class BreakableCudaGraphRunner: return False if forward_batch.replace_embeds is not None: return False + if self._has_inactive_dp_rank(forward_batch): + return False + if ( + forward_batch.global_num_tokens_cpu is not None + and not forward_batch.can_run_dp_breakable_cuda_graph + ): + return False num_tokens = len(forward_batch.input_ids) if forward_batch.return_logprob: for start_len, seq_len in zip( @@ -377,7 +432,7 @@ class BreakableCudaGraphRunner: def _capture_one(self, num_tokens, pool, stream): """Capture a breakable CUDA graph for one token size.""" forward_batch = self._build_capture_forward_batch(num_tokens) - self.model_runner.attn_backend.init_forward_metadata(forward_batch) + self._init_forward_metadata_for_capture(forward_batch, num_tokens) def run_once(): return self._run_forward(forward_batch, num_tokens) @@ -450,7 +505,9 @@ class BreakableCudaGraphRunner: original_layer_forward = self.layer_model.forward self.layer_model.forward = replay_layer_forward try: - self.model_runner.attn_backend.init_forward_metadata(forward_batch) + self._prepare_forward_metadata_for_replay( + forward_batch, static_forward_batch, static_num_tokens + ) with set_forward_context( static_forward_batch, self.attention_layers, diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 75feabdb2..2ad30612a 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -345,6 +345,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # Mirrors ScheduleBatch.all_extend_in_batch; kept for downstream forks. all_extend_in_batch: bool = False can_run_dp_cuda_graph: bool = False + can_run_dp_breakable_cuda_graph: bool = False global_forward_mode: Optional[ForwardMode] = None # For two-batch overlap @@ -647,6 +648,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): is_extend_in_batch=batch.is_extend_in_batch, all_extend_in_batch=batch.all_extend_in_batch, can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph, + can_run_dp_breakable_cuda_graph=batch.can_run_dp_breakable_cuda_graph, global_forward_mode=batch.global_forward_mode, is_prefill_only=batch.is_prefill_only, spec_algorithm=batch.spec_algorithm, diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 0594829bb..7494ce85a 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -27,6 +27,8 @@ from sglang.jit_kernel.dsv4 import ( fused_q_norm_rope, fused_rope_inplace, ) +from sglang.srt.compilation.compilation_config import register_split_op +from sglang.srt.compilation.piecewise_context_manager import get_forward_context from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.distributed import ( get_pp_group, @@ -82,6 +84,12 @@ from sglang.srt.layers.utils.cp_utils import ( ) from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.mem_cache.memory_pool import RadixAttention +from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import ( + eager_on_graph, +) +from sglang.srt.model_executor.breakable_cuda_graph.context import ( + is_in_breakable_cuda_graph, +) from sglang.srt.model_executor.cuda_graph_runner import ( compile_in_capture_mode, get_is_capture_mode, @@ -114,6 +122,7 @@ from sglang.srt.utils import ( log_info_on_rank0, make_layers, ) +from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.hf_transformers_utils import get_rope_config logger = logging.getLogger(__name__) @@ -191,6 +200,57 @@ if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardBatch +@register_custom_op(mutates_args=["output"]) +@register_split_op() +def deepseek_v4_attention_with_output( + query: torch.Tensor, + key_value: torch.Tensor, + output: torch.Tensor, + layer_id: int, + compress_ratio: int, + attn_sink: torch.Tensor, + save_kv_cache: bool, +) -> None: + context = get_forward_context() + forward_batch = context.forward_batch + attention_layers = context.attention_layers + attention_layer = attention_layers[layer_id] + real_num_tokens = forward_batch.num_token_non_padded_cpu + + query = query[:real_num_tokens] + key_value = key_value[:real_num_tokens] + + original_out_cache_loc = forward_batch.out_cache_loc + forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] + + attn_backend = get_attn_backend() + try: + ret = attn_backend.forward( + q=query, + k=key_value, + v=key_value, + layer=attention_layer, + forward_batch=forward_batch, + compress_ratio=compress_ratio, + attn_sink=attn_sink, + save_kv_cache=save_kv_cache, + ) + finally: + forward_batch.out_cache_loc = original_out_cache_loc + + assert ( + output[:real_num_tokens].numel() == ret.numel() + ), f"Output tensor element mismatch: {output[:real_num_tokens].numel()} != {ret.numel()}" + + output[:real_num_tokens].view(ret.shape).copy_(ret) + return + + +bcg_deepseek_v4_attention_with_output = eager_on_graph(True)( + deepseek_v4_attention_with_output +) + + @triton.jit def _rms_normalize_kernel( x_ptr, @@ -889,17 +949,33 @@ class MQALayer(nn.Module): # tell the backend to skip its own store_cache. When `kv is None` # (no DSA-CP), pass `q` as a sentinel for the `k is v` assert; the # attention path doesn't read it once `save_kv_cache=False`. + attn_q = q_padded if q_padded is not None else q attn_k = kv if kv is not None else q - o = attn_backend.forward( - q=q_padded if q_padded is not None else q, - k=attn_k, - v=attn_k, - layer=self.attn_mqa, - forward_batch=forward_batch, - compress_ratio=self.compress_ratio, - attn_sink=self.attn_sink, - save_kv_cache=False, - ) + save_kv_cache = False + if forward_batch.forward_mode.is_extend() and is_in_breakable_cuda_graph(): + o = attn_q.new_empty( + (*attn_q.shape[:-1], self.attn_mqa.v_head_dim), + ) + bcg_deepseek_v4_attention_with_output( + attn_q, + attn_k, + o, + self.attn_mqa.layer_id, + self.compress_ratio, + self.attn_sink, + save_kv_cache, + ) + else: + o = attn_backend.forward( + q=attn_q, + k=attn_k, + v=attn_k, + layer=self.attn_mqa, + forward_batch=forward_batch, + compress_ratio=self.compress_ratio, + attn_sink=self.attn_sink, + save_kv_cache=save_kv_cache, + ) o = o[:, tp_slice, :] fused_rope_inplace( o[..., -self.qk_rope_head_dim :], diff --git a/test/registered/attention/unittests/dsv4/test_deepseek_v4.py b/test/registered/attention/unittests/dsv4/test_deepseek_v4.py index 0166025eb..f2877d153 100644 --- a/test/registered/attention/unittests/dsv4/test_deepseek_v4.py +++ b/test/registered/attention/unittests/dsv4/test_deepseek_v4.py @@ -338,6 +338,172 @@ class TestDSV4AttentionBackendCorrectness(CustomTestCase): run_dsv4_eagle_draft_extend_cuda_graph_runner_case(self, case) +class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase): + """CPU-only checks for the DSV4 BCG metadata replay contract.""" + + def _make_core_metadata(self, base: int): + from sglang.srt.layers.attention.deepseek_v4_backend import DSV4AttnMetadata + + metadata = DSV4AttnMetadata( + page_size=256, + page_table=torch.tensor( + [[base + 1, base + 2], [base + 3, base + 4]], dtype=torch.int32 + ), + raw_out_loc=torch.tensor([base + 5, base + 6], dtype=torch.int32), + cuda_int32_kwargs={"dtype": torch.int32}, + seq_lens_casual=torch.tensor([base + 7, base + 8], dtype=torch.int32), + positions_casual=torch.tensor([base + 9, base + 10], dtype=torch.int32), + swa_page_indices=torch.tensor( + [[base + 11, base + 12], [base + 13, base + 14]], dtype=torch.int32 + ), + swa_topk_lengths=torch.tensor([base + 15, base + 16], dtype=torch.int32), + c4_sparse_topk=128, + ) + metadata.c4_out_loc = torch.tensor([base + 17, base + 18], dtype=torch.int32) + metadata.c128_out_loc = torch.tensor([base + 19, base + 20], dtype=torch.int32) + metadata.c4_topk_lengths_raw = torch.tensor( + [base + 21, base + 22], dtype=torch.int32 + ) + metadata.c4_topk_lengths_clamp1 = torch.tensor( + [base + 23, base + 24], dtype=torch.int32 + ) + metadata.c4_sparse_topk_lengths = torch.tensor( + [base + 25, base + 26], dtype=torch.int32 + ) + metadata.c4_sparse_page_indices = torch.tensor( + [[base + 27, base + 28], [base + 29, base + 30]], dtype=torch.int32 + ) + metadata.c4_sparse_raw_indices = torch.tensor( + [[base + 31, base + 32], [base + 33, base + 34]], dtype=torch.int32 + ) + metadata.c128_page_indices = torch.tensor( + [[base + 35, base + 36], [base + 37, base + 38]], dtype=torch.int32 + ) + metadata.c128_topk_lengths_clamp1 = torch.tensor( + [base + 39, base + 40], dtype=torch.int32 + ) + metadata.c1_flashmla_metadata = object() + metadata.c4_flashmla_metadata = object() + metadata.c128_flashmla_metadata = object() + return metadata + + def test_bcg_is_explicit_and_dsv4_backend_opt_in_only(self): + from sglang.srt.layers.attention.base_attn_backend import AttentionBackend + from sglang.srt.layers.attention.deepseek_v4_backend import ( + DeepseekV4AttnBackend, + ) + from sglang.srt.server_args import ServerArgs + + self.assertFalse(ServerArgs(model_path="dummy").enable_breakable_cuda_graph) + self.assertFalse( + AttentionBackend.use_captured_forward_metadata_for_breakable_cuda_graph + ) + self.assertTrue( + DeepseekV4AttnBackend.use_captured_forward_metadata_for_breakable_cuda_graph + ) + + def test_refresh_replay_metadata_preserves_captured_tensor_storage(self): + capture_metadata = self._make_core_metadata(0) + replay_metadata = self._make_core_metadata(1000) + + tensor_copy_fields = [ + "raw_out_loc", + "seq_lens_casual", + "positions_casual", + "c4_out_loc", + "c128_out_loc", + "c4_topk_lengths_raw", + "c4_topk_lengths_clamp1", + "c4_sparse_topk_lengths", + ] + reference_assign_fields = [ + "page_table", + "swa_page_indices", + "swa_topk_lengths", + "c128_page_indices", + "c128_topk_lengths_clamp1", + "c1_flashmla_metadata", + "c4_flashmla_metadata", + "c128_flashmla_metadata", + ] + + captured_tensor_objects = { + field: getattr(capture_metadata, field) for field in tensor_copy_fields + } + captured_sparse_pages = capture_metadata.c4_sparse_page_indices + captured_sparse_pages_value = captured_sparse_pages.clone() + + capture_metadata.refresh_for_breakable_cuda_graph_replay_(replay_metadata) + + for field in tensor_copy_fields: + self.assertIs( + getattr(capture_metadata, field), captured_tensor_objects[field] + ) + self.assertTrue( + torch.equal( + getattr(capture_metadata, field), getattr(replay_metadata, field) + ), + f"{field} should be copied from the replay metadata", + ) + + for field in reference_assign_fields: + self.assertIs( + getattr(capture_metadata, field), + getattr(replay_metadata, field), + f"{field} should use the replay metadata reference", + ) + + self.assertIs(capture_metadata.c4_sparse_page_indices, captured_sparse_pages) + self.assertTrue( + torch.equal( + capture_metadata.c4_sparse_page_indices, captured_sparse_pages_value + ) + ) + + def test_backend_replay_keeps_captured_metadata_active(self): + from sglang.srt.layers.attention.deepseek_v4_backend import ( + DeepseekV4AttnBackend, + DSV4Metadata, + ) + + capture_metadata = DSV4Metadata( + self._make_core_metadata(0), indexer_metadata=None + ) + replay_metadata = DSV4Metadata( + self._make_core_metadata(1000), indexer_metadata=None + ) + backend = object.__new__(DeepseekV4AttnBackend) + backend.MAX_SEQ_LEN_FOR_CAPTURE = 4096 + calls = [] + + def fake_build_forward_metadata( + forward_batch, *, max_seq_len_override, use_prefill_cuda_graph + ): + calls.append((forward_batch, max_seq_len_override, use_prefill_cuda_graph)) + return replay_metadata + + backend._build_forward_metadata = fake_build_forward_metadata + forward_batch = SimpleNamespace(name="live") + static_forward_batch = SimpleNamespace(name="static") + + backend.prepare_forward_metadata_for_breakable_cuda_graph_replay( + capture_metadata, + forward_batch, + static_forward_batch=static_forward_batch, + ) + + self.assertIs(calls[0][0], static_forward_batch) + self.assertEqual(calls[0][1], backend.MAX_SEQ_LEN_FOR_CAPTURE) + self.assertTrue(calls[0][2]) + self.assertIs(backend.forward_metadata, capture_metadata) + self.assertTrue( + torch.equal( + capture_metadata.core_attn_metadata.seq_lens_casual, + replay_metadata.core_attn_metadata.seq_lens_casual, + ) + ) + + class TestDSV4SwaOutCacheLocResolution(CustomTestCase): """`get_swa_out_cache_loc`: cached fast path vs store-time fallback. 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 f35faf9ce..de4ada42c 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 @@ -156,5 +156,56 @@ class TestDSV4FlashFP4NonMTPB200( kill_process_tree(cls.process.pid) +class TestDSV4FlashFP4BreakableCudaGraphB200( + BasicDecodeCorrectnessMixin, GSM8KMixin, CustomTestCase +): + """BCG recipe: TP=4, DP=4, DeepEP, DP attention, mixed chunk.""" + + 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", + "--dp", + "4", + "--enable-dp-attention", + "--enable-mixed-chunk", + "--enable-breakable-cuda-graph", + "--enforce-piecewise-cuda-graph", + "--moe-a2a-backend", + "deepep", + "--deepep-config", + DEEPEP_CONFIG, + "--chunked-prefill-size", + "4096", + "--piecewise-cuda-graph-max-tokens", + "1024", + "--mem-fraction-static", + "0.80", + "--cuda-graph-max-bs", + "16", + "--max-running-requests", + "128", + "--watchdog-timeout", + "900", + ], + 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()