diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 23ccf51a4..ed28fd019 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -864,12 +864,21 @@ class FlashInferAttnBackend(AttentionBackend): ) else: + swa_window_left = ( + layer.sliding_window_size + if not ( + self.forward_metadata.multi_item_params + and self.forward_metadata.multi_item_params.is_enabled() + ) + else -1 + ) o1, s1 = self.prefill_wrapper_ragged.forward_return_lse( q.view(-1, layer.tp_q_head_num, layer.head_dim), k.view(-1, layer.tp_k_head_num, layer.head_dim), v.view(-1, layer.tp_v_head_num, layer.head_dim), causal=causal, sm_scale=layer.scaling, + window_left=swa_window_left, logits_soft_cap=logits_soft_cap, ) o2, s2 = prefill_wrapper_paged.forward_return_lse( @@ -877,6 +886,7 @@ class FlashInferAttnBackend(AttentionBackend): self.token_to_kv_pool.get_kv_buffer(layer.layer_id), causal=False, sm_scale=layer.scaling, + window_left=swa_window_left, logits_soft_cap=logits_soft_cap, ) @@ -1313,19 +1323,34 @@ class FlashInferIndicesUpdaterPrefill: cross_attention_custom_mask: Optional[torch.Tensor] = None, ): for wrapper_id in range(2): + swa_paged_custom_mask = None if wrapper_id == 0: - # window attention use paged only - paged_kernel_lens = torch.minimum( - seq_lens, - torch.tensor(self.sliding_window_size) + seq_lens - prefix_lens, - ) - paged_kernel_lens_sum = paged_kernel_lens.sum().item() + if use_ragged: + # K for extend tokens is written after the paged wrapper runs, so + # the paged wrapper sees prefix-only. Trim to the last `window` tokens + # (required for SWATokenToKVPoolAllocator; also keeps mask O(window)). + effective_start = torch.clamp( + prefix_lens - self.sliding_window_size, min=0 + ) + paged_kernel_lens = prefix_lens - effective_start + paged_kernel_lens_sum = paged_kernel_lens.sum().item() + kv_start_idx = effective_start + swa_paged_custom_mask = self._build_swa_prefix_custom_mask( + prefix_lens, seq_lens, effective_start + ) + else: + # window attention use paged only + paged_kernel_lens = torch.minimum( + seq_lens, + torch.tensor(self.sliding_window_size) + seq_lens - prefix_lens, + ) + paged_kernel_lens_sum = paged_kernel_lens.sum().item() + kv_start_idx = seq_lens - paged_kernel_lens else: # full attention paged_kernel_lens = seq_lens paged_kernel_lens_sum = seq_lens_sum - - kv_start_idx = seq_lens - paged_kernel_lens + kv_start_idx = seq_lens - paged_kernel_lens use_sliding_window_kv_pool = wrapper_id == 0 and isinstance( self.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator ) @@ -1346,8 +1371,50 @@ class FlashInferIndicesUpdaterPrefill: use_sliding_window_kv_pool=use_sliding_window_kv_pool, fixed_split_size=fixed_split_size, multi_item_params=multi_item_params, + cross_attention_custom_mask=swa_paged_custom_mask, ) + def _build_swa_prefix_custom_mask( + self, + prefix_lens: torch.Tensor, + seq_lens: torch.Tensor, + kv_start_idx: torch.Tensor, + ) -> Optional[torch.Tensor]: + """Custom SWA mask for the paged wrapper in the ragged merge_state EXTEND path. + + Paged KV covers absolute positions [kv_start_idx[i], prefix_lens[i]). + Returns None when every key is in-window for every extend query. + """ + window = self.sliding_window_size + if window is None or window < 0: + return None + + prefix_lens_cpu = prefix_lens.detach().cpu().tolist() + extend_lens_cpu = (seq_lens - prefix_lens).detach().cpu().tolist() + kv_start_cpu = kv_start_idx.detach().cpu().tolist() + if all(p == 0 for p in prefix_lens_cpu): + return None + + device = prefix_lens.device + mask_parts: List[torch.Tensor] = [] + need_mask = False + for prefix_len, extend_len, kv_start in zip( + prefix_lens_cpu, extend_lens_cpu, kv_start_cpu + ): + paged_len = int(prefix_len - kv_start) # = min(prefix_len, window) + if paged_len == 0 or extend_len == 0: + continue + q_abs = torch.arange(extend_len, device=device).view(-1, 1) + prefix_len + k_abs = torch.arange(paged_len, device=device).view(1, -1) + kv_start + block = (k_abs >= (q_abs - window)).to(torch.uint8) + if not bool(block.all()): + need_mask = True + mask_parts.append(block.view(-1)) + + if not need_mask or not mask_parts: + return None + return torch.cat(mask_parts) + def update_cross_attention( self, req_pool_indices: torch.Tensor,