[ROCm][DSV4] Enable breakable CUDA graph prefill (#37810)
Co-authored-by: Duyi-Wang <duyi.wang@amd.com>
This commit is contained in:
co-authored by
Duyi-Wang
parent
ad94978adf
commit
a813224e78
@@ -142,10 +142,11 @@ def expand_prefill_causally(
|
|||||||
seq_lens_casual = torch.nn.functional.pad(
|
seq_lens_casual = torch.nn.functional.pad(
|
||||||
seq_lens_casual, (0, pad_size), value=1
|
seq_lens_casual, (0, pad_size), value=1
|
||||||
)
|
)
|
||||||
req_pool_indices_repeated = torch.nn.functional.pad(
|
req_pool_indices_repeated = torch.cat(
|
||||||
req_pool_indices_repeated,
|
(
|
||||||
(0, pad_size),
|
req_pool_indices_repeated,
|
||||||
value=req_pool_indices_repeated[-1].item(),
|
req_pool_indices_repeated[-1:].expand(pad_size),
|
||||||
|
)
|
||||||
)
|
)
|
||||||
return ExpandPrefillCausallyResult(
|
return ExpandPrefillCausallyResult(
|
||||||
seq_lens_casual=seq_lens_casual,
|
seq_lens_casual=seq_lens_casual,
|
||||||
|
|||||||
@@ -153,6 +153,10 @@ class AttentionBackend(ABC):
|
|||||||
# object during capture, and refresh its dynamic fields before each replay.
|
# object during capture, and refresh its dynamic fields before each replay.
|
||||||
use_captured_forward_metadata_for_breakable_cuda_graph: bool = False
|
use_captured_forward_metadata_for_breakable_cuda_graph: bool = False
|
||||||
|
|
||||||
|
# Backends may keep MIXED prefill eager under DP attention when replaying
|
||||||
|
# the EXTEND graph is a known serving-performance regression.
|
||||||
|
prefer_eager_mixed_prefill_under_dp_attention: bool = False
|
||||||
|
|
||||||
# True when prefill graph metadata can use ForwardBatch.max_seq_len_override.
|
# True when prefill graph metadata can use ForwardBatch.max_seq_len_override.
|
||||||
supports_prefill_cuda_graph_max_context_size: bool = False
|
supports_prefill_cuda_graph_max_context_size: bool = False
|
||||||
|
|
||||||
|
|||||||
@@ -147,6 +147,31 @@ class UnifiedKvMetadata:
|
|||||||
assign_fields=[],
|
assign_fields=[],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def refresh_for_breakable_cuda_graph_replay_(
|
||||||
|
self, other: UnifiedKvMetadata
|
||||||
|
) -> None:
|
||||||
|
copy_metadata(
|
||||||
|
src=other,
|
||||||
|
dst=self,
|
||||||
|
check_eq_fields=[],
|
||||||
|
copy_fields=[
|
||||||
|
"swa_loc",
|
||||||
|
"swa_indices",
|
||||||
|
"swa_indptr",
|
||||||
|
"hca_indices",
|
||||||
|
"hca_indptr",
|
||||||
|
"csa_indices",
|
||||||
|
"csa_indptr",
|
||||||
|
"pf_state_slot",
|
||||||
|
"pf_chunk_start",
|
||||||
|
"pf_cu_q",
|
||||||
|
"pf_final_pos",
|
||||||
|
"verify_store_state_slot",
|
||||||
|
"c4_out_loc",
|
||||||
|
"c128_out_loc",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DSV4AttnMetadata:
|
class DSV4AttnMetadata:
|
||||||
@@ -237,6 +262,51 @@ 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",
|
||||||
|
"swa_out_cache_loc",
|
||||||
|
"c4_out_loc",
|
||||||
|
"c128_out_loc",
|
||||||
|
"page_table",
|
||||||
|
"swa_page_indices",
|
||||||
|
"swa_topk_lengths",
|
||||||
|
"c128_page_indices",
|
||||||
|
"c128_topk_lengths_clamp1",
|
||||||
|
"c128_topk_lengths_raw",
|
||||||
|
"c4_topk_lengths_raw",
|
||||||
|
"c4_topk_lengths_clamp1",
|
||||||
|
"c4_sparse_topk_lengths",
|
||||||
|
"c4_sparse_topk_lengths_raw",
|
||||||
|
"c4_sparse_page_indices",
|
||||||
|
"c4_sparse_raw_indices",
|
||||||
|
]
|
||||||
|
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 src_val is not None and dst_val is not None, (
|
||||||
|
f"{field_name=} {src_val=} {dst_val=}"
|
||||||
|
)
|
||||||
|
dst_val.copy_(src_val)
|
||||||
|
|
||||||
|
if self.unified is None and other.unified is None:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
assert self.unified is not None and other.unified is not None
|
||||||
|
self.unified.refresh_for_breakable_cuda_graph_replay_(other.unified)
|
||||||
|
|
||||||
|
self.c0_flashmla_metadata = other.c0_flashmla_metadata
|
||||||
|
self.c4_flashmla_metadata = other.c4_flashmla_metadata
|
||||||
|
self.c128_flashmla_metadata = other.c128_flashmla_metadata
|
||||||
|
|
||||||
def init_compression_metadata(self, unified_swa_pages: int = 0):
|
def init_compression_metadata(self, unified_swa_pages: int = 0):
|
||||||
assert self.page_table.dim() == 2
|
assert self.page_table.dim() == 2
|
||||||
assert self.raw_out_loc.shape == self.seq_lens_casual.shape, (
|
assert self.raw_out_loc.shape == self.seq_lens_casual.shape, (
|
||||||
@@ -383,6 +453,37 @@ class DSV4Metadata:
|
|||||||
self.c128_compress_metadata, src=other.c128_compress_metadata
|
self.c128_compress_metadata, src=other.c128_compress_metadata
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def refresh_for_breakable_cuda_graph_replay_(self, other: DSV4Metadata) -> None:
|
||||||
|
self.core_attn_metadata.refresh_for_breakable_cuda_graph_replay_(
|
||||||
|
other.core_attn_metadata
|
||||||
|
)
|
||||||
|
maybe_copy_inplace(self.indexer_metadata, src=other.indexer_metadata)
|
||||||
|
maybe_copy_inplace(self.c4_compress_metadata, src=other.c4_compress_metadata)
|
||||||
|
maybe_copy_inplace(
|
||||||
|
self.c128_compress_metadata, src=other.c128_compress_metadata
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.fp4_k_write_metadata is None and other.fp4_k_write_metadata is None:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
assert (
|
||||||
|
self.fp4_k_write_metadata is not None
|
||||||
|
and other.fp4_k_write_metadata is not None
|
||||||
|
)
|
||||||
|
for captured, replay in zip(
|
||||||
|
self.fp4_k_write_metadata,
|
||||||
|
other.fp4_k_write_metadata,
|
||||||
|
strict=True,
|
||||||
|
):
|
||||||
|
captured.copy_(replay)
|
||||||
|
|
||||||
|
if self.fp4_q_positions is None and other.fp4_q_positions is None:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
assert self.fp4_q_positions is not None
|
||||||
|
assert other.fp4_q_positions is not None
|
||||||
|
self.fp4_q_positions.copy_(other.fp4_q_positions)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DSV4RawVerifyMetadata:
|
class DSV4RawVerifyMetadata:
|
||||||
@@ -439,6 +540,9 @@ class DeepseekV4HipRadixBackend(
|
|||||||
# TboAttnBackend reads this to skip children in the *_graph paths only.
|
# TboAttnBackend reads this to skip children in the *_graph paths only.
|
||||||
tbo_supports_cuda_graph = False
|
tbo_supports_cuda_graph = False
|
||||||
supports_ragged_verify_graph: bool = True
|
supports_ragged_verify_graph: bool = True
|
||||||
|
use_captured_forward_metadata_for_breakable_cuda_graph: bool = True
|
||||||
|
# MIXED BCG replay regresses ROCm DSV4 DP-attention serving throughput.
|
||||||
|
prefer_eager_mixed_prefill_under_dp_attention: bool = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -585,13 +689,22 @@ class DeepseekV4HipRadixBackend(
|
|||||||
need_compress=need_compress,
|
need_compress=need_compress,
|
||||||
is_prefill=True,
|
is_prefill=True,
|
||||||
)
|
)
|
||||||
|
# Normal prefill starts with a conservative exact_num_tokens=False.
|
||||||
|
# Its CPU length mirror proves the exact query count without a D2H sync.
|
||||||
|
host_proves_exact_num_tokens = (
|
||||||
|
need_compress
|
||||||
|
and not attach_decode_streams
|
||||||
|
and extend_seq_lens_cpu is not None
|
||||||
|
and sum(extend_seq_lens_cpu) == num_tokens
|
||||||
|
)
|
||||||
self._attach_unified_kv_prefill_meta(
|
self._attach_unified_kv_prefill_meta(
|
||||||
core_attn_metadata,
|
core_attn_metadata,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
|
req_pool_indices_repeated,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
extend_seq_lens,
|
extend_seq_lens,
|
||||||
num_tokens,
|
num_tokens,
|
||||||
exact_num_tokens=exact_num_tokens,
|
exact_num_tokens=exact_num_tokens or host_proves_exact_num_tokens,
|
||||||
)
|
)
|
||||||
if attach_decode_streams:
|
if attach_decode_streams:
|
||||||
# Target-verify runs through the unified_kv DECODE kernel, so build
|
# Target-verify runs through the unified_kv DECODE kernel, so build
|
||||||
@@ -608,33 +721,41 @@ class DeepseekV4HipRadixBackend(
|
|||||||
)
|
)
|
||||||
if not need_compress:
|
if not need_compress:
|
||||||
create = _create_dummy_paged_compress_data
|
create = _create_dummy_paged_compress_data
|
||||||
elif compress_gpu_plan:
|
|
||||||
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=None,
|
|
||||||
extend_lens=extend_seq_lens,
|
|
||||||
extend_lens_cpu=None,
|
|
||||||
num_q_tokens=num_tokens,
|
|
||||||
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
create = functools.partial(
|
|
||||||
create_paged_compressor_data,
|
def create(compress_ratio: Literal[4, 128]):
|
||||||
is_prefill=True,
|
use_graph_plan = use_prefill_cuda_graph and not (
|
||||||
token_to_kv_pool=self.token_to_kv_pool,
|
compress_ratio == 128 and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()
|
||||||
req_to_token=self.req_to_token,
|
)
|
||||||
req_pool_indices=req_pool_indices,
|
if compress_gpu_plan or use_graph_plan:
|
||||||
seq_lens=seq_lens,
|
return create_paged_compressor_data(
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
compress_ratio=compress_ratio,
|
||||||
extend_lens=extend_seq_lens,
|
is_prefill=True,
|
||||||
extend_lens_cpu=extend_seq_lens_cpu,
|
token_to_kv_pool=self.token_to_kv_pool,
|
||||||
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
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,
|
||||||
|
num_q_tokens=(
|
||||||
|
out_cache_loc.shape[0] if use_graph_plan else num_tokens
|
||||||
|
),
|
||||||
|
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
||||||
|
)
|
||||||
|
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=False,
|
||||||
|
)
|
||||||
|
|
||||||
return DSV4Metadata(
|
return DSV4Metadata(
|
||||||
core_attn_metadata,
|
core_attn_metadata,
|
||||||
indexer_metadata,
|
indexer_metadata,
|
||||||
@@ -1147,10 +1268,13 @@ class DeepseekV4HipRadixBackend(
|
|||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
|
def _build_forward_metadata(
|
||||||
if self.mtp_enabled and forward_batch.forward_mode.is_idle():
|
self,
|
||||||
return
|
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
|
||||||
@@ -1158,7 +1282,11 @@ class DeepseekV4HipRadixBackend(
|
|||||||
|
|
||||||
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 = (
|
||||||
|
max_seq_len_override
|
||||||
|
if max_seq_len_override is not None
|
||||||
|
else int(seq_lens_cpu.max().item())
|
||||||
|
)
|
||||||
|
|
||||||
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,
|
||||||
@@ -1211,16 +1339,61 @@ class DeepseekV4HipRadixBackend(
|
|||||||
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,
|
||||||
exact_num_tokens=is_draft,
|
exact_num_tokens=is_draft,
|
||||||
)
|
)
|
||||||
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
|
||||||
|
|
||||||
|
def init_forward_metadata(self, forward_batch: ForwardBatch) -> None:
|
||||||
|
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)
|
self.init_forward_metadata_in_graph(forward_batch)
|
||||||
self._refresh_fp4_prefill_workspace(forward_batch)
|
self._refresh_fp4_prefill_workspace(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,
|
||||||
|
)
|
||||||
|
self.init_forward_metadata_in_graph(forward_batch)
|
||||||
|
self._refresh_fp4_prefill_workspace(forward_batch)
|
||||||
|
assert isinstance(self.forward_metadata, DSV4Metadata)
|
||||||
|
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:
|
||||||
|
replay_batch = (
|
||||||
|
static_forward_batch if static_forward_batch is not None else forward_batch
|
||||||
|
)
|
||||||
|
replay_metadata = self._build_forward_metadata(
|
||||||
|
replay_batch,
|
||||||
|
max_seq_len_override=self.MAX_SEQ_LEN_FOR_CAPTURE,
|
||||||
|
use_prefill_cuda_graph=True,
|
||||||
|
)
|
||||||
|
self.forward_metadata = replay_metadata
|
||||||
|
self.init_forward_metadata_in_graph(replay_batch)
|
||||||
|
|
||||||
|
assert isinstance(capture_metadata, DSV4Metadata)
|
||||||
|
assert isinstance(replay_metadata, DSV4Metadata)
|
||||||
|
capture_metadata.refresh_for_breakable_cuda_graph_replay_(replay_metadata)
|
||||||
|
self.forward_metadata = capture_metadata
|
||||||
|
self._refresh_fp4_prefill_workspace(replay_batch)
|
||||||
|
|
||||||
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[
|
||||||
_GraphBucket,
|
_GraphBucket,
|
||||||
@@ -1327,6 +1500,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
self,
|
self,
|
||||||
core: DSV4AttnMetadata,
|
core: DSV4AttnMetadata,
|
||||||
req_pool_indices: torch.Tensor,
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_pool_indices_repeated: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
extend_seq_lens: torch.Tensor,
|
extend_seq_lens: torch.Tensor,
|
||||||
num_tokens: int,
|
num_tokens: int,
|
||||||
@@ -1352,11 +1526,33 @@ class DeepseekV4HipRadixBackend(
|
|||||||
)
|
)
|
||||||
if core.unified is None:
|
if core.unified is None:
|
||||||
core.unified = UnifiedKvMetadata()
|
core.unified = UnifiedKvMetadata()
|
||||||
core.unified.pf_state_slot = req_pool_indices[bid]
|
state_slot = req_pool_indices[bid]
|
||||||
core.unified.pf_chunk_start = (seq_lens - extend_seq_lens)[bid]
|
chunk_start = (seq_lens - extend_seq_lens)[bid]
|
||||||
cu_q_per_req = torch.cumsum(extend_seq_lens, dim=0) - extend_seq_lens
|
cu_q_per_req = torch.cumsum(extend_seq_lens, dim=0) - extend_seq_lens
|
||||||
core.unified.pf_cu_q = cu_q_per_req[bid]
|
cu_q = cu_q_per_req[bid]
|
||||||
core.unified.pf_final_pos = (seq_lens - 1)[bid]
|
final_pos = (seq_lens - 1)[bid]
|
||||||
|
|
||||||
|
padded_num_tokens = core.positions_casual.shape[0]
|
||||||
|
assert num_tokens <= padded_num_tokens
|
||||||
|
if num_tokens < padded_num_tokens:
|
||||||
|
pad_size = padded_num_tokens - num_tokens
|
||||||
|
state_slot = torch.cat(
|
||||||
|
(state_slot, req_pool_indices_repeated[num_tokens:padded_num_tokens])
|
||||||
|
)
|
||||||
|
chunk_start = F.pad(chunk_start, (0, pad_size), value=0)
|
||||||
|
cu_q = F.pad(cu_q, (0, pad_size), value=0)
|
||||||
|
# Padded positions are zero. final_pos=win makes the SWA store's
|
||||||
|
# `pos <= final_pos - win` guard skip every padded row.
|
||||||
|
final_pos = F.pad(
|
||||||
|
final_pos,
|
||||||
|
(0, pad_size),
|
||||||
|
value=self.token_to_kv_pool.unified_swa_window,
|
||||||
|
)
|
||||||
|
|
||||||
|
core.unified.pf_state_slot = state_slot
|
||||||
|
core.unified.pf_chunk_start = chunk_start
|
||||||
|
core.unified.pf_cu_q = cu_q
|
||||||
|
core.unified.pf_final_pos = final_pos
|
||||||
|
|
||||||
def _forward_unified_kv(
|
def _forward_unified_kv(
|
||||||
self,
|
self,
|
||||||
@@ -1739,16 +1935,12 @@ class DeepseekV4HipRadixBackend(
|
|||||||
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:
|
swa_page_indices = _match_num_queries(swa_page_indices, q.shape[0], value=0)
|
||||||
if swa_page_indices.shape[0] != q.shape[0]:
|
swa_topk_lengths = _match_num_queries(swa_topk_lengths, q.shape[0], value=1)
|
||||||
swa_page_indices = _pad_tensor_to_size(
|
extra_indices = _match_num_queries(extra_indices, q.shape[0], value=-1)
|
||||||
swa_page_indices, q.shape[0], value=0
|
extra_topk_lengths = _match_num_queries(
|
||||||
)
|
extra_topk_lengths, q.shape[0], value=1
|
||||||
|
)
|
||||||
if swa_topk_lengths.shape[0] != q.shape[0]:
|
|
||||||
swa_topk_lengths = _pad_tensor_to_size(
|
|
||||||
swa_topk_lengths, q.shape[0], value=1
|
|
||||||
)
|
|
||||||
|
|
||||||
if q.ndim == 3:
|
if q.ndim == 3:
|
||||||
q = q.unsqueeze(1)
|
q = q.unsqueeze(1)
|
||||||
@@ -1992,6 +2184,33 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
|
|||||||
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
|
||||||
|
):
|
||||||
|
return [
|
||||||
|
self.attn_backends[
|
||||||
|
i
|
||||||
|
].init_forward_metadata_for_breakable_cuda_graph_capture(forward_batch)
|
||||||
|
for i in range(self.speculative_num_steps - 1)
|
||||||
|
]
|
||||||
|
|
||||||
|
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)
|
||||||
@@ -2001,7 +2220,13 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
|
|||||||
backend.on_after_cuda_graph_warmup()
|
backend.on_after_cuda_graph_warmup()
|
||||||
|
|
||||||
|
|
||||||
def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0):
|
def _match_num_queries(
|
||||||
|
tensor: Optional[torch.Tensor], size: int, *, value: int
|
||||||
|
) -> Optional[torch.Tensor]:
|
||||||
|
if tensor is None or tensor.shape[0] == size:
|
||||||
|
return tensor
|
||||||
|
if tensor.shape[0] > size:
|
||||||
|
return tensor[:size]
|
||||||
if value == 0:
|
if value == 0:
|
||||||
return torch.cat(
|
return torch.cat(
|
||||||
[tensor, tensor.new_zeros(size - tensor.shape[0], *tensor.shape[1:])],
|
[tensor, tensor.new_zeros(size - tensor.shape[0], *tensor.shape[1:])],
|
||||||
|
|||||||
@@ -289,9 +289,9 @@ def _local_prefill_cuda_graph_vote(
|
|||||||
model_config,
|
model_config,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""This rank's vote for the prefill graph (min-reduced across dp
|
"""This rank's vote for the prefill graph (min-reduced across dp
|
||||||
ranks). Extend/mixed batches vote their own replayability; a decode
|
ranks). Extend and mixed batches share the runner's rank-local replay
|
||||||
batch eligible for the decode->extend conversion votes as its 1-token-
|
policy. A decode batch eligible for the decode->extend conversion votes as
|
||||||
extend view, so the vote and the post-sync conversion always agree."""
|
its 1-token-extend view, so the vote and post-sync conversion agree."""
|
||||||
if local_batch is None or local_batch.forward_mode.is_idle():
|
if local_batch is None or local_batch.forward_mode.is_idle():
|
||||||
return True
|
return True
|
||||||
if not coordinated_prefill:
|
if not coordinated_prefill:
|
||||||
@@ -350,6 +350,7 @@ def _local_prefill_cuda_graph_vote(
|
|||||||
capture_hidden_mode=None,
|
capture_hidden_mode=None,
|
||||||
return_logprob=return_logprob,
|
return_logprob=return_logprob,
|
||||||
lora_ineligible=prefill_graph_runner.enable_lora,
|
lora_ineligible=prefill_graph_runner.enable_lora,
|
||||||
|
is_mixed=mode == ForwardMode.MIXED,
|
||||||
batch_max_context_len=(
|
batch_max_context_len=(
|
||||||
int(local_batch.seq_lens_cpu.max().item())
|
int(local_batch.seq_lens_cpu.max().item())
|
||||||
if prefill_graph_runner.max_context_size is not None
|
if prefill_graph_runner.max_context_size is not None
|
||||||
|
|||||||
@@ -302,6 +302,15 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
# --- prefill graph config -------------------------------------
|
# --- prefill graph config -------------------------------------
|
||||||
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||||
self.prefill_backend_name = prefill_config.backend
|
self.prefill_backend_name = prefill_config.backend
|
||||||
|
self.prefer_eager_mixed_prefill = (
|
||||||
|
self.prefill_backend_name == Backend.BREAKABLE
|
||||||
|
and get_parallel().enable_dp_attention
|
||||||
|
and getattr(
|
||||||
|
model_runner.attn_backend,
|
||||||
|
"prefer_eager_mixed_prefill_under_dp_attention",
|
||||||
|
False,
|
||||||
|
)
|
||||||
|
)
|
||||||
# bs in prefill carries the captured shape (token count for
|
# bs in prefill carries the captured shape (token count for
|
||||||
# tc_piecewise) — one shape knob per phase.
|
# tc_piecewise) — one shape knob per phase.
|
||||||
capture_tokens = prefill_config.bs
|
capture_tokens = prefill_config.bs
|
||||||
@@ -1199,6 +1208,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
capture_hidden_mode,
|
capture_hidden_mode,
|
||||||
return_logprob: bool,
|
return_logprob: bool,
|
||||||
lora_ineligible: bool = False,
|
lora_ineligible: bool = False,
|
||||||
|
is_mixed: bool = False,
|
||||||
batch_max_context_len: Optional[int] = None,
|
batch_max_context_len: Optional[int] = None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Rank-local replay eligibility: the single source of truth for
|
"""Rank-local replay eligibility: the single source of truth for
|
||||||
@@ -1215,6 +1225,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
# schedule-time vote derives this from enable_lora alone.
|
# schedule-time vote derives this from enable_lora alone.
|
||||||
if lora_ineligible:
|
if lora_ineligible:
|
||||||
return False
|
return False
|
||||||
|
if is_mixed and getattr(self, "prefer_eager_mixed_prefill", False):
|
||||||
|
return False
|
||||||
if input_embeds is not None:
|
if input_embeds is not None:
|
||||||
return False
|
return False
|
||||||
if replace_embeds is not None:
|
if replace_embeds is not None:
|
||||||
@@ -1295,6 +1307,14 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
forward_batch
|
forward_batch
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
|
is_mixed=any(
|
||||||
|
getattr(forward_batch, field, None) == ForwardMode.MIXED
|
||||||
|
for field in (
|
||||||
|
"forward_mode",
|
||||||
|
"global_forward_mode",
|
||||||
|
"_original_forward_mode",
|
||||||
|
)
|
||||||
|
),
|
||||||
batch_max_context_len=batch_max_context_len,
|
batch_max_context_len=batch_max_context_len,
|
||||||
):
|
):
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -0,0 +1,419 @@
|
|||||||
|
import dataclasses
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||||
|
DeepseekV4HipRadixBackend,
|
||||||
|
DeepseekV4MultiStepBackend,
|
||||||
|
DSV4AttnMetadata,
|
||||||
|
DSV4Metadata,
|
||||||
|
UnifiedKvMetadata,
|
||||||
|
_match_num_queries,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import is_hip
|
||||||
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
|
|
||||||
|
register_amd_ci(est_time=5, suite="stage-b-test-1-gpu-small-amd-mi35x")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(is_hip(), "DeepSeek V4 HIP radix backend requires ROCm")
|
||||||
|
class TestDSV4HipBreakableCudaGraphMetadata(unittest.TestCase):
|
||||||
|
@staticmethod
|
||||||
|
def _make_core_metadata(base: int) -> DSV4AttnMetadata:
|
||||||
|
def tensor(offset: int) -> torch.Tensor:
|
||||||
|
return torch.tensor([base + offset], dtype=torch.int32)
|
||||||
|
|
||||||
|
def fill_optional_tensors(metadata, start: int) -> int:
|
||||||
|
for metadata_field in dataclasses.fields(metadata):
|
||||||
|
if "Tensor" not in str(metadata_field.type):
|
||||||
|
continue
|
||||||
|
if getattr(metadata, metadata_field.name, None) is None:
|
||||||
|
setattr(metadata, metadata_field.name, tensor(start))
|
||||||
|
start += 1
|
||||||
|
return start
|
||||||
|
|
||||||
|
metadata = DSV4AttnMetadata(
|
||||||
|
page_size=256,
|
||||||
|
page_table=torch.tensor([[base + 1, base + 2]], dtype=torch.int32),
|
||||||
|
raw_out_loc=torch.tensor([base + 3], dtype=torch.int32),
|
||||||
|
cuda_int32_kwargs={"dtype": torch.int32},
|
||||||
|
seq_lens_casual=torch.tensor([base + 4], dtype=torch.int32),
|
||||||
|
positions_casual=torch.tensor([base + 5], dtype=torch.int32),
|
||||||
|
swa_page_indices=torch.tensor([[base + 6, base + 7]], dtype=torch.int32),
|
||||||
|
swa_topk_lengths=torch.tensor([base + 8], dtype=torch.int32),
|
||||||
|
c4_sparse_topk=512,
|
||||||
|
swa_out_cache_loc=torch.tensor([base + 9], dtype=torch.int32),
|
||||||
|
unified=UnifiedKvMetadata(),
|
||||||
|
)
|
||||||
|
next_offset = fill_optional_tensors(metadata, 10)
|
||||||
|
fill_optional_tensors(metadata.unified, next_offset)
|
||||||
|
metadata.c0_flashmla_metadata = None
|
||||||
|
metadata.c4_flashmla_metadata = None
|
||||||
|
metadata.c128_flashmla_metadata = None
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
def test_backend_opts_into_captured_bcg_metadata(self):
|
||||||
|
self.assertTrue(
|
||||||
|
DeepseekV4HipRadixBackend.use_captured_forward_metadata_for_breakable_cuda_graph
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
DeepseekV4HipRadixBackend.prefer_eager_mixed_prefill_under_dp_attention
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_non_unified_metadata_matches_underfilled_bucket(self):
|
||||||
|
captured = torch.tensor([[1, 2], [3, 4], [5, 6], [7, 8]])
|
||||||
|
replay = _match_num_queries(captured, 3, value=-1)
|
||||||
|
self.assertEqual(replay.tolist(), [[1, 2], [3, 4], [5, 6]])
|
||||||
|
|
||||||
|
short = torch.tensor([9, 10])
|
||||||
|
replay = _match_num_queries(short, 3, value=1)
|
||||||
|
self.assertEqual(replay.tolist(), [9, 10, 1])
|
||||||
|
self.assertIsNone(_match_num_queries(None, 3, value=0))
|
||||||
|
|
||||||
|
def test_unified_prefill_metadata_pads_to_capture_bucket(self):
|
||||||
|
backend = object.__new__(DeepseekV4HipRadixBackend)
|
||||||
|
backend.token_to_kv_pool = SimpleNamespace(unified_swa_window=128)
|
||||||
|
core = self._make_core_metadata(0)
|
||||||
|
core.positions_casual = torch.tensor([0, 1, 2, 0], dtype=torch.int32)
|
||||||
|
core.unified = None
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch(
|
||||||
|
"sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate."
|
||||||
|
"is_unified_kv_triton",
|
||||||
|
return_value=True,
|
||||||
|
),
|
||||||
|
mock.patch(
|
||||||
|
"sglang.srt.layers.attention.deepseek_v4_backend_hip_radix."
|
||||||
|
"torch.repeat_interleave",
|
||||||
|
wraps=torch.repeat_interleave,
|
||||||
|
) as repeat_interleave,
|
||||||
|
):
|
||||||
|
backend._attach_unified_kv_prefill_meta(
|
||||||
|
core,
|
||||||
|
req_pool_indices=torch.tensor([7, 9], dtype=torch.int32),
|
||||||
|
req_pool_indices_repeated=torch.tensor([7, 9, 9, 9], dtype=torch.int32),
|
||||||
|
seq_lens=torch.tensor([1, 3], dtype=torch.int32),
|
||||||
|
extend_seq_lens=torch.tensor([1, 2], dtype=torch.int32),
|
||||||
|
num_tokens=3,
|
||||||
|
exact_num_tokens=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
repeat_interleave.assert_called_once()
|
||||||
|
self.assertEqual(repeat_interleave.call_args.kwargs["output_size"], 3)
|
||||||
|
self.assertEqual(core.unified.pf_state_slot.tolist(), [7, 9, 9, 9])
|
||||||
|
self.assertEqual(core.unified.pf_chunk_start.tolist(), [0, 1, 1, 0])
|
||||||
|
self.assertEqual(core.unified.pf_cu_q.tolist(), [0, 1, 1, 0])
|
||||||
|
self.assertEqual(core.unified.pf_final_pos.tolist(), [0, 2, 2, 128])
|
||||||
|
|
||||||
|
def test_eager_prefill_marks_host_proven_token_count_exact(self):
|
||||||
|
backend = object.__new__(DeepseekV4HipRadixBackend)
|
||||||
|
backend.req_to_token = torch.zeros((2, 8), dtype=torch.int32)
|
||||||
|
backend.token_to_kv_pool = object()
|
||||||
|
core = self._make_core_metadata(0)
|
||||||
|
extend_start_loc = torch.tensor([0, 1], dtype=torch.int32)
|
||||||
|
backend.make_core_attn_metadata = mock.Mock(return_value=core)
|
||||||
|
backend._attach_unified_kv_prefill_meta = mock.Mock()
|
||||||
|
backend.init_forward_metadata_indexer = mock.Mock(return_value=None)
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch(
|
||||||
|
"sglang.kernels.ops.attention.dsv4_attn_metadata_kernels."
|
||||||
|
"ExpandPrefillCausally.execute",
|
||||||
|
return_value=SimpleNamespace(
|
||||||
|
seq_lens_casual=core.seq_lens_casual,
|
||||||
|
req_pool_indices_repeated=torch.tensor(
|
||||||
|
[7, 9, 9], dtype=torch.int32
|
||||||
|
),
|
||||||
|
),
|
||||||
|
) as expand_prefill,
|
||||||
|
mock.patch(
|
||||||
|
"sglang.srt.layers.attention.deepseek_v4_backend_hip_radix."
|
||||||
|
"create_paged_compressor_data",
|
||||||
|
return_value=None,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
backend.init_forward_metadata_prefill(
|
||||||
|
max_seq_len=4096,
|
||||||
|
req_pool_indices=torch.tensor([7, 9], dtype=torch.int32),
|
||||||
|
seq_lens=torch.tensor([1, 3], dtype=torch.int32),
|
||||||
|
seq_lens_cpu=[1, 3],
|
||||||
|
out_cache_loc=torch.zeros(3, dtype=torch.int64),
|
||||||
|
num_tokens=3,
|
||||||
|
extend_seq_lens=torch.tensor([1, 2], dtype=torch.int32),
|
||||||
|
extend_seq_lens_cpu=[1, 2],
|
||||||
|
extend_start_loc=extend_start_loc,
|
||||||
|
use_prefill_cuda_graph=False,
|
||||||
|
exact_num_tokens=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(
|
||||||
|
expand_prefill.call_args.kwargs["extend_start_loc"], extend_start_loc
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
backend._attach_unified_kv_prefill_meta.call_args.kwargs["exact_num_tokens"]
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_prefill_bcg_uses_bucket_sized_gpu_compressor_plans(self):
|
||||||
|
backend = object.__new__(DeepseekV4HipRadixBackend)
|
||||||
|
backend.req_to_token = torch.zeros((2, 8), dtype=torch.int32)
|
||||||
|
backend.token_to_kv_pool = object()
|
||||||
|
core = self._make_core_metadata(0)
|
||||||
|
core.positions_casual = torch.tensor([0, 1, 2, 0], dtype=torch.int32)
|
||||||
|
backend.make_core_attn_metadata = mock.Mock(return_value=core)
|
||||||
|
backend._attach_unified_kv_prefill_meta = mock.Mock()
|
||||||
|
backend.init_forward_metadata_indexer = mock.Mock(return_value=None)
|
||||||
|
|
||||||
|
with mock.patch(
|
||||||
|
"sglang.srt.layers.attention.deepseek_v4_backend_hip_radix."
|
||||||
|
"create_paged_compressor_data",
|
||||||
|
side_effect=lambda compress_ratio, **kwargs: (compress_ratio, kwargs),
|
||||||
|
) as create_plan:
|
||||||
|
backend.init_forward_metadata_prefill(
|
||||||
|
max_seq_len=4096,
|
||||||
|
req_pool_indices=torch.tensor([7, 9], dtype=torch.int32),
|
||||||
|
seq_lens=torch.tensor([1, 3], dtype=torch.int32),
|
||||||
|
seq_lens_cpu=[1, 3],
|
||||||
|
out_cache_loc=torch.zeros(4, dtype=torch.int64),
|
||||||
|
num_tokens=3,
|
||||||
|
extend_seq_lens=torch.tensor([1, 2], dtype=torch.int32),
|
||||||
|
extend_seq_lens_cpu=[1, 2],
|
||||||
|
use_prefill_cuda_graph=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(create_plan.call_count, 2)
|
||||||
|
for call in create_plan.call_args_list:
|
||||||
|
self.assertIsNone(call.kwargs["seq_lens_cpu"])
|
||||||
|
self.assertIsNone(call.kwargs["extend_lens_cpu"])
|
||||||
|
self.assertEqual(call.kwargs["num_q_tokens"], 4)
|
||||||
|
self.assertTrue(call.kwargs["use_prefill_cuda_graph"])
|
||||||
|
|
||||||
|
def test_gpu_compressor_plan_invalidates_bucket_tail(self):
|
||||||
|
from sglang.kernels.ops.attention.dsv4 import CompressorPrefillPlan
|
||||||
|
from sglang.test.kernels.deepseek_v4.common import make_paged_context
|
||||||
|
|
||||||
|
seq_lens = torch.tensor([1, 3], dtype=torch.int64, device="cuda")
|
||||||
|
extend_lens = torch.tensor([1, 2], dtype=torch.int64, device="cuda")
|
||||||
|
for compress_ratio in (4, 128):
|
||||||
|
with self.subTest(compress_ratio=compress_ratio):
|
||||||
|
context = make_paged_context(bs=2, compress_ratio=compress_ratio)
|
||||||
|
plan = CompressorPrefillPlan.generate(
|
||||||
|
compress_ratio=compress_ratio,
|
||||||
|
req_pool_indices=context.req_pool_indices,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
extend_lens=extend_lens,
|
||||||
|
req_to_token=context.req_to_token,
|
||||||
|
full_to_state=context.full_to_swa,
|
||||||
|
swa_page_size=context.swa_page_size,
|
||||||
|
ring_size=context.ring_size,
|
||||||
|
num_q_tokens=4,
|
||||||
|
use_cuda_graph=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(plan.plan_c.shape, (4, 16))
|
||||||
|
self.assertEqual(plan.plan_w.shape, (4, 8))
|
||||||
|
ragged_ids = plan.plan_w.view(torch.uint32).view(-1, 2)[:, 0]
|
||||||
|
self.assertEqual(ragged_ids[:3].cpu().tolist(), [0, 1, 2])
|
||||||
|
self.assertEqual(int(ragged_ids[3].item()), 0xFFFFFFFF)
|
||||||
|
|
||||||
|
def test_capture_builds_graph_compatible_metadata_and_workspace(self):
|
||||||
|
capture_metadata = DSV4Metadata(object(), indexer_metadata=None)
|
||||||
|
backend = object.__new__(DeepseekV4HipRadixBackend)
|
||||||
|
backend.MAX_SEQ_LEN_FOR_CAPTURE = 4096
|
||||||
|
backend._build_forward_metadata = mock.Mock(return_value=capture_metadata)
|
||||||
|
backend.init_forward_metadata_in_graph = mock.Mock()
|
||||||
|
backend._refresh_fp4_prefill_workspace = mock.Mock()
|
||||||
|
forward_batch = SimpleNamespace(name="capture")
|
||||||
|
|
||||||
|
result = backend.init_forward_metadata_for_breakable_cuda_graph_capture(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
|
||||||
|
backend._build_forward_metadata.assert_called_once_with(
|
||||||
|
forward_batch,
|
||||||
|
max_seq_len_override=backend.MAX_SEQ_LEN_FOR_CAPTURE,
|
||||||
|
use_prefill_cuda_graph=True,
|
||||||
|
)
|
||||||
|
backend.init_forward_metadata_in_graph.assert_called_once_with(forward_batch)
|
||||||
|
backend._refresh_fp4_prefill_workspace.assert_called_once_with(forward_batch)
|
||||||
|
self.assertIs(result, capture_metadata)
|
||||||
|
self.assertIs(backend.forward_metadata, capture_metadata)
|
||||||
|
|
||||||
|
def test_refresh_preserves_captured_hip_tensor_storage(self):
|
||||||
|
capture_workspace = object()
|
||||||
|
capture_metadata = DSV4Metadata(
|
||||||
|
self._make_core_metadata(0),
|
||||||
|
indexer_metadata=None,
|
||||||
|
fp4_prefill_workspace=capture_workspace,
|
||||||
|
fp4_k_write_metadata=(
|
||||||
|
torch.tensor([14], dtype=torch.int64),
|
||||||
|
torch.tensor([15], dtype=torch.int64),
|
||||||
|
),
|
||||||
|
fp4_q_positions=torch.tensor([16], dtype=torch.int64),
|
||||||
|
)
|
||||||
|
replay_metadata = DSV4Metadata(
|
||||||
|
self._make_core_metadata(100),
|
||||||
|
indexer_metadata=None,
|
||||||
|
fp4_k_write_metadata=(
|
||||||
|
torch.tensor([114], dtype=torch.int64),
|
||||||
|
torch.tensor([115], dtype=torch.int64),
|
||||||
|
),
|
||||||
|
fp4_q_positions=torch.tensor([116], dtype=torch.int64),
|
||||||
|
)
|
||||||
|
capture_core = capture_metadata.core_attn_metadata
|
||||||
|
replay_core = replay_metadata.core_attn_metadata
|
||||||
|
captured_core_tensors = {
|
||||||
|
field.name: getattr(capture_core, field.name)
|
||||||
|
for field in dataclasses.fields(capture_core)
|
||||||
|
if torch.is_tensor(getattr(capture_core, field.name))
|
||||||
|
}
|
||||||
|
captured_unified_tensors = {
|
||||||
|
field.name: getattr(capture_core.unified, field.name)
|
||||||
|
for field in dataclasses.fields(capture_core.unified)
|
||||||
|
if torch.is_tensor(getattr(capture_core.unified, field.name))
|
||||||
|
}
|
||||||
|
captured_fp4_tensors = {
|
||||||
|
"fp4_k_positions": capture_metadata.fp4_k_write_metadata[0],
|
||||||
|
"fp4_k_slots": capture_metadata.fp4_k_write_metadata[1],
|
||||||
|
"fp4_q_positions": capture_metadata.fp4_q_positions,
|
||||||
|
}
|
||||||
|
expected_core_tensors = {
|
||||||
|
name: getattr(replay_core, name).clone() for name in captured_core_tensors
|
||||||
|
}
|
||||||
|
expected_unified_tensors = {
|
||||||
|
name: getattr(replay_core.unified, name).clone()
|
||||||
|
for name in captured_unified_tensors
|
||||||
|
}
|
||||||
|
replay_fp4_tensors = {
|
||||||
|
"fp4_k_positions": replay_metadata.fp4_k_write_metadata[0],
|
||||||
|
"fp4_k_slots": replay_metadata.fp4_k_write_metadata[1],
|
||||||
|
"fp4_q_positions": replay_metadata.fp4_q_positions,
|
||||||
|
}
|
||||||
|
expected_fp4_tensors = {
|
||||||
|
name: tensor.clone() for name, tensor in replay_fp4_tensors.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
capture_metadata.refresh_for_breakable_cuda_graph_replay_(replay_metadata)
|
||||||
|
|
||||||
|
for field_name, captured_tensor in captured_core_tensors.items():
|
||||||
|
current = getattr(capture_core, field_name)
|
||||||
|
self.assertIs(current, captured_tensor, field_name)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(current, expected_core_tensors[field_name]), field_name
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
getattr(replay_core, field_name),
|
||||||
|
expected_core_tensors[field_name],
|
||||||
|
),
|
||||||
|
f"{field_name} replay source",
|
||||||
|
)
|
||||||
|
for field_name, captured_tensor in captured_unified_tensors.items():
|
||||||
|
current = getattr(capture_core.unified, field_name)
|
||||||
|
self.assertIs(current, captured_tensor, field_name)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(current, expected_unified_tensors[field_name]),
|
||||||
|
field_name,
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
getattr(replay_core.unified, field_name),
|
||||||
|
expected_unified_tensors[field_name],
|
||||||
|
),
|
||||||
|
f"{field_name} replay source",
|
||||||
|
)
|
||||||
|
|
||||||
|
current_fp4_tensors = {
|
||||||
|
"fp4_k_positions": capture_metadata.fp4_k_write_metadata[0],
|
||||||
|
"fp4_k_slots": capture_metadata.fp4_k_write_metadata[1],
|
||||||
|
"fp4_q_positions": capture_metadata.fp4_q_positions,
|
||||||
|
}
|
||||||
|
for name, captured_tensor in captured_fp4_tensors.items():
|
||||||
|
self.assertIs(current_fp4_tensors[name], captured_tensor)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(captured_tensor, expected_fp4_tensors[name]), name
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(replay_fp4_tensors[name], expected_fp4_tensors[name]),
|
||||||
|
f"{name} replay source",
|
||||||
|
)
|
||||||
|
self.assertIs(capture_metadata.fp4_prefill_workspace, capture_workspace)
|
||||||
|
|
||||||
|
def test_replay_refreshes_captured_metadata_and_workspace(self):
|
||||||
|
capture_metadata = DSV4Metadata(object(), indexer_metadata=None)
|
||||||
|
replay_metadata = DSV4Metadata(object(), indexer_metadata=None)
|
||||||
|
capture_metadata.refresh_for_breakable_cuda_graph_replay_ = mock.Mock()
|
||||||
|
|
||||||
|
backend = object.__new__(DeepseekV4HipRadixBackend)
|
||||||
|
backend.MAX_SEQ_LEN_FOR_CAPTURE = 4096
|
||||||
|
backend._build_forward_metadata = mock.Mock(return_value=replay_metadata)
|
||||||
|
backend.init_forward_metadata_in_graph = mock.Mock()
|
||||||
|
backend._refresh_fp4_prefill_workspace = mock.Mock()
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
backend._build_forward_metadata.assert_called_once_with(
|
||||||
|
static_forward_batch,
|
||||||
|
max_seq_len_override=backend.MAX_SEQ_LEN_FOR_CAPTURE,
|
||||||
|
use_prefill_cuda_graph=True,
|
||||||
|
)
|
||||||
|
backend.init_forward_metadata_in_graph.assert_called_once_with(
|
||||||
|
static_forward_batch
|
||||||
|
)
|
||||||
|
capture_metadata.refresh_for_breakable_cuda_graph_replay_.assert_called_once_with(
|
||||||
|
replay_metadata
|
||||||
|
)
|
||||||
|
backend._refresh_fp4_prefill_workspace.assert_called_once_with(
|
||||||
|
static_forward_batch
|
||||||
|
)
|
||||||
|
self.assertIs(backend.forward_metadata, capture_metadata)
|
||||||
|
|
||||||
|
def test_multistep_backend_forwards_bcg_metadata_hooks(self):
|
||||||
|
backend = object.__new__(DeepseekV4MultiStepBackend)
|
||||||
|
backend.speculative_num_steps = 3
|
||||||
|
backend.attn_backends = [mock.Mock(), mock.Mock(), mock.Mock()]
|
||||||
|
forward_batch = SimpleNamespace(name="live")
|
||||||
|
static_forward_batch = SimpleNamespace(name="static")
|
||||||
|
capture_metadata = [object(), object()]
|
||||||
|
|
||||||
|
for index, child in enumerate(backend.attn_backends[:-1]):
|
||||||
|
child.init_forward_metadata_for_breakable_cuda_graph_capture.return_value = f"capture-{index}"
|
||||||
|
|
||||||
|
captured = backend.init_forward_metadata_for_breakable_cuda_graph_capture(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
self.assertEqual(captured, ["capture-0", "capture-1"])
|
||||||
|
|
||||||
|
backend.prepare_forward_metadata_for_breakable_cuda_graph_replay(
|
||||||
|
capture_metadata,
|
||||||
|
forward_batch,
|
||||||
|
static_forward_batch=static_forward_batch,
|
||||||
|
)
|
||||||
|
for index, child in enumerate(backend.attn_backends[:-1]):
|
||||||
|
child.init_forward_metadata_for_breakable_cuda_graph_capture.assert_called_once_with(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
child.prepare_forward_metadata_for_breakable_cuda_graph_replay.assert_called_once_with(
|
||||||
|
capture_metadata[index],
|
||||||
|
forward_batch,
|
||||||
|
static_forward_batch=static_forward_batch,
|
||||||
|
)
|
||||||
|
backend.attn_backends[
|
||||||
|
-1
|
||||||
|
].init_forward_metadata_for_breakable_cuda_graph_capture.assert_not_called()
|
||||||
|
backend.attn_backends[
|
||||||
|
-1
|
||||||
|
].prepare_forward_metadata_for_breakable_cuda_graph_replay.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -131,5 +131,45 @@ class TestDecodeToExtendConversionVote(CustomTestCase):
|
|||||||
self.assertFalse(self._vote(beam=True))
|
self.assertFalse(self._vote(beam=True))
|
||||||
|
|
||||||
|
|
||||||
|
class TestPrefillCudaGraphVote(CustomTestCase):
|
||||||
|
def _vote(self, mode):
|
||||||
|
runner = Mock(spec=dp_attn.PrefillCudaGraphRunner)
|
||||||
|
runner.enable_lora = False
|
||||||
|
runner.max_context_size = None
|
||||||
|
runner.can_replay_locally.return_value = True
|
||||||
|
batch = SimpleNamespace(
|
||||||
|
forward_mode=mode,
|
||||||
|
extend_num_tokens=4,
|
||||||
|
input_embeds=None,
|
||||||
|
replace_embeds=None,
|
||||||
|
prefix_lens=[1, 1],
|
||||||
|
return_logprob=False,
|
||||||
|
batch_size=lambda: 2,
|
||||||
|
)
|
||||||
|
vote = dp_attn._local_prefill_cuda_graph_vote(
|
||||||
|
local_batch=batch,
|
||||||
|
prefill_graph_runner=runner,
|
||||||
|
coordinated_prefill=True,
|
||||||
|
breakable_prefill=True,
|
||||||
|
spec_algorithm=SpeculativeAlgorithm.NONE,
|
||||||
|
model_config=object(),
|
||||||
|
)
|
||||||
|
return vote, runner
|
||||||
|
|
||||||
|
def test_extend_batch_votes_for_prefill_graph(self):
|
||||||
|
vote, runner = self._vote(ForwardMode.EXTEND)
|
||||||
|
|
||||||
|
self.assertTrue(vote)
|
||||||
|
runner.can_replay_locally.assert_called_once()
|
||||||
|
self.assertFalse(runner.can_replay_locally.call_args.kwargs["is_mixed"])
|
||||||
|
|
||||||
|
def test_mixed_batch_delegates_to_runner_policy(self):
|
||||||
|
vote, runner = self._vote(ForwardMode.MIXED)
|
||||||
|
|
||||||
|
self.assertTrue(vote)
|
||||||
|
runner.can_replay_locally.assert_called_once()
|
||||||
|
self.assertTrue(runner.can_replay_locally.call_args.kwargs["is_mixed"])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -31,19 +31,22 @@ class TestPrefillCudaGraphPadding(CustomTestCase):
|
|||||||
runner._capture_chunked_prefix = False
|
runner._capture_chunked_prefix = False
|
||||||
runner.prefill_backend_name = Backend.TC_PIECEWISE
|
runner.prefill_backend_name = Backend.TC_PIECEWISE
|
||||||
runner.has_mha_companion_layers = False
|
runner.has_mha_companion_layers = False
|
||||||
|
runner.prefer_eager_mixed_prefill = False
|
||||||
runner.capture_hidden_mode = CaptureHiddenMode.NULL
|
runner.capture_hidden_mode = CaptureHiddenMode.NULL
|
||||||
runner.capture_num_tokens = [4, 16]
|
runner.capture_num_tokens = [4, 16]
|
||||||
runner.max_context_size = None
|
runner.max_context_size = None
|
||||||
runner.max_num_tokens = 16
|
runner.max_num_tokens = 16
|
||||||
return runner
|
return runner
|
||||||
|
|
||||||
def _make_forward_batch(self, num_tokens):
|
def _make_forward_batch(self, num_tokens, mode=ForwardMode.EXTEND):
|
||||||
return SimpleNamespace(
|
return SimpleNamespace(
|
||||||
batch_size=1,
|
batch_size=1,
|
||||||
input_embeds=None,
|
input_embeds=None,
|
||||||
replace_embeds=None,
|
replace_embeds=None,
|
||||||
mm_inputs=None,
|
mm_inputs=None,
|
||||||
forward_mode=ForwardMode.EXTEND,
|
forward_mode=mode,
|
||||||
|
global_forward_mode=None,
|
||||||
|
_original_forward_mode=None,
|
||||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
capture_hidden_mode=CaptureHiddenMode.NULL,
|
||||||
global_num_tokens_cpu=None,
|
global_num_tokens_cpu=None,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
@@ -63,6 +66,14 @@ class TestPrefillCudaGraphPadding(CustomTestCase):
|
|||||||
|
|
||||||
self.assertTrue(runner.can_run_graph(self._make_forward_batch(8)))
|
self.assertTrue(runner.can_run_graph(self._make_forward_batch(8)))
|
||||||
|
|
||||||
|
def test_mixed_batch_uses_scoped_runner_policy(self):
|
||||||
|
runner = self._make_runner()
|
||||||
|
batch = self._make_forward_batch(8, mode=ForwardMode.MIXED)
|
||||||
|
|
||||||
|
self.assertTrue(runner.can_run_graph(batch))
|
||||||
|
runner.prefer_eager_mixed_prefill = True
|
||||||
|
self.assertFalse(runner.can_run_graph(batch))
|
||||||
|
|
||||||
def test_replay_snapshot_uses_padded_token_count(self):
|
def test_replay_snapshot_uses_padded_token_count(self):
|
||||||
runner = self._make_runner()
|
runner = self._make_runner()
|
||||||
runner.use_captured_attn_metadata = False
|
runner.use_captured_attn_metadata = False
|
||||||
|
|||||||
Reference in New Issue
Block a user