Fix TRTLLM MHA draft decode cache seqlens replay (#26655)

This commit is contained in:
Lianmin Zheng
2026-05-28 20:58:16 -07:00
committed by GitHub
parent 79c844527c
commit dc4e7bc479
@@ -458,16 +458,19 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# Draft Decode # Draft Decode
# Here we only support topk = 1 for now. # Here we only support topk = 1 for now.
metadata = self.decode_cuda_graph_metadata[bs] metadata = self.decode_cuda_graph_metadata[bs]
max_len = seq_lens_cpu.max().item() metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
metadata.max_seq_len_k = max_len + self.speculative_step_id + 1 "cache_seqlens"
][:bs]
metadata.cache_seqlens_int32.copy_(
seq_lens + self.speculative_step_id + 1
)
metadata.max_seq_len_k = seq_lens.max().item() + (
self.speculative_step_id + 1
)
max_seq_pages = ( max_seq_pages = (
metadata.max_seq_len_k + self.page_size - 1 metadata.max_seq_len_k + self.page_size - 1
) // self.page_size ) // self.page_size
metadata.cache_seqlens_int32.copy_(
seq_lens + self.speculative_step_id + 1
)
else: else:
# Normal Decode # Normal Decode
metadata = self.decode_cuda_graph_metadata[bs] metadata = self.decode_cuda_graph_metadata[bs]