[Spec] Remove deprecated EAGLE v1 DRAFT_EXTEND forward mode (#28129)

This commit is contained in:
Cheng Wan
2026-06-13 13:00:21 -07:00
committed by GitHub
parent d988d5d681
commit bde6bccf39
56 changed files with 148 additions and 2484 deletions
@@ -423,13 +423,9 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
_fa_cp_attn,
)
elif (
(
forward_batch.extend_prefix_lens_cpu is not None
and any(forward_batch.extend_prefix_lens_cpu)
)
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
):
forward_batch.extend_prefix_lens_cpu is not None
and any(forward_batch.extend_prefix_lens_cpu)
) or forward_batch.forward_mode.is_target_verify():
result = flash_attn_with_kvcache(
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache=key_cache,
@@ -513,7 +509,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
if (
forward_batch.attn_attend_prefix_cache is not None
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
):
if forward_batch.attn_attend_prefix_cache:
assert not get_global_server_args().disable_chunked_prefix_cache
@@ -463,7 +463,6 @@ class AscendAttnBackend(AttentionBackend):
self.forward_metadata.seq_lens_cpu_int = forward_batch.seq_lens_cpu.int()
if (
not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_target_verify()
):
seq_lens_list_cumsum = np.cumsum(forward_batch.extend_seq_lens_cpu)
@@ -480,7 +479,7 @@ class AscendAttnBackend(AttentionBackend):
if (
self.use_mla
and forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_target_verify()
and sum(forward_batch.extend_prefix_lens_cpu) > 0
):
@@ -570,11 +569,7 @@ class AscendAttnBackend(AttentionBackend):
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
metadata.seq_lens = seq_lens
if (
forward_mode.is_target_verify()
or forward_mode.is_draft_extend_v2()
or forward_mode.is_draft_extend()
):
if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2():
metadata.actual_seq_lengths_q = torch.arange(
self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens
@@ -980,7 +975,6 @@ class AscendAttnBackend(AttentionBackend):
is_prefill = (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_target_verify()
)
@@ -1014,12 +1008,6 @@ class AscendAttnBackend(AttentionBackend):
.to(q.device)
.to(torch.int32)
)
elif forward_batch.forward_mode.is_draft_extend():
actual_seq_qlen = (
forward_batch.extend_seq_lens.cumsum()
.to(q.device)
.to(torch.int32)
)
else:
actual_seq_qlen = (
torch.arange(1, q.shape[0] + 1).to(q.device).to(torch.int32)
@@ -1118,7 +1106,6 @@ class AscendAttnBackend(AttentionBackend):
)
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend_v2()
):
return self.forward_mtp(
@@ -1888,10 +1875,7 @@ class AscendAttnBackend(AttentionBackend):
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
)
if (
forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend_v2()
):
if forward_batch.forward_mode.is_draft_extend_v2():
actual_seq_lengths = (
np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist()
)
@@ -2004,16 +1988,11 @@ class AscendAttnBackend(AttentionBackend):
actual_seq_lengths_kv = (
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
)
if forward_batch.forward_mode.is_draft_extend():
actual_seq_lengths = (
np.array(forward_batch.extend_seq_lens_cpu).cumsum().tolist()
)
else:
actual_seq_lengths = np.arange(
self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens + q_nope.shape[0],
self.speculative_num_draft_tokens,
)
actual_seq_lengths = np.arange(
self.speculative_num_draft_tokens,
self.speculative_num_draft_tokens + q_nope.shape[0],
self.speculative_num_draft_tokens,
)
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
q_nope,
@@ -77,7 +77,7 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
forward_batch: ForwardBatch,
in_capture: bool = False,
):
if forward_batch.forward_mode.is_draft_extend(True):
if forward_batch.forward_mode.is_draft_extend_v2():
return
super().init_forward_metadata_out_graph(forward_batch, in_capture=in_capture)
self.prepare_gdn_inputs(
@@ -88,7 +88,7 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
self.graph_mode = True
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_draft_extend(True):
if forward_batch.forward_mode.is_draft_extend_v2():
return
super().init_forward_metadata(forward_batch)
self.prepare_gdn_inputs(
@@ -459,7 +459,7 @@ def forward_dsa_core_npu(
if (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_target_verify()
):
attn_output = attn_output.transpose(0, 1)
@@ -65,7 +65,7 @@ def _build_generalized_req_id_per_token(
if forward_mode.is_target_verify():
per_req = int(forward_batch.spec_info.draft_token_num)
result = _expand_uniform(generalized_req_ids_per_row, per_req)
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
per_req = int(forward_batch.spec_info.num_tokens_per_req)
result = _expand_uniform(generalized_req_ids_per_row, per_req)
elif forward_mode.is_extend():
@@ -1117,88 +1117,6 @@ class AiterAttnBackend(AttentionBackend):
self.indices_updater_prefill.max_q_len,
self.indices_updater_prefill.max_kv_len,
)
elif forward_batch.forward_mode.is_draft_extend():
# EAGLE V1: DRAFT_EXTEND mode - uses spec_info.num_accept_tokens
if self.use_mla:
kv_indices, kv_indptr, qo_indptr, custom_mask = (
spec_info.generate_attn_arg_prefill(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
self.req_to_token,
)
)
if _use_mla_ps_kernel:
max_seqlen_qo = max(forward_batch.extend_seq_lens_cpu)
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = self.make_mla_decode_meta_data_buffer(max_seqlen_qo, bs)
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
self.kv_last_page_len[:bs],
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_seqlen_qo,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
# self.mla_indices_updater_prefill.kv_last_page_len,
self.kv_last_page_len[:bs],
max(forward_batch.extend_seq_lens_cpu),
forward_batch.seq_lens_cpu.max().item(),
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
run_graph=False,
)
else:
# Non-MLA draft_extend: use triton extend kernel with causal masking
kv_indices, kv_indptr, qo_indptr, custom_mask = (
spec_info.generate_attn_arg_prefill(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
self.req_to_token,
)
)
kv_indices = kv_indices.to(torch.int64)
draft_max_extend_len = torch.max(spec_info.num_accept_tokens).item()
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
None,
draft_max_extend_len,
None,
custom_mask=custom_mask,
mask_indptr=None,
max_extend_len=draft_max_extend_len,
)
elif forward_batch.forward_mode.is_target_verify():
if self.use_mla:
draft_num = spec_info.draft_token_num
@@ -1941,72 +1859,6 @@ class AiterAttnBackend(AttentionBackend):
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
)
elif forward_mode.is_draft_extend():
# EAGLE V1: Uses spec_info.num_accept_tokens
num_tokens_per_bs = self.speculative_num_steps + 1
seq_lens = seq_lens[:bs]
extend_lens = spec_info.num_accept_tokens[:bs]
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[1 : bs + 1] = torch.cumsum(extend_lens, dim=0)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_bs
if self.use_mla and _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
kv_last_page_len,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
max_kv_len,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
)
else:
raise ValueError("Invalid forward mode")
@@ -2123,7 +1975,6 @@ class AiterAttnBackend(AttentionBackend):
if (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
):
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
@@ -2272,10 +2123,7 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits=num_kv_splits,
)
return o
elif (
forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend_v2()
):
elif forward_batch.forward_mode.is_draft_extend_v2():
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
work_info_set = self.forward_metadata.work_info_set
@@ -2347,10 +2195,7 @@ class AiterAttnBackend(AttentionBackend):
f"Invalid forward mode for MLA prefill: {forward_batch.forward_mode=}"
)
else:
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
):
if forward_batch.forward_mode.is_target_verify():
if layer.qk_head_dim != layer.v_head_dim:
o = q.new_empty(
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim)
@@ -412,7 +412,7 @@ class _GraphBucket(enum.Enum):
return cls.DECODE_OR_IDLE
if forward_mode.is_target_verify():
return cls.TARGET_VERIFY
if forward_mode.is_draft_extend(include_v2=True):
if forward_mode.is_draft_extend_v2():
return cls.DRAFT_EXTEND
raise NotImplementedError(f"unsupported {forward_mode=}")
@@ -1016,7 +1016,7 @@ class DeepseekV4AttnBackend(
and extend_seq_lens is not None
and extend_seq_lens_cpu is not None
)
is_draft = forward_batch.forward_mode.is_draft_extend(include_v2=True)
is_draft = forward_batch.forward_mode.is_draft_extend_v2()
metadata = self.init_forward_metadata_prefill(
max_seq_len=max_seq_len,
req_pool_indices=req_pool_indices,
@@ -372,7 +372,7 @@ class _GraphBucket(enum.Enum):
return cls.DECODE_OR_IDLE
if forward_mode.is_target_verify():
return cls.TARGET_VERIFY
if forward_mode.is_draft_extend(include_v2=True):
if forward_mode.is_draft_extend_v2():
return cls.DRAFT_EXTEND
raise NotImplementedError(f"unsupported {forward_mode=}")
@@ -927,7 +927,7 @@ class DeepseekV4HipRadixBackend(
and extend_seq_lens is not None
and extend_seq_lens_cpu is not None
)
is_draft = forward_batch.forward_mode.is_draft_extend(include_v2=True)
is_draft = forward_batch.forward_mode.is_draft_extend_v2()
metadata = self.init_forward_metadata_prefill(
max_seq_len=max_seq_len,
req_pool_indices=req_pool_indices,
@@ -105,10 +105,6 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
return self._precompute_target_verify_mode(
bs, req_pool_indices, seq_lens, seq_lens_cpu
)
elif forward_mode.is_draft_extend():
return self._precompute_draft_extend_mode(
bs, req_pool_indices, seq_lens, seq_lens_cpu, spec_info
)
else:
raise ValueError(f"Unsupported forward mode: {forward_mode}")
@@ -594,7 +594,7 @@ class Indexer(MultiPlatformOp):
blocksize = page_size
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
seqlens_32 = metadata.get_seqlens_expanded()
else:
@@ -1505,7 +1505,7 @@ class Indexer(MultiPlatformOp):
if (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
topk_result = self._get_topk_paged(
forward_batch, layer_id, q_fp8, weights, metadata
@@ -1612,7 +1612,6 @@ class Indexer(MultiPlatformOp):
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
)
bs = q_lora.shape[0]
@@ -1799,7 +1798,6 @@ class Indexer(MultiPlatformOp):
if (
forward_batch.forward_mode.is_draft_extend_v2()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
):
num_draft_tokens = get_attn_backend().speculative_num_draft_tokens
actual_seq_lengths_q = torch.arange(
@@ -83,14 +83,6 @@ def verify_single_backend_fused_metadata_copy(
ref_page_table_1[:, : precomputed.max_seqlen_k].copy_(precomputed.page_indices)
ref_dsa_seqlens_expanded.copy_(precomputed.seqlens_expanded)
ref_dsa_cache_seqlens.copy_(precomputed.dsa_cache_seqlens)
elif forward_mode.is_draft_extend():
# Draft extend mode
rows = precomputed.page_indices.shape[0]
cols = precomputed.max_seqlen_k
ref_page_table_1[:rows, :cols].copy_(precomputed.page_indices)
size = precomputed.seqlens_expanded_size
ref_dsa_seqlens_expanded[:size].copy_(precomputed.seqlens_expanded)
ref_dsa_cache_seqlens[:size].copy_(precomputed.dsa_cache_seqlens)
# Copy DSA cu_seqlens
size = precomputed.seqlens_expanded_size
@@ -141,14 +133,6 @@ def verify_single_backend_fused_metadata_copy(
fused_page_table_1[:, : precomputed.max_seqlen_k],
ref_page_table_1[:, : precomputed.max_seqlen_k],
)
elif forward_mode.is_draft_extend():
rows = precomputed.page_indices.shape[0]
cols = precomputed.max_seqlen_k
check_tensor_equal(
"page_table_1",
fused_page_table_1[:rows, :cols],
ref_page_table_1[:rows, :cols],
)
# Compare dsa_cache_seqlens only for the region that was updated
if forward_mode.is_decode_or_idle():
@@ -157,7 +141,7 @@ def verify_single_backend_fused_metadata_copy(
fused_dsa_cache_seqlens,
ref_dsa_cache_seqlens,
)
else: # TARGET_VERIFY or DRAFT_EXTEND
else: # TARGET_VERIFY
size = precomputed.seqlens_expanded_size
check_tensor_equal(
"dsa_cache_seqlens",
@@ -165,8 +149,8 @@ def verify_single_backend_fused_metadata_copy(
ref_dsa_cache_seqlens[:size],
)
# Compare dsa_seqlens_expanded only for TARGET_VERIFY and DRAFT_EXTEND
if forward_mode.is_target_verify() or forward_mode.is_draft_extend():
# Compare dsa_seqlens_expanded only for TARGET_VERIFY
if forward_mode.is_target_verify():
size = precomputed.seqlens_expanded_size
check_tensor_equal(
"dsa_seqlens_expanded",
@@ -406,9 +406,7 @@ class DeepseekSparseAttnBackend(
and is_sm100_supported()
):
return cache_seqlens_int32.view(-1, 1).expand(-1, next_n).contiguous()
if forward_mode.is_target_verify() or forward_mode.is_draft_extend(
include_v2=True
):
if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2():
return _to_2d_context_lens(seqlens_expanded, batch_size)
return _to_2d_context_lens(cache_seqlens_int32, batch_size)
@@ -488,7 +486,7 @@ class DeepseekSparseAttnBackend(
if (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
)
else self.dsa_prefill_impl
)
@@ -531,7 +529,7 @@ class DeepseekSparseAttnBackend(
page_table = torch.repeat_interleave(
page_table, repeats=self.speculative_num_draft_tokens, dim=0
)
elif forward_batch.forward_mode.is_draft_extend(include_v2=True):
elif forward_batch.forward_mode.is_draft_extend_v2():
assert (
forward_batch.extend_seq_lens_cpu is not None
and forward_batch.extend_seq_lens is not None
@@ -614,11 +612,7 @@ class DeepseekSparseAttnBackend(
)
page_table = page_table[bs_idx, :max_seqlen_k]
if (
any(forward_batch.extend_prefix_lens_cpu)
or forward_batch.forward_mode == ForwardMode.DRAFT_EXTEND
or bs_idx_cpu is not None
):
if any(forward_batch.extend_prefix_lens_cpu) or bs_idx_cpu is not None:
max_seqlen_q = (
max(extend_seq_lens_cpu) if len(extend_seq_lens_cpu) != 0 else 1
)
@@ -690,7 +684,7 @@ class DeepseekSparseAttnBackend(
if is_cuda() and (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
paged_mqa_ctx_lens_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
forward_batch.forward_mode,
@@ -905,9 +899,7 @@ class DeepseekSparseAttnBackend(
)
else:
flashmla_metadata = None
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(
include_v2=True
):
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2():
cache_seqlens_int32 = (seq_lens + self.speculative_num_draft_tokens).to(
torch.int32
)
@@ -975,7 +967,7 @@ class DeepseekSparseAttnBackend(
if is_cuda() and (
forward_mode.is_decode_or_idle()
or forward_mode.is_target_verify()
or forward_mode.is_draft_extend(include_v2=True)
or forward_mode.is_draft_extend_v2()
):
paged_mqa_ctx_lens_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
forward_mode, cache_seqlens_int32, seqlens_expanded, bs
@@ -1094,7 +1086,7 @@ class DeepseekSparseAttnBackend(
seqlens_expanded, self.dsa_index_topk
)
metadata.dsa_cache_seqlens_int32.copy_(dsa_cache_seqlens)
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
max_seqlen_k = int(seq_lens_cpu.max().item())
cache_seqlens = seq_lens.to(torch.int32)
metadata.cache_seqlens_int32.copy_(cache_seqlens)
@@ -1133,9 +1125,9 @@ class DeepseekSparseAttnBackend(
if is_cuda() and (
forward_mode.is_decode_or_idle()
or forward_mode.is_target_verify()
or forward_mode.is_draft_extend(include_v2=True)
or forward_mode.is_draft_extend_v2()
):
if forward_mode.is_draft_extend(include_v2=True):
if forward_mode.is_draft_extend_v2():
schedule_seqlens_expanded = metadata.dsa_seqlens_expanded
else:
schedule_seqlens_expanded = seqlens_expanded
@@ -1227,8 +1219,6 @@ class DeepseekSparseAttnBackend(
mode_int = 0 # DECODE
elif forward_mode.is_target_verify():
mode_int = 1 # TARGET_VERIFY
elif forward_mode.is_draft_extend():
mode_int = 2 # DRAFT_EXTEND
else:
raise ValueError(f"Unsupported forward_mode: {forward_mode}")
@@ -1314,18 +1304,6 @@ class DeepseekSparseAttnBackend(
metadata.dsa_seqlens_expanded.copy_(precomputed.seqlens_expanded)
metadata.dsa_cache_seqlens_int32.copy_(precomputed.dsa_cache_seqlens)
elif forward_mode.is_draft_extend():
# Draft extend mode
rows = precomputed.page_indices.shape[0]
cols = precomputed.max_seqlen_k
metadata.page_table_1[:rows, :cols].copy_(precomputed.page_indices)
size = precomputed.seqlens_expanded_size
metadata.dsa_seqlens_expanded[:size].copy_(precomputed.seqlens_expanded)
metadata.dsa_cache_seqlens_int32[:size].copy_(
precomputed.dsa_cache_seqlens
)
# Copy DSA cu_seqlens
size = precomputed.seqlens_expanded_size
metadata.dsa_cu_seqlens_k[1 : 1 + size].copy_(
@@ -1400,7 +1378,7 @@ class DeepseekSparseAttnBackend(
self.dsa_decode_impl
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
)
else self.dsa_prefill_impl
)
@@ -352,7 +352,7 @@ class FlashAttentionBackend(AttentionBackend):
self._sched_meta_buf[n:] = 0
metadata.scheduler_metadata = self._sched_meta_buf[:n]
if forward_mode.is_draft_extend(include_v2=True):
if forward_mode.is_draft_extend_v2():
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel
@@ -645,9 +645,10 @@ class FlashAttentionBackend(AttentionBackend):
forward_batch.req_pool_indices, : metadata.max_seq_len_k
]
if any(
forward_batch.extend_prefix_lens_cpu
) or forward_batch.forward_mode.is_draft_extend(include_v2=True):
if (
any(forward_batch.extend_prefix_lens_cpu)
or forward_batch.forward_mode.is_draft_extend_v2()
):
extend_seq_lens = forward_batch.extend_seq_lens
metadata.max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
metadata.cu_seqlens_q = torch.nn.functional.pad(
@@ -1056,7 +1057,7 @@ class FlashAttentionBackend(AttentionBackend):
if (
forward_batch.attn_attend_prefix_cache is not None
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
):
# Do multi-head attention with chunked prefix cache
if forward_batch.attn_attend_prefix_cache:
@@ -2024,7 +2025,7 @@ class FlashAttentionBackend(AttentionBackend):
self.target_verify_metadata_topk_swa[bs] = metadata_swa
metadata.swa_spec_metadata = metadata_swa
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = num_tokens // bs
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs
@@ -2355,40 +2356,6 @@ class FlashAttentionBackend(AttentionBackend):
metadata, metadata_expand, metadata_swa
)
elif forward_mode.is_draft_extend():
metadata = self.draft_extend_metadata[bs]
metadata.cache_seqlens_int32.copy_(seq_lens)
metadata.max_seq_len_k = seq_lens_cpu.max().item()
metadata.cu_seqlens_k[1:].copy_(
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
)
extend_lens = spec_info.num_accept_tokens[:bs]
if spec_info.num_accept_tokens_cpu:
metadata.max_seq_len_q = max(spec_info.num_accept_tokens_cpu)
else:
metadata.max_seq_len_q = 1
metadata.cu_seqlens_q[1:].copy_(
torch.cumsum(extend_lens, dim=0, dtype=torch.int32)
)
max_seq_pages = (
metadata.max_seq_len_k + self.page_size - 1
) // self.page_size
page_indices = self.req_to_token[
req_pool_indices[:, None],
self.draft_extend_metadata["strided_indices"][:max_seq_pages],
]
if self.use_sliding_window_kv_pool and metadata.swa_page_table is not None:
swa_page_indices = self.token_to_kv_pool.translate_loc_from_full_to_swa(
page_indices
)
metadata.swa_page_table[:, :max_seq_pages].copy_(
swa_page_indices // self.page_size
)
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
elif forward_mode.is_draft_extend_v2():
metadata = self.draft_extend_metadata[bs]
metadata.cache_seqlens_int32.copy_(seq_lens)
@@ -556,7 +556,7 @@ class FlashInferAttnBackend(AttentionBackend):
fixed_split_size=None,
disable_split_kv=self.disable_cuda_graph_kv_split,
)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
elif forward_mode.is_target_verify():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
@@ -628,24 +628,6 @@ class FlashInferAttnBackend(AttentionBackend):
self.forward_metadata = DecodeMetadata(
self.decode_wrappers, swa_out_cache_loc=swa_out_cache_loc
)
elif forward_batch.forward_mode.is_draft_extend():
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_cpu,
forward_batch.seq_lens_sum,
prefix_lens=None,
prefill_wrappers=self.prefill_wrappers_paged,
use_ragged=False,
encoder_lens=forward_batch.encoder_lens,
spec_info=forward_batch.spec_info,
)
self.forward_metadata = PrefillMetadata(
self.prefill_wrappers_paged,
False,
False,
swa_out_cache_loc=swa_out_cache_loc,
)
elif forward_batch.forward_mode.is_target_verify():
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
@@ -812,11 +794,7 @@ class FlashInferAttnBackend(AttentionBackend):
decode_wrappers = self._create_decode_wrappers(bs, num_tokens)
self.decode_cuda_graph_metadata[bs] = decode_wrappers
self.forward_metadata = DecodeMetadata(decode_wrappers)
elif (
forward_mode.is_target_verify()
or forward_mode.is_draft_extend()
or forward_mode.is_dllm_extend()
):
elif forward_mode.is_target_verify() or forward_mode.is_dllm_extend():
use_custom_mask = (
forward_mode.is_target_verify()
and spec_info is not None
@@ -336,7 +336,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
# fast_mla_decode_plan needs _cached_module from the initial
# begin_forward above, so install it only after that call completes.
decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
elif forward_mode.is_target_verify():
prefill_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
use_cuda_graph=True,
@@ -381,17 +381,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
init_metadata_replay=False,
)
self.forward_metadata = DecodeMetadata(self.decode_wrapper)
elif forward_batch.forward_mode.is_draft_extend():
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
prefix_lens=None,
prefill_wrapper_paged=self.prefill_wrapper_paged,
use_ragged=False,
spec_info=forward_batch.spec_info,
)
self.forward_metadata = PrefillMetadata(self.prefill_wrapper_paged, False)
elif forward_batch.forward_mode.is_target_verify():
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
@@ -493,7 +482,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
spec_info=spec_info,
**self.fast_decode_kwargs,
)
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
elif forward_mode.is_target_verify():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
@@ -385,7 +385,6 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
):
if forward_batch.forward_mode in (
ForwardMode.EXTEND,
ForwardMode.DRAFT_EXTEND,
ForwardMode.DRAFT_EXTEND_V2,
):
return super().forward_extend(q, k, v, layer, forward_batch, save_kv_cache)
@@ -37,12 +37,12 @@ class HybridAttnBackend(AttentionBackend):
Note:
- decode_or_idle: Always uses decode backend
- target_verify or draft_extend: Uses decode backend if speculative_attention_mode is "decode", otherwise prefill backend
- target_verify: Uses decode backend if speculative_attention_mode is "decode", otherwise prefill backend
- prefill: Always uses prefill backend
"""
if forward_mode.is_decode_or_idle():
return self.decode_backend
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
elif forward_mode.is_target_verify():
return (
self.decode_backend
if self.model_runner.server_args.speculative_attention_mode == "decode"
@@ -49,7 +49,7 @@ class MambaAttnBackendBase(AttentionBackend):
if (
not forward_batch.forward_mode.is_extend()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
or self.is_draft_worker
):
return
@@ -454,31 +454,25 @@ class TritonAttnBackend(AttentionBackend):
dtype=torch.int32,
device=self.device,
)
if forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2).
# Triton extend kernel receives extend K/V as separate tensors, so
# kv_indptr/kv_indices must cover only the prefix portion.
# extend_seq_lens_tensor is only attached to spec_info at real
# replay (eagle_draft_extend_cuda_graph_runner.replay); during the
# capture-time warmup it's absent, so fall back to zeros (matches
# the pre-unification capture path in #26651). Clamp at 0 because
# padded rows (raw_bs..bs) leave seq_lens at the fill value (1)
# while extend_seq_lens stays at num_tokens_per_bs, which would
# otherwise produce negative kv_lens; padded rows reference
# reserved req-pool slot 0 and their output is discarded.
if (
spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
):
extend_seq_lens = spec_info.extend_seq_lens_tensor[:bs].to(torch.int32)
else:
extend_seq_lens = torch.zeros(
bs, dtype=torch.int32, device=seq_lens.device
)
kv_lens = torch.clamp(seq_lens - extend_seq_lens, min=0).to(torch.int32)
# DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2).
# Triton extend kernel receives extend K/V as separate tensors, so
# kv_indptr/kv_indices must cover only the prefix portion.
# extend_seq_lens_tensor is only attached to spec_info at real
# replay (eagle_draft_extend_cuda_graph_runner.replay); during the
# capture-time warmup it's absent, so fall back to zeros (matches
# the pre-unification capture path in #26651). Clamp at 0 because
# padded rows (raw_bs..bs) leave seq_lens at the fill value (1)
# while extend_seq_lens stays at num_tokens_per_bs, which would
# otherwise produce negative kv_lens; padded rows reference
# reserved req-pool slot 0 and their output is discarded.
if (
spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
):
extend_seq_lens = spec_info.extend_seq_lens_tensor[:bs].to(torch.int32)
else:
# DRAFT_EXTEND_V1: seq_lens = prefix only.
kv_lens = seq_lens
extend_seq_lens = torch.zeros(bs, dtype=torch.int32, device=seq_lens.device)
kv_lens = torch.clamp(seq_lens - extend_seq_lens, min=0).to(torch.int32)
kv_indptr = self._fill_kv_indptr_and_indices(
bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices
)
@@ -693,32 +687,6 @@ class TritonAttnBackend(AttentionBackend):
attn_logits = None
attn_lse = None
elif forward_batch.forward_mode.is_draft_extend():
# Eager only (CG replay bypasses init); explicit D2H here instead of
# letting torch.empty inside generate_attn_arg_prefill .item() on a
# GPU cumsum tensor.
seq_lens_sum = (
forward_batch.seq_lens_sum
if forward_batch.seq_lens_sum is not None
else int(forward_batch.seq_lens.sum())
)
kv_indices, kv_indptr, qo_indptr, custom_mask = (
spec_info.generate_attn_arg_prefill(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum,
self.req_to_token,
)
)
kv_indices = kv_indices.to(torch.int64)
mask_indptr = None
# TODO(FIXME): This will trigger an invalid Eagle tree when using
# `max(spec_info.num_accept_tokens_cpu)`.
# It might have been forgotten to update somewhere.
max_extend_len = torch.max(spec_info.num_accept_tokens).item()
num_kv_splits = None
attn_logits = None
attn_lse = None
else:
# gpu_only leaves _cpu unset; ub-allocate is safe (ragged write
# from GPU tensor, extra tail unused).
@@ -944,7 +912,7 @@ class TritonAttnBackend(AttentionBackend):
window_kv_offsets=self.cuda_graph_window_kv_offsets if swa else None,
swa_out_cache_loc=swa_out_cache_loc,
)
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
return ForwardMetadata(
attn_logits=None,
attn_lse=None,
@@ -1000,7 +968,7 @@ class TritonAttnBackend(AttentionBackend):
self._update_target_verify_buffers(
bs, seq_lens, req_pool_indices, spec_info
)
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
self._update_draft_extend_buffers(
bs, seq_lens, req_pool_indices, forward_mode, spec_info
)
@@ -409,7 +409,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
bs,
)
self.target_verify_metadata[bs] = metadata
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = num_tokens // bs
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs
@@ -506,7 +506,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
]
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
self._copy_swa_page_table(metadata, page_indices, max_seq_pages)
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
metadata = self.draft_extend_metadata[bs]
metadata.cache_seqlens_int32.copy_(seq_lens)
@@ -620,12 +620,6 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_draft_extend():
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it
# to max(num_accept_tokens_cpu) which is None/empty at capture
# time, falling back to 1. Restore the correct upper bound so
# the kernel sees num_tokens_per_bs (not 1) for all replays.
self.forward_metadata.max_seq_len_q = num_tokens // bs
else:
self._apply_cuda_graph_metadata(
bs=bs,
@@ -725,9 +719,10 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
forward_batch.req_pool_indices, : metadata.max_seq_len_k
]
if any(
forward_batch.extend_prefix_lens_cpu
) or forward_batch.forward_mode.is_draft_extend(include_v2=True):
if (
any(forward_batch.extend_prefix_lens_cpu)
or forward_batch.forward_mode.is_draft_extend_v2()
):
extend_seq_lens = forward_batch.extend_seq_lens
# NOTE: in piecewise CUDA graph warmup, extend_seq_lens_cpu is a torch.Tensor;
# Python's max() returns a 0-d tensor, but flashinfer expects an int.
@@ -344,7 +344,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if forward_mode.is_target_verify():
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = num_tokens // bs
metadata.max_seq_len_q = num_tokens_per_bs
metadata.sum_seq_lens_q = num_tokens_per_bs * bs
@@ -386,7 +386,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if forward_mode.is_target_verify():
seq_lens = seq_lens[:bs] + self.num_draft_tokens
metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32))
elif forward_mode.is_draft_extend(include_v2=True):
elif forward_mode.is_draft_extend_v2():
num_tokens_per_bs = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_bs
metadata.sum_seq_lens_q = num_tokens_per_bs * bs
@@ -447,7 +447,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if (
not forward_mode.is_decode_or_idle()
and not forward_mode.is_target_verify()
and not forward_mode.is_draft_extend(include_v2=True)
and not forward_mode.is_draft_extend_v2()
):
return super().init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
@@ -483,7 +483,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
):
# For extend batch with prefix length > 0, fallback to ragged kernel implemented in flashinfer MLA backend
# when chunked prefix cache is disabled.
@@ -514,7 +514,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
elif (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
bs = forward_batch.batch_size
self.forward_decode_metadata = TRTLLMMLADecodeMetadata()
@@ -522,7 +522,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
# and forward_prefill_metadata from a previous regular extend call could still be set.
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
self.forward_prefill_metadata = None
# Get maximum sequence length.
@@ -537,7 +537,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
max_seq = max_seq + self.num_draft_tokens
seq_lens = seq_lens + self.num_draft_tokens
self.forward_decode_metadata.seq_lens_k = seq_lens.to(torch.int32)
elif forward_batch.forward_mode.is_draft_extend(include_v2=True):
elif forward_batch.forward_mode.is_draft_extend_v2():
sum_seq_lens_q = sum(forward_batch.extend_seq_lens_cpu)
max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
cu_seqlens_q = torch.nn.functional.pad(
@@ -879,7 +879,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
metadata = (
getattr(forward_batch, "decode_trtllm_mla_metadata", None)
@@ -284,23 +284,6 @@ class WaveAttnBackend(AttentionBackend):
num_kv_splits = None
attn_logits = None
attn_lse = None
elif forward_batch.forward_mode.is_draft_extend():
kv_indices, kv_indptr, qo_indptr, custom_mask = (
spec_info.generate_attn_arg_prefill(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
None,
self.req_to_token,
)
)
mask_indptr = None
# TODO(FIXME): This will trigger an invalid Eagle tree when using
# `max(spec_info.num_accept_tokens_cpu)`.
# It might have been forgotten to update somewhere.
max_extend_len = torch.max(spec_info.num_accept_tokens).item()
num_kv_splits = None
attn_logits = None
attn_lse = None
else:
kv_indptr[1 : bs + 1] = torch.cumsum(
forward_batch.extend_prefix_lens, dim=0
@@ -341,10 +341,7 @@ class XPUAttentionBackend(AttentionBackend):
forward_batch.req_pool_indices, : metadata.max_seq_len_k
]
if (
any(forward_batch.extend_prefix_lens_cpu)
or forward_batch.forward_mode == ForwardMode.DRAFT_EXTEND
):
if any(forward_batch.extend_prefix_lens_cpu):
extend_seq_lens = forward_batch.extend_seq_lens
metadata.max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
metadata.cu_seqlens_q = torch.nn.functional.pad(
@@ -637,7 +634,6 @@ class XPUAttentionBackend(AttentionBackend):
if (
forward_batch.attn_attend_prefix_cache is not None
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
):
# Do multi-head attention with chunked prefix cache
if forward_batch.attn_attend_prefix_cache:
-1
View File
@@ -296,7 +296,6 @@ class AttnTpContext:
self.allow_input_scattered
and forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
and forward_batch.input_ids is not None
and not forward_batch.can_run_tbo
)
@@ -94,8 +94,6 @@ class ForwardMode(IntEnum):
# Used in speculative decoding: verify a batch in the target model.
TARGET_VERIFY = auto()
# Used in speculative decoding: extend a batch in the draft model.
DRAFT_EXTEND = auto()
DRAFT_EXTEND_V2 = auto()
# Used in disaggregated decode worker
@@ -115,7 +113,6 @@ class ForwardMode(IntEnum):
return (
self == ForwardMode.EXTEND
or self == ForwardMode.MIXED
or self == ForwardMode.DRAFT_EXTEND
or (include_draft_extend_v2 and self == ForwardMode.DRAFT_EXTEND_V2)
or self == ForwardMode.TARGET_VERIFY
or self == ForwardMode.SPLIT_PREFILL
@@ -148,11 +145,6 @@ class ForwardMode(IntEnum):
def is_target_verify(self):
return self == ForwardMode.TARGET_VERIFY
def is_draft_extend(self, include_v2: bool = False):
return self == ForwardMode.DRAFT_EXTEND or (
include_v2 and self == ForwardMode.DRAFT_EXTEND_V2
)
def is_draft_extend_v2(self):
# For fixed shape logits output in eagle v2 worker
return self == ForwardMode.DRAFT_EXTEND_V2
@@ -160,7 +152,6 @@ class ForwardMode(IntEnum):
def is_extend_or_draft_extend_or_mixed(self, include_draft_extend_v2: bool = False):
return (
self == ForwardMode.EXTEND
or self == ForwardMode.DRAFT_EXTEND
or self == ForwardMode.MIXED
or self == ForwardMode.SPLIT_PREFILL
or (include_draft_extend_v2 and self == ForwardMode.DRAFT_EXTEND_V2)
@@ -181,11 +172,7 @@ class ForwardMode(IntEnum):
return self == ForwardMode.SPLIT_PREFILL
def is_extend_without_speculative(self):
return (
self.is_extend()
and not self.is_target_verify()
and not self.is_draft_extend()
)
return self.is_extend() and not self.is_target_verify()
def is_prebuilt(self):
return self == ForwardMode.PREBUILT
@@ -943,47 +930,26 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
device = model_runner.device
mm_inputs = batch.multimodal_inputs
if batch.forward_mode.is_draft_extend(): # draft_extend_after_decode
mrope_deltas = []
extend_lens = []
for batch_idx in range(batch_size):
extend_seq_len = batch.extend_lens[batch_idx]
extend_lens.append(extend_seq_len)
mrope_delta = (
# target_verify or draft_decode
seq_positions = batch.spec_info.positions.view(batch_size, -1)
# Split text-only and mixed batches here because SpecV2 text-only batches can avoid an extra D2H.
if all(mm_input is None for mm_input in mm_inputs):
mrope_delta_tensor = torch.zeros(
(batch_size, 1), dtype=torch.int64, device=device
)
else:
mrope_deltas = [
(
torch.zeros(1, dtype=torch.int64)
if mm_inputs[batch_idx] is None
else mm_inputs[batch_idx].mrope_position_delta.squeeze(0)
if mm_inputs[i] is None
else mm_inputs[i].mrope_position_delta.squeeze(0)
)
mrope_deltas.append(mrope_delta.to(device=device))
position_chunks = torch.split(batch.spec_info.positions, extend_lens)
mrope_positions_list = [
pos_chunk + delta
for pos_chunk, delta in zip(position_chunks, mrope_deltas)
for i in range(batch_size)
]
next_input_positions = (
torch.cat(mrope_positions_list, dim=0).unsqueeze(0).repeat(3, 1)
)
else: # target_verify or draft_decode
seq_positions = batch.spec_info.positions.view(batch_size, -1)
# Split text-only and mixed batches here because SpecV2 text-only batches can avoid an extra D2H.
if all(mm_input is None for mm_input in mm_inputs):
mrope_delta_tensor = torch.zeros(
(batch_size, 1), dtype=torch.int64, device=device
)
else:
mrope_deltas = [
(
torch.zeros(1, dtype=torch.int64)
if mm_inputs[i] is None
else mm_inputs[i].mrope_position_delta.squeeze(0)
)
for i in range(batch_size)
]
mrope_delta_tensor = torch.stack(mrope_deltas, dim=0).to(device=device)
next_input_positions = (
(seq_positions + mrope_delta_tensor).flatten().unsqueeze(0).repeat(3, 1)
)
mrope_delta_tensor = torch.stack(mrope_deltas, dim=0).to(device=device)
next_input_positions = (
(seq_positions + mrope_delta_tensor).flatten().unsqueeze(0).repeat(3, 1)
)
self.mrope_positions = next_input_positions
@@ -1135,7 +1101,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if (
self.forward_mode.is_decode()
or self.forward_mode.is_target_verify()
or self.forward_mode.is_draft_extend(include_v2=True)
or self.forward_mode.is_draft_extend_v2()
or self.forward_mode.is_idle()
):
# Mamba-hybrid families need the fabricated-row idle conversion
@@ -1352,13 +1318,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
:num_tokens
]
logits_output.hidden_states = logits_output.hidden_states[:num_tokens]
elif self.forward_mode.is_draft_extend(): # draft extend
self.spec_info.num_correct_drafts = self.spec_info.num_correct_drafts[
:bs
]
self.spec_info.num_accept_tokens = self.spec_info.num_accept_tokens[:bs]
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
logits_output.hidden_states = logits_output.hidden_states[:bs]
elif self.forward_mode.is_draft_extend_v2(): # draft extend_v2
bs = bs * self.spec_info.num_tokens_per_req
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
@@ -43,7 +43,6 @@ def handle_attention_ascend(attn, forward_batch):
if (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
):
if hasattr(attn, "indexer"):
+1 -1
View File
@@ -1726,7 +1726,7 @@ class DeepseekV2AttentionMLA(
attention_backend = get_global_server_args().decode_attention_backend
elif (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
):
# Use the specified backend for speculative operations (both verify and draft extend)
if get_global_server_args().speculative_attention_mode == "decode":
+1 -1
View File
@@ -258,7 +258,7 @@ class Eagle3MLAModel(nn.Module):
if (
forward_batch.forward_mode.is_extend()
and forward_batch.contains_mm_inputs()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
):
assert embeds is not None
last_indices = (
+1 -1
View File
@@ -195,7 +195,7 @@ class LlamaModel(nn.Module):
if (
forward_batch.forward_mode.is_extend()
and forward_batch.contains_mm_inputs()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
):
assert embeds is not None
last_indices = (
-1
View File
@@ -253,7 +253,6 @@ class MindSporeForCausalLM(torch.nn.Module):
is_prefill = (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_target_verify()
)
if forward_batch.extend_prefix_lens is not None:
+1 -1
View File
@@ -161,7 +161,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
if (
forward_batch.forward_mode.is_extend()
and forward_batch.contains_mm_inputs()
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
and not forward_batch.forward_mode.is_draft_extend_v2()
):
assert input_embeds is not None
last_indices = (
+1 -1
View File
@@ -802,7 +802,7 @@ class CCA(nn.Module):
if forward_batch.forward_mode.is_decode_or_idle():
return self._forward_decode(hidden_states, forward_batch)
# EXTEND / MIXED / DRAFT_EXTEND / DLLM_EXTEND all share the prefill loop.
# EXTEND / MIXED / DLLM_EXTEND all share the prefill loop.
return self._forward_extend(hidden_states, forward_batch)
@@ -82,12 +82,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
):
# Parse args
self.eagle_worker = eagle_worker
if not hasattr(eagle_worker, "model_runner"):
self.model_runner = model_runner = eagle_worker.draft_runner
self.forward_mode = ForwardMode.DRAFT_EXTEND_V2
else:
self.model_runner = model_runner = eagle_worker.model_runner
self.forward_mode = ForwardMode.DRAFT_EXTEND
self.model_runner = model_runner = eagle_worker.draft_runner
self.forward_mode = ForwardMode.DRAFT_EXTEND_V2
# Fields the parent's capture() reads:
self.device = model_runner.device
@@ -211,11 +207,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
next_token_logits_buffer = torch.zeros(
(
(
self.max_bs * self.num_tokens_per_bs
if self.forward_mode == ForwardMode.DRAFT_EXTEND_V2
else self.max_bs
),
self.max_bs * self.num_tokens_per_bs,
vocab_size,
),
dtype=torch.float,
@@ -312,15 +304,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
)
num_correct_drafts = buffers.num_correct_drafts[:bs]
num_accept_tokens = buffers.num_accept_tokens[:bs]
next_token_logits_buffer = buffers.next_token_logits_buffer[
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
]
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
# V1 (DRAFT_EXTEND): pruned_states = bs (last token per seq)
# V2 (DRAFT_EXTEND_V2): pruned_states = num_tokens (all tokens)
num_tokens_for_logprob = (
num_tokens if self.forward_mode.is_draft_extend_v2() else bs
)
# pruned_states = num_tokens (all tokens)
num_tokens_for_logprob = num_tokens
if self.require_mlp_tp_gather:
global_num_tokens_cpu = [num_tokens] * self.dp_size
@@ -497,12 +484,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer:
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
if self.forward_mode.is_draft_extend_v2():
buffers.global_num_tokens_for_logprob_gpu.fill_(
bs * self.num_tokens_per_bs
)
else:
buffers.global_num_tokens_for_logprob_gpu.fill_(bs)
buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
@@ -559,26 +541,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
with timer_ctx:
out = self._replay_graph(shape_key, forward_batch)
if self.forward_mode == ForwardMode.DRAFT_EXTEND_V2:
unpadding_bs = num_tokens
elif bs != raw_bs:
forward_batch.spec_info.num_correct_drafts = buffers.num_correct_drafts[
:raw_bs
]
forward_batch.spec_info.num_accept_tokens = buffers.num_accept_tokens[
:raw_bs
]
unpadding_bs = raw_bs
else:
unpadding_bs = None
if unpadding_bs is not None:
out_copy = out
out = LogitsProcessorOutput(
next_token_logits=out.next_token_logits[:unpadding_bs],
hidden_states=out.hidden_states[:unpadding_bs],
)
if self.forward_mode != ForwardMode.DRAFT_EXTEND_V2:
out.topk_p = out_copy.topk_p[:raw_bs]
out.topk_index = out_copy.topk_index[:raw_bs]
out = LogitsProcessorOutput(
next_token_logits=out.next_token_logits[:num_tokens],
hidden_states=out.hidden_states[:num_tokens],
)
return out
@@ -327,9 +327,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
positions = buffers.positions[:num_tokens]
mrope_positions = buffers.mrope_positions[:, :num_tokens]
hidden_states = buffers.hidden_states[:num_tokens]
next_token_logits_buffer = buffers.next_token_logits_buffer[
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
]
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
if self.require_mlp_tp_gather:
global_num_tokens_cpu = [num_tokens] * self.dp_size
@@ -327,7 +327,7 @@ class MockModelRunner(ModelRunner):
speculative_num_draft_tokens = (
max(case.input_lens)
if case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend(include_v2=True)
or case.forward_mode.is_draft_extend_v2()
else 0
)
self.server_args = make_mock_server_args(
@@ -294,8 +294,9 @@ class DSAMockModelRunner(ModelRunner):
# `kAlignedBatchSize=0U`, which fails to compile. We auto-derive
# the draft-token count from `case.extend_lens` so the
# speculative paths produce a non-empty `seqlens_expanded`.
if case.forward_mode.is_target_verify() or case.forward_mode.is_draft_extend(
include_v2=True
if (
case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend_v2()
):
spec_num_draft_tokens = max(case.extend_lens) if case.extend_lens else 1
else:
@@ -1463,7 +1464,7 @@ def run_dsa_sparse_speculative_forward_mode_case(
) -> None:
"""Run a sparse case with a speculative forward mode (TARGET_VERIFY,
DRAFT_EXTEND, or DRAFT_EXTEND_V2). DSA dispatches both
`is_target_verify()` and `is_draft_extend(include_v2=True)` through
`is_target_verify()` and `is_draft_extend_v2()` through
`dsa_decode_impl` (`dsa_backend.py:1352-1358`), so the kernel
selection matches plain DECODE but `seqlens_expanded` is computed
differently per forward mode (`dsa_backend.py:469-529`).
@@ -1472,8 +1473,7 @@ def run_dsa_sparse_speculative_forward_mode_case(
speculative modes so deep_gemm's `paged_mqa_logits_metadata` JIT
compiles with a non-zero `kAlignedBatchSize`."""
if not (
case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend(include_v2=True)
case.forward_mode.is_target_verify() or case.forward_mode.is_draft_extend_v2()
):
raise ValueError(
"run_dsa_sparse_speculative_forward_mode_case expects a "
@@ -313,8 +313,9 @@ class MockDSV4ModelRunner:
# case's per-request input length (target_verify uses the draft count
# directly; draft_extend uses the accepted-token count). Non-spec cases
# leave it at 0 so the backend skips the speculative branches.
if case.forward_mode.is_target_verify() or case.forward_mode.is_draft_extend(
include_v2=True
if (
case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend_v2()
):
speculative_num_draft_tokens = case.input_lens[0] if case.input_lens else 0
speculative_eagle_topk = 1
@@ -1492,8 +1493,8 @@ def run_dsv4_draft_extend_attention_case(
"`deepseek_v4_backend.py:636-663` and the 'Production-Unsupported' "
"section in dsv4/README.md."
)
assert case.forward_mode.is_draft_extend(
include_v2=True
assert (
case.forward_mode.is_draft_extend_v2()
), f"run_dsv4_draft_extend_attention_case requires DRAFT_EXTEND; got {case.forward_mode}"
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
_make_eagle_draft_extend_input,
@@ -216,7 +216,7 @@ class MockGDNModelRunner(ModelRunner):
speculative_num_draft_tokens = (
case.input_lens[0]
if case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend(include_v2=True)
or case.forward_mode.is_draft_extend_v2()
else 0
)
self.server_args = make_mock_server_args(
@@ -222,7 +222,7 @@ class MockKDAModelRunner(ModelRunner):
speculative_num_draft_tokens = (
case.input_lens[0]
if case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend(include_v2=True)
or case.forward_mode.is_draft_extend_v2()
else 0
)
self.server_args = make_mock_server_args(
@@ -231,7 +231,7 @@ class MockLightningModelRunner(ModelRunner):
speculative_num_draft_tokens = (
case.input_lens[0]
if case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend(include_v2=True)
or case.forward_mode.is_draft_extend_v2()
else 0
)
self.server_args = make_mock_server_args(
@@ -329,8 +329,9 @@ class MockMamba2ModelRunner(ModelRunner):
# `intermediate_ssm` / `intermediate_conv_window` buffers when
# `speculative_num_draft_tokens is not None`, so auto-derive the
# count from `case.extend_lens` for the speculative modes.
if case.forward_mode.is_target_verify() or case.forward_mode.is_draft_extend(
include_v2=True
if (
case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend_v2()
):
speculative_num_draft_tokens = (
max(case.extend_lens) if case.extend_lens else 1
@@ -566,7 +567,7 @@ class ProjectedMamba2Attention(nn.Module):
# state support. The dense-extend path leaves it False.
use_triton_causal_conv = (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
or forward_batch.forward_mode.is_draft_extend_v2()
)
self.backend.forward(
self.mixer,
@@ -236,7 +236,7 @@ class MockMLAModelRunner(ModelRunner):
speculative_num_draft_tokens = (
max(case.input_lens)
if case.forward_mode.is_target_verify()
or case.forward_mode.is_draft_extend(include_v2=True)
or case.forward_mode.is_draft_extend_v2()
else 0
)
self.server_args = make_mock_server_args(
@@ -1861,7 +1861,7 @@ def run_dsa_eagle_draft_cuda_graph_runner_case(
# 1. Multi-query-per-request: `num_input_tokens = sum(input_lens)`.
# 2. Routes through `forward_extend` rather than `forward_decode`.
# Production picks `dsa_decode_impl` (default `flashmla_kv`)
# because `is_draft_extend(include_v2=True)` is in the
# because `is_draft_extend_v2()` is in the
# decode-impl branch (`dsa_backend.py:1352-1358`).
# 3. DraftBackendFactory returns a single `DeepseekSparseAttnBackend`
# (not a multi-step wrapper) via `_create_dsa_prefill_backend`.
@@ -20,10 +20,7 @@ from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner i
run_dense_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_dense_draft_extend_cuda_graph_case,
run_dense_draft_extend_v2_cuda_graph_case,
run_dense_eagle_draft_extend_case,
run_dense_eagle_draft_extend_cuda_graph_runner_case,
run_dense_eagle_draft_extend_v2_cuda_graph_runner_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
@@ -64,62 +61,6 @@ class TestFA3DenseAttentionBackendCorrectness(CustomTestCase):
prefix_lens=(14, 15, 16),
),
)
DRAFT_EXTEND_CASES = (
(
DenseAttentionCase(
name="runner_fa3_eagle_draft_extend",
backend="fa3",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
DenseAttentionCase(
name="runner_fa3_frozen_kv_mtp_draft_extend",
backend="fa3",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
DRAFT_EXTEND_CUDA_GRAPH_CASES = (
(
DenseAttentionCase(
name="runner_cuda_graph_fa3_eagle_draft_extend",
backend="fa3",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
DenseAttentionCase(
name="runner_cuda_graph_fa3_frozen_kv_mtp_draft_extend",
backend="fa3",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
DRAFT_EXTEND_V2_CUDA_GRAPH_CASES = (
DenseAttentionCase(
name="runner_cuda_graph_fa3_eagle_draft_extend_v2_fixed_tokens",
@@ -247,18 +188,6 @@ class TestFA3DenseAttentionBackendCorrectness(CustomTestCase):
"ngram",
),
)
EAGLE_DRAFT_EXTEND_RUNNER_CASES = (
DenseAttentionCase(
name="runner_fa3_eagle_draft_extend_cuda_graph_runner",
backend="fa3",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
)
EAGLE_DRAFT_EXTEND_V2_RUNNER_CASES = (
DenseAttentionCase(
name="runner_fa3_eagle_draft_extend_v2_cuda_graph_runner_fixed_tokens",
@@ -450,42 +379,6 @@ class TestFA3DenseAttentionBackendCorrectness(CustomTestCase):
pad_style=pad_style,
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_dense_eagle_draft_extend_case(
self,
case,
spec_kind=spec_kind,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_draft_extend_cuda_graph_cases(self):
for case, spec_kind in self.DRAFT_EXTEND_CUDA_GRAPH_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_dense_draft_extend_cuda_graph_case(
self,
case,
spec_kind=spec_kind,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_eagle_draft_extend_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dense_eagle_draft_extend_cuda_graph_runner_case(
self,
case,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_eagle_draft_extend_v2_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_V2_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
@@ -19,10 +19,7 @@ from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner i
run_dense_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_dense_draft_extend_cuda_graph_case,
run_dense_draft_extend_v2_cuda_graph_case,
run_dense_eagle_draft_extend_case,
run_dense_eagle_draft_extend_cuda_graph_runner_case,
run_dense_eagle_draft_extend_v2_cuda_graph_runner_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
@@ -59,62 +56,6 @@ class TestFA4DenseAttentionBackendCorrectness(CustomTestCase):
prefix_lens=(14, 15, 16),
),
)
DRAFT_EXTEND_CASES = (
(
DenseAttentionCase(
name="runner_fa4_eagle_draft_extend",
backend="fa4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
DenseAttentionCase(
name="runner_fa4_frozen_kv_mtp_draft_extend",
backend="fa4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
DRAFT_EXTEND_CUDA_GRAPH_CASES = (
(
DenseAttentionCase(
name="runner_cuda_graph_fa4_eagle_draft_extend",
backend="fa4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
DenseAttentionCase(
name="runner_cuda_graph_fa4_frozen_kv_mtp_draft_extend",
backend="fa4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
DRAFT_EXTEND_V2_CUDA_GRAPH_CASES = (
DenseAttentionCase(
name="runner_cuda_graph_fa4_eagle_draft_extend_v2_fixed_tokens",
@@ -242,18 +183,6 @@ class TestFA4DenseAttentionBackendCorrectness(CustomTestCase):
"ngram",
),
)
EAGLE_DRAFT_EXTEND_RUNNER_CASES = (
DenseAttentionCase(
name="runner_fa4_eagle_draft_extend_cuda_graph_runner",
backend="fa4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
)
EAGLE_DRAFT_EXTEND_V2_RUNNER_CASES = (
DenseAttentionCase(
name="runner_fa4_eagle_draft_extend_v2_cuda_graph_runner_fixed_tokens",
@@ -439,42 +368,6 @@ class TestFA4DenseAttentionBackendCorrectness(CustomTestCase):
pad_style=pad_style,
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_dense_eagle_draft_extend_case(
self,
case,
spec_kind=spec_kind,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_draft_extend_cuda_graph_cases(self):
for case, spec_kind in self.DRAFT_EXTEND_CUDA_GRAPH_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_dense_draft_extend_cuda_graph_case(
self,
case,
spec_kind=spec_kind,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_eagle_draft_extend_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dense_eagle_draft_extend_cuda_graph_runner_case(
self,
case,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_eagle_draft_extend_v2_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_V2_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
@@ -19,11 +19,6 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_dense_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_dense_draft_extend_cuda_graph_case,
run_dense_eagle_draft_extend_case,
run_dense_eagle_draft_extend_cuda_graph_runner_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
run_dense_eagle_draft_cuda_graph_runner_case,
run_dense_frozen_kv_mtp_cuda_graph_runner_case,
@@ -251,74 +246,6 @@ class TestFlashInferDenseAttentionBackendCorrectness(CustomTestCase):
"ngram",
),
)
EAGLE_DRAFT_EXTEND_CASES = (
(
DenseAttentionCase(
name="runner_eagle_draft_extend_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(2, 5),
extend_lens=(1, 3),
),
"eagle",
),
(
DenseAttentionCase(
name="runner_frozen_kv_mtp_draft_extend_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(2, 5),
extend_lens=(1, 3),
),
"frozen_kv_mtp",
),
)
DRAFT_EXTEND_CUDA_GRAPH_CASES = (
(
DenseAttentionCase(
name="runner_cuda_graph_eagle_draft_extend_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(2, 5),
extend_lens=(1, 3),
),
"eagle",
),
(
DenseAttentionCase(
name="runner_cuda_graph_frozen_kv_mtp_draft_extend_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(2, 5),
extend_lens=(1, 3),
),
"frozen_kv_mtp",
),
)
EAGLE_DRAFT_EXTEND_RUNNER_CASES = (
DenseAttentionCase(
name="runner_eagle_draft_extend_cuda_graph_runner_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(2, 5),
extend_lens=(2, 4),
),
)
EAGLE_DRAFT_RUNNER_CASES = (
(
DenseAttentionCase(
@@ -475,46 +402,6 @@ class TestFlashInferDenseAttentionBackendCorrectness(CustomTestCase):
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.EAGLE_DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name,
backend=case.backend,
spec_kind=spec_kind,
):
run_dense_eagle_draft_extend_case(
self,
case,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
spec_kind=spec_kind,
)
def test_runner_mode_draft_extend_cuda_graph_cases(self):
for case, spec_kind in self.DRAFT_EXTEND_CUDA_GRAPH_CASES:
with self.subTest(
case=case.name,
backend=case.backend,
spec_kind=spec_kind,
):
run_dense_draft_extend_cuda_graph_case(
self,
case,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
spec_kind=spec_kind,
)
def test_runner_mode_eagle_draft_extend_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dense_eagle_draft_extend_cuda_graph_runner_case(
self,
case,
head_dim=self.HEAD_DIM,
hidden_size=self.HIDDEN_SIZE,
)
def test_runner_mode_eagle_draft_cuda_graph_runner_cases(self):
for case, topk, num_draft_tokens in self.EAGLE_DRAFT_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend, topk=topk):
@@ -31,9 +31,6 @@ from sglang.test.kits.attention_unittest.attention_methods.dsa_attention import
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_dsa_sparse_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_dsa_eagle_draft_extend_cuda_graph_runner_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
run_dsa_eagle_draft_cuda_graph_runner_case,
)
@@ -217,8 +214,8 @@ class TestDSAAttentionBackendCorrectness(CustomTestCase):
self, self.DECODE_IMPL_CASE, impl
)
# Speculative forward-mode coverage. TARGET_VERIFY, DRAFT_EXTEND,
# and DRAFT_EXTEND_V2 all route through the `dsa_decode_impl`
# Speculative forward-mode coverage. TARGET_VERIFY and
# DRAFT_EXTEND_V2 both route through the `dsa_decode_impl`
# dispatcher (the same kernel selection as plain DECODE) but
# produce different `seqlens_expanded` and `cu_seqlens_q` from
# `dsa_backend.py:469-529`. `DSAMockModelRunner.__init__` derives
@@ -235,16 +232,6 @@ class TestDSAAttentionBackendCorrectness(CustomTestCase):
prefix_lens=(128,),
extend_lens=(3,),
),
DSAAttentionCase(
name="dsa_sparse_draft_extend",
backend="dsa",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=1,
page_size=DSA_PAGE_SIZE,
prefix_lens=(128,),
extend_lens=(3,),
),
DSAAttentionCase(
name="dsa_sparse_draft_extend_v2",
backend="dsa",
@@ -382,29 +369,6 @@ class TestDSAAttentionBackendCorrectness(CustomTestCase):
with self.subTest(case=case.name, backend=case.backend):
run_dsa_eagle_draft_cuda_graph_runner_case(self, case)
# EAGLE production draft-extend CUDA-graph runner. Routes through
# `DraftBackendFactory._create_dsa_prefill_backend` which returns a
# single `DeepseekSparseAttnBackend` (not multi-step); the forward
# goes through `forward_extend` with `dsa_decode_impl` selected via
# `is_draft_extend(include_v2=True)`.
EAGLE_DRAFT_EXTEND_CASES = (
DSAAttentionCase(
name="runner_eagle_draft_extend_cuda_graph_dsa",
backend="dsa",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
num_kv_heads=1,
page_size=DSA_PAGE_SIZE,
prefix_lens=(128, 192),
extend_lens=(2, 3),
),
)
def test_runner_mode_eagle_draft_extend_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dsa_eagle_draft_extend_cuda_graph_runner_case(self, case)
# CG decode replay with FP8 KV cache. Captures and replays through
# `flashmla_kv` (the only FP8-compatible decode kernel). The
# `_clone_dsa_sparse_cache` hook is reused as-is — it snapshots the
@@ -34,16 +34,11 @@ from sglang.test.kits.attention_unittest.attention_methods.dsv4_attention import
make_dsv4_cases,
run_dsv4_attention_case,
run_dsv4_compress_attention_case,
run_dsv4_draft_extend_attention_case,
run_dsv4_target_verify_attention_case,
)
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import ( # noqa: E402
run_dsv4_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import ( # noqa: E402
run_dsv4_eagle_draft_extend_cuda_graph_case,
run_dsv4_eagle_draft_extend_cuda_graph_runner_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import ( # noqa: E402
run_dsv4_eagle_draft_cuda_graph_runner_case,
)
@@ -246,48 +241,6 @@ class TestDSV4AttentionBackendCorrectness(CustomTestCase):
):
run_dsv4_eagle_verify_cuda_graph_case(self, case, topk=1)
# EAGLE DRAFT_EXTEND is SWA-only for DSV4 (see runner docstring).
DRAFT_EXTEND_CASES = (
DSV4AttentionCase(
name="dsv4_swa_eagle_draft_extend",
backend="dsv4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=64,
page_size=DSV4_PAGE_SIZE,
prefix_lens=(64, 96),
extend_lens=(2, 4),
),
)
def test_eagle_draft_extend_cases(self):
for case in self.DRAFT_EXTEND_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dsv4_draft_extend_attention_case(self, case)
# CUDA-graph capture/replay for EAGLE DRAFT_EXTEND — SWA only
# (init_forward_metadata_draft_extend uses need_compress=False; see
# `Production-Unsupported` in dsv4/README.md). Uniform `extend_lens`
# because DSV4 `forward(compress_ratio=0)` asserts
# `swa_page_indices.shape[0] == q.shape[0]` and the graph metadata
# builder uses uniform `num_tokens_per_bs = max_num_tokens // max_bs`
# (see `deepseek_v4_backend.py:646-647`).
EAGLE_DRAFT_EXTEND_CUDA_GRAPH_CASES = (
DSV4AttentionCase(
name="runner_cuda_graph_dsv4_swa_eagle_draft_extend",
backend="dsv4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=64,
page_size=DSV4_PAGE_SIZE,
prefix_lens=(64, 96),
extend_lens=(4, 4),
),
)
def test_runner_mode_eagle_draft_extend_cuda_graph_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_CUDA_GRAPH_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dsv4_eagle_draft_extend_cuda_graph_case(self, case)
# Production EAGLE draft graph runner (chain only, SWA only). The runner
# routes through `DeepseekV4MultiStepBackend` (one `DeepseekV4AttnBackend`
# per draft step), captures a fixed batch, and replays distinct request
@@ -309,34 +262,6 @@ class TestDSV4AttentionBackendCorrectness(CustomTestCase):
with self.subTest(case=case.name, backend=case.backend):
run_dsv4_eagle_draft_cuda_graph_runner_case(self, case)
# Production EAGLE draft-extend graph runner (SWA only). Routes through
# the prefill-side `DeepseekV4AttnBackend` (single backend, not
# multi-step); `init_forward_metadata_draft_extend` forces
# `need_compress=False` so C4/C128 is structurally unreachable for this
# path.
# Uniform `extend_lens` because the DSV4 graph contract requires
# `q.shape[0] == swa_page_indices.shape[0]` and the
# `init_forward_metadata_draft_extend` graph path uses
# `num_tokens_per_bs = max_num_tokens // max_bs` (see
# `deepseek_v4_backend.py:646-647`). Same constraint as the metadata-
# style draft_extend CG case.
PRODUCTION_EAGLE_DRAFT_EXTEND_RUNNER_CASES = (
DSV4AttentionCase(
name="runner_production_eagle_draft_extend_dsv4_swa",
backend="dsv4",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=64,
page_size=DSV4_PAGE_SIZE,
prefix_lens=(64, 96),
extend_lens=(4, 4),
),
)
def test_runner_mode_production_eagle_draft_extend_cuda_graph_runner_cases(self):
for case in self.PRODUCTION_EAGLE_DRAFT_EXTEND_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_dsv4_eagle_draft_extend_cuda_graph_runner_case(self, case)
class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
"""CPU-only checks for the DSV4 BCG metadata replay contract."""
@@ -25,9 +25,6 @@ from sglang.test.kits.attention_unittest.attention_methods.gdn_attention import
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_gdn_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_gdn_eagle_draft_extend_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import (
run_gdn_eagle_verify_case,
run_gdn_eagle_verify_cuda_graph_case,
@@ -306,48 +303,6 @@ class TestTritonGDNBackendCorrectness(CustomTestCase):
self, case, topk=topk, spec_kind=spec_kind
)
# EAGLE / Frozen-KV MTP DRAFT_EXTEND eager — `HybridLinearAttnBackend`
# raises `ValueError("Invalid forward mode")` for DRAFT_EXTEND CG
# capture (`hybrid_linear_attn_backend.py:509,572`), so CG is
# structurally blocked across the family (GDN/KDA/Lightning/Mamba2).
# The EXTEND-style gated-delta recurrence reference doubles as the
# DRAFT_EXTEND reference across both spec kinds.
EAGLE_DRAFT_EXTEND_CASES = (
(
GDNAttentionCase(
name="runner_eagle_draft_extend_gdn",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_k_heads=2,
num_v_heads=2,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
GDNAttentionCase(
name="runner_frozen_kv_mtp_draft_extend_gdn",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_k_heads=2,
num_v_heads=2,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.EAGLE_DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_gdn_eagle_draft_extend_case(self, case, spec_kind=spec_kind)
# Spy directly on each sub-backend's `init_forward_metadata*` so
# dispatch-layer slice mutations show up as a missing call, which
# forward-output assertions can miss when the fixture happens to
@@ -19,9 +19,6 @@ from sglang.test.kits.attention_unittest.attention_methods.kda_attention import
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_kda_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_kda_eagle_draft_extend_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import (
run_kda_eagle_verify_case,
run_kda_eagle_verify_cuda_graph_case,
@@ -239,44 +236,6 @@ class TestTritonKDABackendCorrectness(CustomTestCase):
32,
),
)
# EAGLE / Frozen-KV MTP DRAFT_EXTEND eager. CG is structurally
# blocked across the HybridLinearAttn family
# (`hybrid_linear_attn_backend.py:509,572`).
EAGLE_DRAFT_EXTEND_CASES = (
(
KDAAttentionCase(
name="runner_eagle_draft_extend_kda",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_k_heads=2,
num_v_heads=2,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
KDAAttentionCase(
name="runner_frozen_kv_mtp_draft_extend_kda",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_k_heads=2,
num_v_heads=2,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.EAGLE_DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_kda_eagle_draft_extend_case(self, case, spec_kind=spec_kind)
@unittest.skipIf(
is_hip(),
@@ -18,9 +18,6 @@ from sglang.test.kits.attention_unittest.attention_methods.lightning_attention i
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_lightning_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_lightning_eagle_draft_extend_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import (
run_lightning_eagle_verify_case,
run_lightning_eagle_verify_cuda_graph_case,
@@ -186,42 +183,6 @@ class TestTritonLightningBackendCorrectness(CustomTestCase):
):
run_lightning_eagle_verify_cuda_graph_case(self, case, topk=topk)
# EAGLE / Frozen-KV MTP DRAFT_EXTEND eager — CG is structurally
# blocked across the HybridLinearAttn family.
EAGLE_DRAFT_EXTEND_CASES = (
(
LightningAttentionCase(
name="runner_eagle_draft_extend_lightning",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=2,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"eagle",
),
(
LightningAttentionCase(
name="runner_frozen_kv_mtp_draft_extend_lightning",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=2,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
"frozen_kv_mtp",
),
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.EAGLE_DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_lightning_eagle_draft_extend_case(self, case, spec_kind=spec_kind)
# PCG/BCG split-op extend is deliberately NOT covered. Lightning's
# backend `forward_extend` flattens the output via `o.view(-1,
# tp_q_head_num * v_head_dim)` (`lightning_backend.py:335`), so eager
@@ -32,9 +32,6 @@ from sglang.test.kits.attention_unittest.attention_methods.mamba2_attention impo
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_mamba2_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_mamba2_eagle_draft_extend_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import (
run_mamba2_eagle_verify_case,
run_mamba2_eagle_verify_cuda_graph_case,
@@ -128,33 +125,6 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase):
1,
),
)
# EAGLE / Frozen-KV MTP DRAFT_EXTEND eager — `HybridLinearAttnBackend`
# raises `ValueError("Invalid forward mode")` for DRAFT_EXTEND CG
# capture (`hybrid_linear_attn_backend.py:509,572`), so CG is
# structurally blocked; only eager is exercised. Same EXTEND-style
# chunked-scan recurrence reference doubles as the DRAFT_EXTEND
# reference, like the verify path.
EAGLE_DRAFT_EXTEND_CASES = tuple(
(
Mamba2AttentionCase(
name=f"runner_{spec_kind}_draft_extend_mamba2",
backend="triton",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=DEFAULT_NUM_HEADS,
head_dim=DEFAULT_HEAD_DIM,
state_size=DEFAULT_STATE_SIZE,
n_groups=DEFAULT_N_GROUPS,
conv_kernel=DEFAULT_CONV_KERNEL,
mamba_chunk_size=DEFAULT_MAMBA_CHUNK_SIZE,
hidden_size=DEFAULT_HIDDEN_SIZE,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
),
spec_kind,
)
for spec_kind in ("eagle", "frozen_kv_mtp")
)
def test_projected_mamba2_attention_cases(self):
for case in self.CASES:
@@ -194,13 +164,6 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase):
with self.subTest(case=case.name, backend=case.backend, topk=topk):
run_mamba2_eagle_verify_cuda_graph_case(self, case, topk=topk)
def test_runner_mode_eagle_draft_extend_cases(self):
for case, spec_kind in self.EAGLE_DRAFT_EXTEND_CASES:
with self.subTest(
case=case.name, backend=case.backend, spec_kind=spec_kind
):
run_mamba2_eagle_draft_extend_case(self, case, spec_kind=spec_kind)
# PCG/BCG split-op extend is deliberately NOT covered. The
# `MambaMixer2.forward` asserts `num_actual_tokens ==
# projected_states.shape[0]` (`mamba.py:467`) — the projection step
@@ -17,11 +17,6 @@ from sglang.test.kits.attention_unittest.attention_methods.mla_attention import
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
run_mla_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_mla_draft_extend_cuda_graph_case,
run_mla_eagle_draft_extend_case,
run_mla_eagle_draft_extend_cuda_graph_runner_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
run_mla_eagle_draft_cuda_graph_runner_case,
)
@@ -101,39 +96,6 @@ class TestFlashInferMLAAttentionBackendCorrectness(CustomTestCase):
1,
),
)
DRAFT_EXTEND_CASES = (
MLAAttentionCase(
name="runner_eagle_draft_extend_mla_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
page_size=16,
prefix_lens=(5, 8),
extend_lens=(2, 4),
),
)
DRAFT_EXTEND_CUDA_GRAPH_CASES = (
MLAAttentionCase(
name="runner_cuda_graph_eagle_draft_extend_mla_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
page_size=16,
prefix_lens=(5, 8),
extend_lens=(2, 4),
),
)
EAGLE_DRAFT_EXTEND_RUNNER_CASES = (
MLAAttentionCase(
name="runner_eagle_draft_extend_mla_cuda_graph_runner_ragged_accept",
backend="flashinfer",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
page_size=16,
prefix_lens=(5, 8),
extend_lens=(2, 4),
),
)
EAGLE_DRAFT_RUNNER_CASES = (
(
MLAAttentionCase(
@@ -259,29 +221,6 @@ class TestFlashInferMLAAttentionBackendCorrectness(CustomTestCase):
**MLA_SHAPE_KWARGS,
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case in self.DRAFT_EXTEND_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_mla_eagle_draft_extend_case(self, case, **MLA_SHAPE_KWARGS)
def test_runner_mode_eagle_draft_extend_cuda_graph_cases(self):
for case in self.DRAFT_EXTEND_CUDA_GRAPH_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_mla_draft_extend_cuda_graph_case(
self,
case,
**MLA_SHAPE_KWARGS,
)
def test_runner_mode_eagle_draft_extend_cuda_graph_runner_cases(self):
for case in self.EAGLE_DRAFT_EXTEND_RUNNER_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_mla_eagle_draft_extend_cuda_graph_runner_case(
self,
case,
**MLA_SHAPE_KWARGS,
)
def test_runner_mode_eagle_draft_cuda_graph_runner_cases(self):
# Backend gate (KNOWN_FAILURES.md §3): FlashInfer MLA multi-step
# draft CG capture/replay produces numerically wrong outputs on
@@ -21,9 +21,6 @@ from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner i
_init_cuda_graph_replay_metadata,
run_mla_cuda_graph_decode_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_extend_runner import (
run_mla_eagle_draft_extend_case,
)
from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import (
run_mla_eagle_draft_cuda_graph_runner_case,
)
@@ -203,17 +200,6 @@ class TestFlashMLAAttentionBackendCorrectness(CustomTestCase):
1,
),
)
DRAFT_EXTEND_CASES = (
MLAAttentionCase(
name="runner_eagle_draft_extend_mla_flashmla_ragged_accept",
backend="flashmla",
forward_mode=ForwardMode.DRAFT_EXTEND,
num_heads=4,
page_size=64,
prefix_lens=(5, 8),
extend_lens=(2, 4),
),
)
EAGLE_DRAFT_RUNNER_CASES = (
(
MLAAttentionCase(
@@ -338,11 +324,6 @@ class TestFlashMLAAttentionBackendCorrectness(CustomTestCase):
**MLA_SHAPE_KWARGS,
)
def test_runner_mode_eagle_draft_extend_cases(self):
for case in self.DRAFT_EXTEND_CASES:
with self.subTest(case=case.name, backend=case.backend):
run_mla_eagle_draft_extend_case(self, case, **MLA_SHAPE_KWARGS)
@unittest.skipIf(_DECODE_REQUIRES_SM90A, _DECODE_SKIP_REASON)
def test_runner_mode_eagle_draft_cuda_graph_runner_cases(self):
for case, topk, num_draft_tokens in self.EAGLE_DRAFT_RUNNER_CASES:
@@ -26,7 +26,7 @@ class TestTokenOracleManager(CustomTestCase):
) -> None:
"""Verify EAGLE draft extend maps one request row to every draft token."""
forward_batch = SimpleNamespace(
forward_mode=ForwardMode.DRAFT_EXTEND,
forward_mode=ForwardMode.DRAFT_EXTEND_V2,
spec_info=SimpleNamespace(num_tokens_per_req=4),
rids_int=torch.tensor([3, 7], dtype=torch.int64, device=self.device),
bootstrap_room_ids_int=None,