From dc4e7bc479df5a93cb8790bcd7fc1ae07780b771 Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Thu, 28 May 2026 20:58:16 -0700 Subject: [PATCH] Fix TRTLLM MHA draft decode cache seqlens replay (#26655) --- .../srt/layers/attention/trtllm_mha_backend.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index b74a11846..e270f2ccb 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -458,16 +458,19 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): # Draft Decode # Here we only support topk = 1 for now. metadata = self.decode_cuda_graph_metadata[bs] - max_len = seq_lens_cpu.max().item() - metadata.max_seq_len_k = max_len + self.speculative_step_id + 1 + metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ + "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 = ( metadata.max_seq_len_k + self.page_size - 1 ) // self.page_size - - metadata.cache_seqlens_int32.copy_( - seq_lens + self.speculative_step_id + 1 - ) else: # Normal Decode metadata = self.decode_cuda_graph_metadata[bs]