[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:
@@ -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,
|
||||
|
||||
@@ -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 :],
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user