From e381312664b1227aca2a18c7322d4819804b2079 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Thu, 28 May 2026 18:45:33 -0700 Subject: [PATCH] Revert "Fix FA DRAFT_EXTEND_V2 cache extent" (#26628) Co-authored-by: Claude Sonnet 4.6 --- .../attention/flashattention_backend.py | 55 +++---------------- .../attention/unittest/KNOWN_FAILURES.md | 2 + .../attention/unittest/dense/README.md | 4 +- .../attention/unittest/dense/test_fa3.py | 6 ++ .../attention/unittest/dense/test_fa4.py | 6 ++ 5 files changed, 25 insertions(+), 48 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index db3e1ce45..1d0d1b8ba 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -503,33 +503,10 @@ class FlashAttentionBackend(AttentionBackend): elif forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed( include_draft_extend_v2=True ): - # 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.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) + metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() metadata.cu_seqlens_k = torch.nn.functional.pad( - torch.cumsum(effective_cache_seqlens, dim=0, dtype=torch.int32), - (1, 0), + torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) ) # MLA/MHA CP: prepare_mlp_sync_batch pads extend tokens up to @@ -2290,8 +2267,13 @@ 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: @@ -2311,25 +2293,6 @@ 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: diff --git a/test/registered/attention/unittest/KNOWN_FAILURES.md b/test/registered/attention/unittest/KNOWN_FAILURES.md index 3e67c9c0a..64c679d23 100644 --- a/test/registered/attention/unittest/KNOWN_FAILURES.md +++ b/test/registered/attention/unittest/KNOWN_FAILURES.md @@ -161,6 +161,8 @@ fixture investigation; no test in the suite). | Backend | Mode | Status | Root cause | |---|---|---|---| +| FA3 | `DRAFT_EXTEND_V2` CUDA-graph replay | `[gated]` `dense/test_fa3.py::test_runner_mode_eagle_draft_extend_v2_cuda_graph_cases` | `init_forward_metadata_replay_cuda_graph` sets `cache_seqlens_int32 = seq_lens` (prefix only); effective KV extent must be `prefix + extend` for DRAFT_EXTEND_V2. Kernel reads prefix-only KV → ~82% wrong values. Fix: `cache_seqlens = seq_lens + extend_seq_lens` in `flashattention_backend.py`. | +| FA4 | `DRAFT_EXTEND_V2` CUDA-graph replay | `[gated]` `dense/test_fa4.py::test_runner_mode_eagle_draft_extend_v2_cuda_graph_cases` | Same root cause as FA3 (FA4 inherits the same `init_forward_metadata_replay_cuda_graph` path). | | FlashInfer MLA | EAGLE draft CG, chain | `[gated on SM≥10]` `mla/test_flashinfer.py::test_runner_mode_eagle_draft_cuda_graph_runner_cases` | FlashInfer MLA decode kernel in container targets SM9x; on Blackwell falls back to a generic path that doesn't restore metadata buffers under graph replay (~22 abs-diff vs reference) | | FlashMLA | MLA `DRAFT_EXTEND` CUDA-graph replay | `[no test]` (`mla/README.md` Next Work) | Capture falls through to `FlashInferMLAAttnBackend.init_forward_metadata_capture_cuda_graph` (1D `cuda_graph_kv_indices`); FlashMLA decode uses 2D `[max_bs, (max_context + PAGE_SIZE) // PAGE_SIZE]` layout — buffer mismatch | | GDN / KDA / Lightning / Mamba2 | `DRAFT_EXTEND` and `DRAFT_EXTEND_V2` graph capture | `[no test]` for CG; eager-only paths covered | `HybridLinearAttnBackend` raises `ValueError("Invalid forward mode")` at `hybrid_linear_attn_backend.py:509,572` | diff --git a/test/registered/attention/unittest/dense/README.md b/test/registered/attention/unittest/dense/README.md index 0b2f12678..09d07ef22 100644 --- a/test/registered/attention/unittest/dense/README.md +++ b/test/registered/attention/unittest/dense/README.md @@ -17,8 +17,8 @@ Columns are runner modes; rows are attention backends. Cells use: | `torch_native` | ✓ full MHA/GQA/MQA input sweep + decode/extend runner-eager cases | — (no `init_cuda_graph_state` / capture / replay hooks) | — (no CG path) | — (no CG path) | deferred: extend-metadata mismatch in `TARGET_VERIFY` reference | — | — | — | — | — | — | — | | `triton` | ✓ MHA/GQA/MQA + 10 input layouts (page 1/16/32, prefix/decode edges) | ✓ MHA/GQA/MQA decode page-boundary | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | ✓ EAGLE chain+tree, Frozen-KV-MTP chain, DFlash chain, NGRAM chain | ✓ EAGLE tree, DFlash chain, NGRAM chain | deferred: Triton `DRAFT_EXTEND` HF-ref mismatch on narrow accept layouts | — (V1 not enabled; Triton uses V2) | ✓ fixed-tokens-per-req | ✓ chain (topk=1) + tree (topk=2) | ✓ via `DRAFT_EXTEND_V2` graph runner | — (production dispatcher only wires Frozen-KV-MTP through FlashInfer-style draft backends) | | `flashinfer` | ✓ MHA/GQA/MQA + 10 input layouts (`head_dim=64` for SM90 prefill constraints) | ✓ MHA/GQA/MQA decode page-boundary | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | ✓ EAGLE chain+tree, Frozen-KV-MTP chain, DFlash chain, NGRAM chain | ✓ EAGLE tree, Frozen-KV-MTP chain, DFlash chain | ✓ EAGLE ragged-accept, Frozen-KV-MTP ragged-accept | ✓ EAGLE ragged-accept, Frozen-KV-MTP ragged-accept | blocked: `is_draft_extend()` default `include_v2=False` → `raise ValueError` (`flashinfer_backend.py:651,748`) | ✓ chain (topk=1) + tree (topk=2) | ✓ EAGLE ragged-accept (V1) | ✓ chain (topk=1) | -| `fa3` | ✓ MHA/GQA/MQA input sweep (FA-friendly `head_dim=64`) | ✓ MHA decode page-boundary | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | deferred: EAGLE tree (topk=2) eager diffs ~0.16 vs the bf16 HF reference — kernel-level drift, not a CG issue | deferred: same kernel-level drift | — | — | deferred: FA's eager `DRAFT_EXTEND_V2` itself diverges by ~0.55 vs HF-ref when `seq_lens=prefix_lens` (the production convention for V2 — the eager `init_forward_metadata` at `flashattention_backend.py:506` sets `cache_seqlens_int32 = seqlens_in_batch` which treats `seq_lens` as full cache length, but for V2 it's prefix only). Triton handles this correctly; needs production-side fix in FA's V2 metadata path. | — | — | — | -| `fa4` | ✓ MHA/GQA/MQA input sweep (FA-friendly `head_dim=64`) | ✓ MHA decode page-boundary | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | deferred: same EAGLE tree eager drift as fa3 | deferred: same | — | — | deferred: same `DRAFT_EXTEND_V2` issue as fa3 | — | — | — | +| `fa3` | ✓ MHA/GQA/MQA input sweep (FA-friendly `head_dim=64`) | ✓ MHA decode page-boundary | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | deferred: EAGLE tree (topk=2) eager diffs ~0.16 vs the bf16 HF reference — kernel-level drift, not a CG issue | deferred: same kernel-level drift | — | — | deferred: `init_forward_metadata_replay_cuda_graph` sets `cache_seqlens = prefix` only; needs `prefix + extend` for DRAFT_EXTEND_V2. Kernel reads stale KV → ~82% wrong values. See KNOWN_FAILURES.md §C.3. | — | — | — | +| `fa4` | ✓ MHA/GQA/MQA input sweep (FA-friendly `head_dim=64`) | ✓ MHA decode page-boundary | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | deferred: same EAGLE tree eager drift as fa3 | deferred: same | — | — | deferred: same `DRAFT_EXTEND_V2` cache-extent issue as fa3 (see KNOWN_FAILURES.md §C.3) | — | — | — | | `flex_attention` | ✓ MHA/GQA/MQA input sweep | blocked: no `init_cuda_graph_state` / capture / replay hooks (`torch_flex_backend.py`) | ✓ MHA ragged, GQA cross-page | ✓ MHA ragged, GQA cross-page | blocked: no CG capture/replay path | blocked: no CG capture/replay path | — | blocked: no CG capture/replay path | blocked: no CG capture/replay path | blocked: no CG capture/replay path | blocked: no CG capture/replay path | blocked: no CG capture/replay path | | `trtllm_mha` | ✓ decode-only MHA/GQA/MQA + page-32 boundary (prefill blocked by `Unsupported architecture`) | deferred: replay mismatches HF-ref on SM90 | — (no extend backend) | — (no extend backend) | blocked: `topk=1` only (`server_args.py:2391-2392`, `trtllm_mha_backend.py:459,492`) | blocked: same `topk=1` constraint | — | — | — | deferred: requires chain-only graph capture wiring | — | — | diff --git a/test/registered/attention/unittest/dense/test_fa3.py b/test/registered/attention/unittest/dense/test_fa3.py index f4d8e6df1..79bb7d3dd 100644 --- a/test/registered/attention/unittest/dense/test_fa3.py +++ b/test/registered/attention/unittest/dense/test_fa3.py @@ -428,6 +428,12 @@ class TestFA3DenseAttentionBackendCorrectness(CustomTestCase): ) def test_runner_mode_eagle_draft_extend_v2_cuda_graph_cases(self): + self.skipTest( + "deferred: FA3 DRAFT_EXTEND_V2 CUDA-graph replay requires " + "`cache_seqlens = prefix + extend` in init_forward_metadata_replay_cuda_graph " + "(flashattention_backend.py); without that fix the kernel reads prefix-only KV " + "and produces ~82 % wrong attention values. Tracked in KNOWN_FAILURES.md §C.3." + ) for case in self.DRAFT_EXTEND_V2_CUDA_GRAPH_CASES: with self.subTest(case=case.name, backend=case.backend): run_dense_draft_extend_v2_cuda_graph_case( diff --git a/test/registered/attention/unittest/dense/test_fa4.py b/test/registered/attention/unittest/dense/test_fa4.py index 4ea21e951..dabacf530 100644 --- a/test/registered/attention/unittest/dense/test_fa4.py +++ b/test/registered/attention/unittest/dense/test_fa4.py @@ -417,6 +417,12 @@ class TestFA4DenseAttentionBackendCorrectness(CustomTestCase): ) def test_runner_mode_eagle_draft_extend_v2_cuda_graph_cases(self): + self.skipTest( + "deferred: FA4 DRAFT_EXTEND_V2 CUDA-graph replay requires " + "`cache_seqlens = prefix + extend` in init_forward_metadata_replay_cuda_graph " + "(flashattention_backend.py); without that fix the kernel reads prefix-only KV " + "and produces ~82 % wrong attention values. Tracked in KNOWN_FAILURES.md §C.3." + ) for case in self.DRAFT_EXTEND_V2_CUDA_GRAPH_CASES: with self.subTest(case=case.name, backend=case.backend): run_dense_draft_extend_v2_cuda_graph_case(