From b24c8f10e71b936dbb2a7fbfd1c76cb4afe53c10 Mon Sep 17 00:00:00 2001 From: "lynn.lin" <44727020+hn1209@users.noreply.github.com> Date: Wed, 2 Sep 2026 05:08:22 +0800 Subject: [PATCH] [FlashInfer] Avoid D2H sync for sliding-window lengths (#32218) Co-authored-by: llilian73 <204300658+llilian73@users.noreply.github.com> Co-authored-by: hnyls2002 Co-authored-by: Liangsheng Yin --- .../srt/layers/attention/flashinfer_backend.py | 16 ++++++++++++++-- .../attention/unittests/swa/test_flashinfer.py | 16 ++++++++++++++++ 2 files changed, 30 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 4d89de822..93e360cd0 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -1964,14 +1964,19 @@ class FlashInferIndicesUpdaterPrefill: ) if prefix_lens is None: num_accept_tokens = getattr(spec_info, "num_accept_tokens", None) + # Spec verify keeps its query block outside seq_lens, so an unset + # prefix means the whole sequence is already-cached prefix. + prefix_is_full_seq = num_accept_tokens is None prefix_lens = ( seq_lens - if num_accept_tokens is None + if prefix_is_full_seq else seq_lens - num_accept_tokens[: seq_lens.shape[0]].to( device=seq_lens.device, dtype=seq_lens.dtype ) ) + else: + prefix_is_full_seq = False sliding_window_size = self.sliding_window_size assert sliding_window_size is not None for wrapper_id in range(2): @@ -1997,7 +2002,14 @@ class FlashInferIndicesUpdaterPrefill: seq_lens, sliding_window_size + seq_lens - prefix_lens, ) - paged_kernel_lens_sum = paged_kernel_lens.sum().item() + if prefix_is_full_seq and seq_lens_cpu is not None: + # prefix_lens is seq_lens, so the trim is min(seq_lens, window); + # summing the host mirror avoids draining the stream. + paged_kernel_lens_sum = int( + torch.clamp(seq_lens_cpu, max=sliding_window_size).sum() + ) + else: + paged_kernel_lens_sum = paged_kernel_lens.sum().item() kv_start_idx = seq_lens - paged_kernel_lens else: # full attention diff --git a/test/registered/attention/unittests/swa/test_flashinfer.py b/test/registered/attention/unittests/swa/test_flashinfer.py index 6d4bb602a..3f9b1ccbe 100644 --- a/test/registered/attention/unittests/swa/test_flashinfer.py +++ b/test/registered/attention/unittests/swa/test_flashinfer.py @@ -120,6 +120,22 @@ class TestFlashInferSWAAttentionBackendCorrectness(CustomTestCase): 1, "dflash", ), + ( + DenseAttentionCase( + name="runner_dflash_verify_swa_window_edges", + backend="flashinfer", + forward_mode=ForwardMode.TARGET_VERIFY, + num_heads=4, + num_kv_heads=4, + page_size=16, + # Straddle the window: one request below, one at, one above. + prefix_lens=(1, 4, 9), + extend_lens=(3, 3, 3), + sliding_window_size=4, + ), + 1, + "dflash", + ), ) SPEC_VERIFY_CUDA_GRAPH_CASES = ( (