diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 418ed488b..cffc95da1 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -345,7 +345,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_v2(): - num_tokens_per_bs = num_tokens // bs + 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 metadata.cu_seqlens_q = torch.arange( @@ -385,24 +385,13 @@ 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)) + metadata.seq_lens_k.copy_(seq_lens) 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 - metadata.cu_seqlens_q[: bs + 1].copy_( - torch.arange( - 0, - bs * num_tokens_per_bs + 1, - step=num_tokens_per_bs, - dtype=torch.int32, - device=seq_lens.device, - ) - ) - metadata.seq_lens_q[:bs].fill_(num_tokens_per_bs) - # see NOTE(draft_extend seq_len handling) - seq_lens = seq_lens[:bs] - metadata.seq_lens_q[:bs] + metadata.max_seq_len_q - metadata.seq_lens_k.copy_(seq_lens.to(torch.int32)) + seq_lens = seq_lens[:bs] + metadata.seq_lens_k.copy_(seq_lens) # Update block indices for new sequences. create_flashmla_kv_indices_triton[