[Refactor] Make DeepSeek-V4 attention backend tolerate an absent CPU seq_lens mirror (#30695)
This commit is contained in:
@@ -457,6 +457,8 @@ class DeepseekV4AttnBackend(
|
||||
):
|
||||
use_captured_forward_metadata_for_breakable_cuda_graph: bool = True
|
||||
|
||||
needs_cpu_seq_lens: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_runner: ModelRunner,
|
||||
@@ -509,6 +511,11 @@ class DeepseekV4AttnBackend(
|
||||
DSV4RawDecodeMetadata,
|
||||
] = None
|
||||
self.online_c128_mtp = OnlineC128MTPController(self)
|
||||
# Draft-extend and online-c128 verify metadata are host-planned, so
|
||||
# spec runs keep the relay publish (the mirror only exists under
|
||||
# spec-v2; without spec the flag has no consumer either way).
|
||||
if model_runner.server_args.speculative_algorithm is not None:
|
||||
self.needs_cpu_seq_lens = True
|
||||
self.sparse_prefill_workspace = SparsePrefillWorkspace(self.device)
|
||||
|
||||
def _move_to_device(self, x: List[int]) -> torch.Tensor:
|
||||
@@ -706,9 +713,7 @@ class DeepseekV4AttnBackend(
|
||||
if envs.SGLANG_PREP_IN_CUDA_GRAPH.get():
|
||||
assert out_cache_loc is not None
|
||||
seq_lens_cpu_list = (
|
||||
seq_lens.detach().cpu().tolist()
|
||||
if seq_lens_cpu is None
|
||||
else seq_lens_cpu.tolist()
|
||||
seq_lens_cpu.tolist() if seq_lens_cpu is not None else None
|
||||
)
|
||||
if not hasattr(self, "extend_seq_lens_buffer"):
|
||||
self.extend_seq_lens_buffer = torch.tensor(
|
||||
@@ -984,14 +989,13 @@ class DeepseekV4AttnBackend(
|
||||
)
|
||||
out_cache_loc = torch.zeros(bs, dtype=torch.int64, device=device)
|
||||
|
||||
assert seq_lens_cpu is not None
|
||||
seq_lens = seq_lens[:bs]
|
||||
seq_lens_cpu = seq_lens_cpu[:bs]
|
||||
req_pool_indices = req_pool_indices[:bs]
|
||||
|
||||
actual_max_seq_len = seq_lens_cpu.max().item()
|
||||
chosen_max_seq_len = self.MAX_SEQ_LEN_FOR_CAPTURE
|
||||
assert actual_max_seq_len <= chosen_max_seq_len
|
||||
if seq_lens_cpu is not None:
|
||||
seq_lens_cpu = seq_lens_cpu[:bs]
|
||||
actual_max_seq_len = seq_lens_cpu.max().item()
|
||||
assert actual_max_seq_len <= chosen_max_seq_len
|
||||
|
||||
if bucket == _GraphBucket.DECODE_OR_IDLE:
|
||||
assert out_cache_loc is not None
|
||||
@@ -1060,11 +1064,14 @@ class DeepseekV4AttnBackend(
|
||||
mode="constant",
|
||||
value=0,
|
||||
)
|
||||
draft_extend_seq_lens_cpu = (
|
||||
seq_lens_cpu.tolist() if seq_lens_cpu is not None else seq_lens.tolist()
|
||||
)
|
||||
temp_metadata = self.init_forward_metadata_draft_extend(
|
||||
max_seq_len=chosen_max_seq_len,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu.tolist(),
|
||||
seq_lens_cpu=draft_extend_seq_lens_cpu,
|
||||
num_tokens_per_bs=num_tokens_per_bs,
|
||||
out_cache_loc=out_cache_loc,
|
||||
use_prefill_cuda_graph=True,
|
||||
@@ -1112,14 +1119,14 @@ class DeepseekV4AttnBackend(
|
||||
assert self.req_to_token_pool.req_to_token is self.req_to_token
|
||||
|
||||
assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0
|
||||
assert seq_lens_cpu is not None
|
||||
if max_seq_len_override is None:
|
||||
max_seq_len_override = getattr(forward_batch, "max_seq_len_override", None)
|
||||
max_seq_len = (
|
||||
int(seq_lens_cpu.max().item())
|
||||
if max_seq_len_override is None
|
||||
else max_seq_len_override
|
||||
)
|
||||
if max_seq_len_override is not None:
|
||||
max_seq_len = max_seq_len_override
|
||||
elif seq_lens_cpu is not None:
|
||||
max_seq_len = int(seq_lens_cpu.max().item())
|
||||
else:
|
||||
max_seq_len = int(seq_lens.max().item())
|
||||
verify_bs = _get_target_verify_bs(forward_batch)
|
||||
online_c128_state_slot_offset = self.online_c128_mtp.prepare_forward(
|
||||
logical_forward_mode,
|
||||
@@ -1159,16 +1166,18 @@ class DeepseekV4AttnBackend(
|
||||
extend_seq_lens = forward_batch.extend_seq_lens
|
||||
assert (
|
||||
seq_lens is not None
|
||||
and seq_lens_cpu is not None
|
||||
and extend_seq_lens is not None
|
||||
and extend_seq_lens_cpu is not None
|
||||
)
|
||||
is_draft = forward_batch.forward_mode.is_draft_extend_v2()
|
||||
prefill_seq_lens_cpu = (
|
||||
seq_lens_cpu.tolist() if seq_lens_cpu is not None else seq_lens.tolist()
|
||||
)
|
||||
metadata = self.init_forward_metadata_prefill(
|
||||
max_seq_len=max_seq_len,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu.tolist(),
|
||||
seq_lens_cpu=prefill_seq_lens_cpu,
|
||||
out_cache_loc=forward_batch.out_cache_loc,
|
||||
num_tokens=sum(extend_seq_lens_cpu),
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
|
||||
Reference in New Issue
Block a user