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