[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_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<int8_t>()
.with_dtype<int8_t, fp8_e4m3_t>()
.with_device(device)
.verify(buf_x);
// 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
# 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",
@@ -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()
@@ -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)
@@ -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,
)
@@ -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,
)
@@ -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,
@@ -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:
@@ -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,
@@ -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,
+86 -10
View File
@@ -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 :],