[Fix] Unify ForwardBatch extend lens cpu fields to their declared list type (#30896)

This commit is contained in:
Liangsheng Yin
2026-07-11 17:30:53 -05:00
committed by GitHub
parent d8ef76682e
commit 4884f6fbee
6 changed files with 12 additions and 29 deletions
@@ -714,12 +714,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
)
# Query-side max length, sourced from the host-resident extend lengths
# (sync-free); for plain prefill these equal the full 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.
max_q = max(forward_batch.extend_seq_lens_cpu)
metadata.max_seq_len_q = (
int(max_q.item()) if isinstance(max_q, torch.Tensor) else int(max_q)
)
metadata.max_seq_len_q = int(max(forward_batch.extend_seq_lens_cpu))
if (
forward_batch.extend_prefix_lens_cpu is not None
and any(forward_batch.extend_prefix_lens_cpu)
@@ -1267,8 +1267,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
self.extend_start_loc = torch.arange(
bs, dtype=torch.int32, device=self.seq_lens.device
)
self.extend_prefix_lens_cpu = self.extend_prefix_lens.cpu()
self.extend_seq_lens_cpu = self.extend_seq_lens.cpu()
self.extend_prefix_lens_cpu = self.extend_prefix_lens.cpu().tolist()
self.extend_seq_lens_cpu = self.extend_seq_lens.cpu().tolist()
self.extend_logprob_start_lens_cpu = self.extend_prefix_lens_cpu
else:
if self.spec_info is not None:
@@ -728,11 +728,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
extend_seq_lens=shape_inputs["extend_seq_lens"],
extend_prefix_lens=shape_inputs["extend_prefix_lens"],
extend_start_loc=shape_inputs["extend_start_loc"],
extend_prefix_lens_cpu=torch.zeros(
(bs,), dtype=torch.int64, device="cpu"
),
extend_seq_lens_cpu=torch.tensor(lens_cpu, device="cpu"),
extend_logprob_start_lens_cpu=torch.tensor(lens_cpu, device="cpu"),
extend_prefix_lens_cpu=[0] * bs,
extend_seq_lens_cpu=list(lens_cpu),
extend_logprob_start_lens_cpu=list(lens_cpu),
positions=_slot("positions"),
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,