Fix DRAFT_EXTEND_V2 CG metadata: align test fixture and Triton with production seq_lens convention (#26651)

Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-29 02:46:45 -07:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent eb5d4827e8
commit ec075d8bc5
6 changed files with 52 additions and 38 deletions
@@ -704,12 +704,26 @@ class TritonAttnBackend(AttentionBackend):
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
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 = (
spec_info.extend_seq_lens_tensor[:bs].to(torch.int32)
if spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
else torch.zeros(bs, dtype=torch.int32, device=self.device)
)
kv_lens = (seq_lens - extend_seq_lens).to(torch.int32)
else:
# DRAFT_EXTEND_V1: seq_lens = prefix only.
kv_lens = seq_lens
kv_indptr[1 : bs + 1] = torch.cumsum(kv_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_lens,
kv_indptr,
None,
kv_indices,
@@ -859,12 +873,31 @@ class TritonAttnBackend(AttentionBackend):
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
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.
# 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.
assert (
spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
), "DRAFT_EXTEND_V2 replay requires spec_info.extend_seq_lens_tensor"
kv_lens = torch.clamp(
seq_lens - spec_info.extend_seq_lens_tensor[:bs].to(torch.int32),
min=0,
).to(torch.int32)
else:
# DRAFT_EXTEND_V1: seq_lens = prefix only.
kv_lens = seq_lens
kv_indptr[1 : bs + 1] = torch.cumsum(kv_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_lens,
kv_indptr,
None,
kv_indices,
@@ -226,10 +226,12 @@ def _make_eagle_draft_extend_v2_input(case, batch, *, device: str):
def _set_draft_extend_v2_prefix_lens(batch, case, *, device: str):
prefix_lens = torch.tensor(case.prefix_lens, dtype=torch.int32, device=device)
batch.seq_lens = prefix_lens
batch.seq_lens_cpu = torch.tensor(case.prefix_lens, dtype=torch.int32, device="cpu")
batch.seq_lens_sum = sum(case.prefix_lens)
# Production sets seq_lens = prefix + extend before init_forward_metadata
# (eagle_info_v2.py bumps seq_lens by num_draft_tokens). Match that here.
seq_lens = tuple(p + e for p, e in zip(case.prefix_lens, case.input_lens))
batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device)
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32, device="cpu")
batch.seq_lens_sum = sum(seq_lens)
def _prepare_draft_extend_batch(
@@ -1415,9 +1417,12 @@ def _set_draft_extend_v2_prefix_lens(
*,
device: str,
) -> None:
batch.seq_lens = torch.tensor(case.prefix_lens, dtype=torch.int32, device=device)
batch.seq_lens_cpu = torch.tensor(case.prefix_lens, dtype=torch.int32, device="cpu")
batch.seq_lens_sum = sum(case.prefix_lens)
# Production sets seq_lens = prefix + extend before init_forward_metadata
# (eagle_info_v2.py bumps seq_lens by num_draft_tokens). Match that here.
seq_lens = tuple(p + e for p, e in zip(case.prefix_lens, case.input_lens))
batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device)
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32, device="cpu")
batch.seq_lens_sum = sum(seq_lens)
def _make_dense_eagle_draft_extend_forward_batch(