[BCG] Support breakable CUDA graph for DeepSeek V4 DP attention (#25195)

Co-authored-by: Yuwei An <ayw.sirius19@gmail.com>
This commit is contained in:
YAMY
2026-06-08 13:54:58 -07:00
committed by GitHub
co-authored by Yuwei An
parent 801fe5e0f2
commit ca66e6fb5e
13 changed files with 726 additions and 66 deletions
@@ -155,8 +155,10 @@ struct MegaMoEPreDispatchKernel {
.with_dtype<float>() .with_dtype<float>()
.with_device(device) .with_device(device)
.verify(topk_weights); .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 TensorMatcher({P, H}) // buf.x
.with_dtype<int8_t>() .with_dtype<int8_t, fp8_e4m3_t>()
.with_device(device) .with_device(device)
.verify(buf_x); .verify(buf_x);
// buf.x_sf is the contiguous row-major int32 view from DeepGEMM's mega // buf.x_sf is the contiguous row-major int32 view from DeepGEMM's mega
@@ -381,6 +381,14 @@ class TboDPAttentionPreparer:
self.enable_two_batch_overlap = enable_two_batch_overlap 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: if local_batch is not None:
token_num_per_seq = get_token_num_per_seq( token_num_per_seq = get_token_num_per_seq(
forward_mode=local_batch.forward_mode, spec_info=local_batch.spec_info forward_mode=local_batch.forward_mode, spec_info=local_batch.spec_info
@@ -692,6 +700,7 @@ class TboForwardBatchPreparer:
"all_extend_in_batch", "all_extend_in_batch",
"return_logprob", "return_logprob",
"can_run_dp_cuda_graph", "can_run_dp_cuda_graph",
"can_run_dp_breakable_cuda_graph",
"dp_padding_mode", "dp_padding_mode",
"global_forward_mode", "global_forward_mode",
"is_prefill_only", "is_prefill_only",
@@ -85,10 +85,40 @@ class AttentionBackend(ABC):
# Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum. # Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum.
needs_cpu_seq_lens: bool = True 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): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
"""Init the global shared states for cuda graph.""" """Init the global shared states for cuda graph."""
raise NotImplementedError() 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): def get_cuda_graph_seq_len_fill_value(self):
"""Get the fill value for padded seq lens. Typically, it is 0 or 1.""" """Get the fill value for padded seq lens. Typically, it is 0 or 1."""
raise NotImplementedError() raise NotImplementedError()
@@ -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): def init_compression_metadata(self):
assert self.page_table.dim() == 2 assert self.page_table.dim() == 2
assert ( assert (
@@ -312,6 +353,24 @@ class DSV4Metadata:
) )
self.sparse_prefill_cache = None 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 @dataclass
class DSV4RawVerifyMetadata: class DSV4RawVerifyMetadata:
@@ -360,6 +419,8 @@ class _GraphBucket(enum.Enum):
class DeepseekV4AttnBackend( class DeepseekV4AttnBackend(
AttentionBackend, C4IndexerBackendMixin, CompressorBackendMixin AttentionBackend, C4IndexerBackendMixin, CompressorBackendMixin
): ):
use_captured_forward_metadata_for_breakable_cuda_graph: bool = True
def __init__( def __init__(
self, self,
model_runner: ModelRunner, model_runner: ModelRunner,
@@ -477,6 +538,7 @@ class DeepseekV4AttnBackend(
num_tokens: int, num_tokens: int,
extend_seq_lens: torch.Tensor, extend_seq_lens: torch.Tensor,
extend_seq_lens_cpu: List[int], extend_seq_lens_cpu: List[int],
extend_start_loc: Optional[torch.Tensor] = None,
need_compress: bool = True, need_compress: bool = True,
use_prefill_cuda_graph: bool = False, use_prefill_cuda_graph: bool = False,
) -> DSV4Metadata: ) -> DSV4Metadata:
@@ -486,6 +548,9 @@ class DeepseekV4AttnBackend(
extend_seq_lens=extend_seq_lens_cpu, extend_seq_lens=extend_seq_lens_cpu,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
padded_num_tokens=out_cache_loc.shape[0], 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( core_attn_metadata = self.make_core_attn_metadata(
req_to_token=self.req_to_token, req_to_token=self.req_to_token,
@@ -504,23 +569,48 @@ class DeepseekV4AttnBackend(
if not need_compress: if not need_compress:
create = _create_dummy_paged_compress_data create = _create_dummy_paged_compress_data
else: else:
create = functools.partial(
create_paged_compressor_data, def create(compress_ratio: Literal[4, 128]):
is_prefill=True, # Online c128 uses a different planner that cannot be created in
token_to_kv_pool=self.token_to_kv_pool, # prefill cuda-graph mode. Keep c4 graph-friendly while matching
req_to_token=self.req_to_token, # c128's existing online path.
req_pool_indices=req_pool_indices, use_graph_plan = use_prefill_cuda_graph and not (
seq_lens=seq_lens, compress_ratio == 128 and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()
seq_lens_cpu=seq_lens_cpu, )
extend_lens=extend_seq_lens, if use_graph_plan:
extend_lens_cpu=extend_seq_lens_cpu, return create_paged_compressor_data(
use_prefill_cuda_graph=use_prefill_cuda_graph, 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( return DSV4Metadata(
core_attn_metadata, core_attn_metadata,
indexer_metadata, indexer_metadata,
c4_compress_metadata=create(compress_ratio=4), c4_compress_metadata=c4_compress_metadata,
c128_compress_metadata=create(compress_ratio=128), c128_compress_metadata=c128_compress_metadata,
) )
def init_forward_metadata_target_verify( def init_forward_metadata_target_verify(
@@ -582,6 +672,7 @@ class DeepseekV4AttnBackend(
num_tokens=num_tokens, num_tokens=num_tokens,
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
extend_start_loc=None,
need_compress=True, need_compress=True,
use_prefill_cuda_graph=use_prefill_cuda_graph, use_prefill_cuda_graph=use_prefill_cuda_graph,
) )
@@ -689,6 +780,7 @@ class DeepseekV4AttnBackend(
num_tokens=num_tokens, num_tokens=num_tokens,
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
extend_start_loc=None,
need_compress=False, need_compress=False,
use_prefill_cuda_graph=use_prefill_cuda_graph, 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(): if self.mtp_enabled and forward_batch.forward_mode.is_idle():
return 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 req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens.to(torch.int32) seq_lens = forward_batch.seq_lens.to(torch.int32)
seq_lens_cpu = forward_batch.seq_lens_cpu 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 self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0
assert seq_lens_cpu is not None 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(): if forward_batch.forward_mode.is_decode_or_idle():
# DSv4 bakes this step's KV write target (c4/c128) into metadata, # 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), num_tokens=sum(extend_seq_lens_cpu),
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
extend_start_loc=forward_batch.extend_start_loc,
need_compress=not is_draft, need_compress=not is_draft,
use_prefill_cuda_graph=use_prefill_cuda_graph,
) )
else: else:
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
self.forward_metadata = metadata return metadata
self.init_forward_metadata_in_graph(forward_batch)
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: def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None:
self.cuda_graph_metadata_of_bucket_and_bs: Dict[ 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_page_indices = core_attn_metadata.swa_page_indices
swa_topk_lengths = core_attn_metadata.swa_topk_lengths swa_topk_lengths = core_attn_metadata.swa_topk_lengths
if self.mtp_enabled: def match_num_queries(x, value):
if swa_page_indices.shape[0] != q.shape[0]: if x is None or x.shape[0] == q.shape[0]:
swa_page_indices = _pad_tensor_to_size( return x
swa_page_indices, q.shape[0], value=0 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_page_indices = match_num_queries(swa_page_indices, value=0)
swa_topk_lengths = _pad_tensor_to_size( swa_topk_lengths = match_num_queries(swa_topk_lengths, value=1)
swa_topk_lengths, q.shape[0], 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: if q.ndim == 3:
q = q.unsqueeze(1) q = q.unsqueeze(1)
@@ -1281,7 +1418,24 @@ class DeepseekV4AttnBackend(
extend_seq_lens: List[int], extend_seq_lens: List[int],
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
padded_num_tokens: Optional[int], 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]: ) -> 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) seq_lens_casual = torch.empty(num_tokens, **self.cuda_int32_kwargs)
idx_to_req_repeated = torch.empty(num_tokens, **self.cuda_int32_kwargs) idx_to_req_repeated = torch.empty(num_tokens, **self.cuda_int32_kwargs)
offset = 0 offset = 0
@@ -1309,6 +1463,48 @@ class DeepseekV4AttnBackend(
return seq_lens_casual, req_pool_indices_repeated 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( def expand_extend_with_same_length(
self, self,
bs: int, bs: int,
@@ -1467,6 +1663,35 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend):
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata(forward_batch) 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): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
for i in range(self.speculative_num_steps): for i in range(self.speculative_num_steps):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
@@ -179,6 +179,8 @@ class CompressorBackendMixin:
if compressor.ratio == 4 if compressor.ratio == 4
else core_metadata.c128_out_loc 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(): if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
token_to_kv_pool.set_extra_key_buffer_fused( token_to_kv_pool.set_extra_key_buffer_fused(
layer_id=layer_id, layer_id=layer_id,
@@ -202,16 +204,19 @@ class CompressorBackendMixin:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
new_compressed_kv = compressor(x, forward_batch, attn_backend=self) 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: if self.enable_deepseek_v4_fp4_indexer:
token_to_kv_pool.set_index_k_fp4( token_to_kv_pool.set_index_k_fp4(
layer_id=layer_id, layer_id=layer_id,
loc=self.forward_metadata.core_metadata.c4_out_loc, loc=out_loc,
cache_k=new_compressed_kv, cache_k=new_compressed_kv,
) )
elif envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get(): elif envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
token_to_kv_pool.set_index_k_fused( token_to_kv_pool.set_index_k_fused(
layer_id=layer_id, layer_id=layer_id,
loc=self.forward_metadata.core_metadata.c4_out_loc, loc=out_loc,
cache_k=new_compressed_kv, cache_k=new_compressed_kv,
) )
else: else:
@@ -220,7 +225,7 @@ class CompressorBackendMixin:
) )
token_to_kv_pool.set_index_k_scale_buffer( token_to_kv_pool.set_index_k_scale_buffer(
layer_id=layer_id, layer_id=layer_id,
loc=self.forward_metadata.core_metadata.c4_out_loc, loc=out_loc,
index_k=new_compressed_kv_fp8, index_k=new_compressed_kv_fp8,
index_k_scale=new_compressed_kv_scale, index_k_scale=new_compressed_kv_scale,
) )
@@ -455,13 +455,22 @@ class C4IndexerBackendMixin:
assert isinstance(indexer_metadata, PagedIndexerMetadata) 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: if enable_multi_stream:
q_indexer, weights, c4_indexer_kv_cache = ( q_indexer, weights, c4_indexer_kv_cache = (
self._forward_prepare_multi_stream( self._forward_prepare_multi_stream(
x=x, x=x,
q_lora=q_lora, q_lora=q_lora,
c4_indexer=c4_indexer, c4_indexer=c4_indexer,
positions=core_metadata.positions, positions=positions,
forward_batch=forward_batch, forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool, token_to_kv_pool=token_to_kv_pool,
alt_streams=alt_streams, alt_streams=alt_streams,
@@ -474,7 +483,7 @@ class C4IndexerBackendMixin:
x=x, x=x,
q_lora=q_lora, q_lora=q_lora,
c4_indexer=c4_indexer, c4_indexer=c4_indexer,
positions=core_metadata.positions, positions=positions,
forward_batch=forward_batch, forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool, token_to_kv_pool=token_to_kv_pool,
skip_compressor=skip_compressor, skip_compressor=skip_compressor,
@@ -519,7 +528,22 @@ class C4IndexerBackendMixin:
else: else:
from deep_gemm import fp8_paged_mqa_logits as fn 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 = ( _use_tilelang = (
envs.SGLANG_OPT_USE_TILELANG_INDEXER.get() and not use_fp4_indexer envs.SGLANG_OPT_USE_TILELANG_INDEXER.get() and not use_fp4_indexer
) )
@@ -531,7 +555,7 @@ class C4IndexerBackendMixin:
c4_indexer_kv_cache, c4_indexer_kv_cache,
weights, weights,
_c4sl, _c4sl,
indexer_metadata.page_table, page_table,
indexer_metadata.deep_gemm_metadata, indexer_metadata.deep_gemm_metadata,
indexer_metadata.max_c4_seq_len, indexer_metadata.max_c4_seq_len,
False, False,
@@ -551,10 +575,10 @@ class C4IndexerBackendMixin:
raw_indices = None raw_indices = None
if capture_enabled: 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: elif hisparse_decode:
raw_indices = hisparse_coordinator.raw_indices_buffer[ 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: elif core_metadata.c4_sparse_raw_indices is not None:
raw_indices = core_metadata.c4_sparse_raw_indices raw_indices = core_metadata.c4_sparse_raw_indices
@@ -562,27 +586,27 @@ class C4IndexerBackendMixin:
if envs.SGLANG_TOPK_TRANSFORM_512_TORCH.get(): if envs.SGLANG_TOPK_TRANSFORM_512_TORCH.get():
topk_transform_512_pytorch_vectorized( topk_transform_512_pytorch_vectorized(
logits, logits,
indexer_metadata.c4_seq_lens, c4_seq_lens,
core_metadata.page_table, page_table,
core_metadata.c4_sparse_page_indices, c4_sparse_page_indices,
indexer_metadata.c4_page_size, indexer_metadata.c4_page_size,
raw_indices, raw_indices,
) )
elif envs.SGLANG_OPT_USE_TOPK_V2.get() and raw_indices is None: elif envs.SGLANG_OPT_USE_TOPK_V2.get() and raw_indices is None:
topk_transform_512_v2( topk_transform_512_v2(
logits, logits,
indexer_metadata.c4_seq_lens, c4_seq_lens,
core_metadata.page_table, page_table,
core_metadata.c4_sparse_page_indices, c4_sparse_page_indices,
indexer_metadata.c4_page_size, indexer_metadata.c4_page_size,
indexer_metadata.topk_metadata, indexer_metadata.topk_metadata,
) )
else: else:
topk_transform_512( topk_transform_512(
logits, logits,
indexer_metadata.c4_seq_lens, c4_seq_lens,
core_metadata.page_table, page_table,
core_metadata.c4_sparse_page_indices, c4_sparse_page_indices,
indexer_metadata.c4_page_size, indexer_metadata.c4_page_size,
raw_indices, raw_indices,
) )
@@ -1667,6 +1667,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
is_extend_in_batch: bool = False is_extend_in_batch: bool = False
all_extend_in_batch: bool = False # plumbing for downstream forks (PR #19639) all_extend_in_batch: bool = False # plumbing for downstream forks (PR #19639)
can_run_dp_cuda_graph: bool = False can_run_dp_cuda_graph: bool = False
can_run_dp_breakable_cuda_graph: bool = False
tbo_split_seq_index: Optional[int] = None tbo_split_seq_index: Optional[int] = None
# For processing logprobs # For processing logprobs
@@ -2749,6 +2750,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
global_num_tokens=self.global_num_tokens, global_num_tokens=self.global_num_tokens,
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob, 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_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, is_extend_in_batch=self.is_extend_in_batch,
all_extend_in_batch=self.all_extend_in_batch, all_extend_in_batch=self.all_extend_in_batch,
is_prefill_only=self.is_prefill_only, is_prefill_only=self.is_prefill_only,
@@ -39,6 +39,7 @@ class MLPSyncBatchInfo:
is_extend_in_batch: bool is_extend_in_batch: bool
local_can_run_tbo: bool local_can_run_tbo: bool
local_forward_mode: int local_forward_mode: int
can_run_breakable_cuda_graph: bool
# some gathered elements # some gathered elements
tp0_info: torch.Tensor = None tp0_info: torch.Tensor = None
@@ -57,6 +58,7 @@ class MLPSyncBatchInfo:
int(self.is_extend_in_batch), int(self.is_extend_in_batch),
int(self.local_can_run_tbo), int(self.local_can_run_tbo),
self.local_forward_mode, self.local_forward_mode,
int(self.can_run_breakable_cuda_graph),
], ],
device=device, device=device,
dtype=dtype, dtype=dtype,
@@ -71,6 +73,7 @@ class MLPSyncBatchInfo:
0, # is_extend_in_batch 0, # is_extend_in_batch
1, # local_can_run_tbo 1, # local_can_run_tbo
ForwardMode.IDLE.value, # local_forward_mode ForwardMode.IDLE.value, # local_forward_mode
0, # can_run_breakable_cuda_graph
], ],
device=device, device=device,
dtype=dtype, dtype=dtype,
@@ -79,7 +82,7 @@ class MLPSyncBatchInfo:
def all_gather(self, device, group: torch.distributed.ProcessGroup): def all_gather(self, device, group: torch.distributed.ProcessGroup):
local_info_tensor = self._get_local_tensor(device=device) local_info_tensor = self._get_local_tensor(device=device)
global_info_tensor = torch.empty( 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, dtype=torch.int64,
device=device, device=device,
) )
@@ -95,7 +98,7 @@ class MLPSyncBatchInfo:
tp_active_ranks = get_tp_group().active_ranks tp_active_ranks = get_tp_group().active_ranks
# Set fallback values for inactive 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) tp_info[tp_active_ranks == 0] = self._get_fallback_tensor(device=device)
tp0_info = global_info_tensor[:, 0, :] tp0_info = global_info_tensor[:, 0, :]
@@ -106,6 +109,7 @@ class MLPSyncBatchInfo:
self.global_num_tokens_for_logprob = cpu_data[:, 1].tolist() self.global_num_tokens_for_logprob = cpu_data[:, 1].tolist()
self.can_cuda_graph = bool(tp0_info[:, 2].min().item()) self.can_cuda_graph = bool(tp0_info[:, 2].min().item())
self.is_extend_in_batch = bool(tp0_info[:, 3].max().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: if _ENABLE_METRICS_DP_ATTENTION:
self.dp_cooperation_info = DPCooperationInfo.create(tp0_info[:, 5].tolist()) self.dp_cooperation_info = DPCooperationInfo.create(tp0_info[:, 5].tolist())
@@ -132,6 +136,7 @@ def _update_gather_batch(
# Check forward mode for cuda graph # Check forward mode for cuda graph
batch.can_run_dp_cuda_graph = mlp_sync_info.can_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( 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_decode_or_idle()
or local_batch.forward_mode.is_prebuilt() or local_batch.forward_mode.is_prebuilt()
) and not disable_cuda_graph ) 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 is_extend_in_batch = local_batch.forward_mode.is_extend() if local_batch else False
if local_batch is not None: if local_batch is not None:
@@ -206,6 +216,7 @@ def prepare_mlp_sync_batch_raw(
is_extend_in_batch=is_extend_in_batch, is_extend_in_batch=is_extend_in_batch,
local_can_run_tbo=local_can_run_tbo, local_can_run_tbo=local_can_run_tbo,
local_forward_mode=local_forward_mode, local_forward_mode=local_forward_mode,
can_run_breakable_cuda_graph=can_run_breakable_cuda_graph,
) )
if not skip_all_gather: if not skip_all_gather:
@@ -36,10 +36,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
set_graph_pool_id, set_graph_pool_id,
) )
from sglang.srt.distributed.parallel_state import graph_capture from sglang.srt.distributed.parallel_state import graph_capture
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import set_dp_buffer_len, set_is_extend_in_batch
set_dp_buffer_len,
set_is_extend_in_batch,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.pooler import EmbeddingPoolerOutput from sglang.srt.layers.pooler import EmbeddingPoolerOutput
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import ( 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.attention_layers = model_runner.attention_layers
self.moe_layers = model_runner.moe_layers self.moe_layers = model_runner.moe_layers
self.moe_fusions = model_runner.moe_fusions 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 # Resolve the inner transformer-stack module (the same boundary PCG draws
# via patch_model). At replay we monkey-patch this module's forward with # via patch_model). At replay we monkey-patch this module's forward with
@@ -167,6 +168,18 @@ class BreakableCudaGraphRunner:
self.raw_num_tokens = 0 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): def _init_buffers(self, model_runner):
"""Initialize input buffers.""" """Initialize input buffers."""
from sglang.srt.model_executor.cuda_graph_buffer_registry import ( from sglang.srt.model_executor.cuda_graph_buffer_registry import (
@@ -317,12 +330,47 @@ class BreakableCudaGraphRunner:
"""Warmup the model with a forward pass.""" """Warmup the model with a forward pass."""
num_tokens = self.capture_num_tokens[0] num_tokens = self.capture_num_tokens[0]
forward_batch = self._build_capture_forward_batch(num_tokens) forward_batch = self._build_capture_forward_batch(num_tokens)
with forward_context( with (
ForwardContext(attn_backend=self.model_runner.attn_backend) 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) 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): def _capture_all(self):
"""Capture breakable CUDA graphs for all token sizes.""" """Capture breakable CUDA graphs for all token sizes."""
with ( with (
@@ -364,6 +412,13 @@ class BreakableCudaGraphRunner:
return False return False
if forward_batch.replace_embeds is not None: if forward_batch.replace_embeds is not None:
return False 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) num_tokens = len(forward_batch.input_ids)
if forward_batch.return_logprob: if forward_batch.return_logprob:
for start_len, seq_len in zip( for start_len, seq_len in zip(
@@ -377,7 +432,7 @@ class BreakableCudaGraphRunner:
def _capture_one(self, num_tokens, pool, stream): def _capture_one(self, num_tokens, pool, stream):
"""Capture a breakable CUDA graph for one token size.""" """Capture a breakable CUDA graph for one token size."""
forward_batch = self._build_capture_forward_batch(num_tokens) 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(): def run_once():
return self._run_forward(forward_batch, num_tokens) return self._run_forward(forward_batch, num_tokens)
@@ -450,7 +505,9 @@ class BreakableCudaGraphRunner:
original_layer_forward = self.layer_model.forward original_layer_forward = self.layer_model.forward
self.layer_model.forward = replay_layer_forward self.layer_model.forward = replay_layer_forward
try: 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( with set_forward_context(
static_forward_batch, static_forward_batch,
self.attention_layers, self.attention_layers,
@@ -345,6 +345,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Mirrors ScheduleBatch.all_extend_in_batch; kept for downstream forks. # Mirrors ScheduleBatch.all_extend_in_batch; kept for downstream forks.
all_extend_in_batch: bool = False all_extend_in_batch: bool = False
can_run_dp_cuda_graph: bool = False can_run_dp_cuda_graph: bool = False
can_run_dp_breakable_cuda_graph: bool = False
global_forward_mode: Optional[ForwardMode] = None global_forward_mode: Optional[ForwardMode] = None
# For two-batch overlap # For two-batch overlap
@@ -647,6 +648,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
is_extend_in_batch=batch.is_extend_in_batch, is_extend_in_batch=batch.is_extend_in_batch,
all_extend_in_batch=batch.all_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_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, global_forward_mode=batch.global_forward_mode,
is_prefill_only=batch.is_prefill_only, is_prefill_only=batch.is_prefill_only,
spec_algorithm=batch.spec_algorithm, spec_algorithm=batch.spec_algorithm,
+86 -10
View File
@@ -27,6 +27,8 @@ from sglang.jit_kernel.dsv4 import (
fused_q_norm_rope, fused_q_norm_rope,
fused_rope_inplace, 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.configs.deepseek_v4 import DeepSeekV4Config
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, 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.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.mem_cache.memory_pool import RadixAttention 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 ( from sglang.srt.model_executor.cuda_graph_runner import (
compile_in_capture_mode, compile_in_capture_mode,
get_is_capture_mode, get_is_capture_mode,
@@ -114,6 +122,7 @@ from sglang.srt.utils import (
log_info_on_rank0, log_info_on_rank0,
make_layers, make_layers,
) )
from sglang.srt.utils.custom_op import register_custom_op
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -191,6 +200,57 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch 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 @triton.jit
def _rms_normalize_kernel( def _rms_normalize_kernel(
x_ptr, x_ptr,
@@ -889,17 +949,33 @@ class MQALayer(nn.Module):
# tell the backend to skip its own store_cache. When `kv is None` # 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 # (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`. # 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 attn_k = kv if kv is not None else q
o = attn_backend.forward( save_kv_cache = False
q=q_padded if q_padded is not None else q, if forward_batch.forward_mode.is_extend() and is_in_breakable_cuda_graph():
k=attn_k, o = attn_q.new_empty(
v=attn_k, (*attn_q.shape[:-1], self.attn_mqa.v_head_dim),
layer=self.attn_mqa, )
forward_batch=forward_batch, bcg_deepseek_v4_attention_with_output(
compress_ratio=self.compress_ratio, attn_q,
attn_sink=self.attn_sink, attn_k,
save_kv_cache=False, 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, :] o = o[:, tp_slice, :]
fused_rope_inplace( fused_rope_inplace(
o[..., -self.qk_rope_head_dim :], o[..., -self.qk_rope_head_dim :],
@@ -338,6 +338,172 @@ class TestDSV4AttentionBackendCorrectness(CustomTestCase):
run_dsv4_eagle_draft_extend_cuda_graph_runner_case(self, case) 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): class TestDSV4SwaOutCacheLocResolution(CustomTestCase):
"""`get_swa_out_cache_loc`: cached fast path vs store-time fallback. """`get_swa_out_cache_loc`: cached fast path vs store-time fallback.
@@ -156,5 +156,56 @@ class TestDSV4FlashFP4NonMTPB200(
kill_process_tree(cls.process.pid) 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__": if __name__ == "__main__":
unittest.main() unittest.main()