From e5f5d84780f8926cad7921cc060a7a75a99a0ba9 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Thu, 28 May 2026 00:56:23 -0700 Subject: [PATCH] Fix FA DRAFT_EXTEND_V2 cache extent (#26512) Co-authored-by: Claude Opus 4.7 (1M context) --- .../attention/flashattention_backend.py | 55 ++++++++++++++++--- 1 file changed, 46 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 1d0d1b8ba..db3e1ce45 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -503,10 +503,33 @@ class FlashAttentionBackend(AttentionBackend): elif forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed( include_draft_extend_v2=True ): - metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + # DRAFT_EXTEND_V2: seq_lens = prefix_lens; effective KV extent is prefix + extend. + if forward_batch.forward_mode.is_draft_extend_v2(): + effective_cache_seqlens = ( + seqlens_in_batch + forward_batch.extend_seq_lens + ) + seq_lens_cpu = forward_batch.seq_lens_cpu + extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu + if extend_seq_lens_cpu is not None: + extend_cpu_tensor = torch.as_tensor( + extend_seq_lens_cpu, dtype=seq_lens_cpu.dtype + ) + effective_max_seq_len_k = int( + (seq_lens_cpu + extend_cpu_tensor) + .max() + .item() # per-request sum, not max+max + ) + else: + effective_max_seq_len_k = int(effective_cache_seqlens.max().item()) + else: + effective_cache_seqlens = seqlens_in_batch + effective_max_seq_len_k = int(forward_batch.seq_lens_cpu.max().item()) + + metadata.cache_seqlens_int32 = effective_cache_seqlens.to(torch.int32) + metadata.max_seq_len_k = effective_max_seq_len_k metadata.cu_seqlens_k = torch.nn.functional.pad( - torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) + torch.cumsum(effective_cache_seqlens, dim=0, dtype=torch.int32), + (1, 0), ) # MLA/MHA CP: prepare_mlp_sync_batch pads extend tokens up to @@ -2267,13 +2290,8 @@ class FlashAttentionBackend(AttentionBackend): elif forward_mode.is_draft_extend_v2(): metadata = self.draft_extend_metadata[bs] - metadata.cache_seqlens_int32.copy_(seq_lens) - - metadata.max_seq_len_k = seq_lens_cpu.max().item() - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) - ) + # DRAFT_EXTEND_V2: seq_lens = prefix_lens; effective KV extent is prefix + extend. extend_seq_lens_tensor = getattr(spec_info, "extend_seq_lens_tensor", None) extend_seq_lens_cpu = getattr(spec_info, "extend_seq_lens_cpu", None) if extend_seq_lens_tensor is not None: @@ -2293,6 +2311,25 @@ class FlashAttentionBackend(AttentionBackend): ) extend_seq_lens_cpu = [default_extend] * bs + effective_cache_seqlens = seq_lens.to(torch.int32) + extend_seq_lens + metadata.cache_seqlens_int32.copy_(effective_cache_seqlens) + + if extend_seq_lens_cpu is not None: + extend_cpu_tensor = torch.as_tensor( + extend_seq_lens_cpu, dtype=seq_lens_cpu.dtype + ) + metadata.max_seq_len_k = int( + (seq_lens_cpu + extend_cpu_tensor) + .max() + .item() # per-request sum, not max+max + ) + else: + metadata.max_seq_len_k = int(effective_cache_seqlens.max().item()) + + metadata.cu_seqlens_k[1:].copy_( + torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) + ) + if extend_seq_lens_cpu: metadata.max_seq_len_q = int(max(extend_seq_lens_cpu)) else: