[SWA] Cache full→SWA out_cache_loc per forward across attention backends (#27617)

This commit is contained in:
Cheng Wan
2026-06-09 22:57:51 -07:00
committed by GitHub
parent 08ceb96ea5
commit 758fd4bb9a
14 changed files with 730 additions and 69 deletions
@@ -263,7 +263,17 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
elif not self.use_mla:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -276,7 +286,16 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
) )
if is_cp_mode: if is_cp_mode:
cp_allgather_and_save_kv_cache( cp_allgather_and_save_kv_cache(
forward_batch, layer, k, v, self.attn_cp_size forward_batch,
layer,
k,
v,
self.attn_cp_size,
swa_loc=(
self.forward_metadata.swa_out_cache_loc
if self.use_sliding_window_kv_pool
else None
),
) )
metadata = self.forward_metadata metadata = self.forward_metadata
@@ -654,7 +673,17 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
elif not self.use_mla:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -24,6 +24,7 @@ from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast
@@ -59,6 +60,9 @@ class ForwardMetadata:
# mapped block_tables for swa # mapped block_tables for swa
block_tables_swa: Optional[torch.Tensor] = None block_tables_swa: Optional[torch.Tensor] = None
# pre-translated full->SWA write target for SWAKVPool.set_kv_buffer
swa_out_cache_loc: Optional[torch.Tensor] = None
# seq len inputs # seq len inputs
extend_seq_lens_cpu_int: Optional[torch.Tensor] = None extend_seq_lens_cpu_int: Optional[torch.Tensor] = None
seq_lens_cpu_int: Optional[torch.Tensor] = None seq_lens_cpu_int: Optional[torch.Tensor] = None
@@ -222,7 +226,7 @@ class AscendAttnMaskBuilder:
def _cp_allgather_and_save_kv_npu( def _cp_allgather_and_save_kv_npu(
forward_batch, layer, k, v, cp_size, token_to_kv_pool forward_batch, layer, k, v, cp_size, token_to_kv_pool, swa_loc=None
): ):
"""NPU-compatible CP KV all-gather with merged K/V communication. """NPU-compatible CP KV all-gather with merged K/V communication.
@@ -234,6 +238,9 @@ def _cp_allgather_and_save_kv_npu(
Equivalent to cp_allgather_and_save_kv_cache() in cp_utils.py, but uses Equivalent to cp_allgather_and_save_kv_cache() in cp_utils.py, but uses
a single all-gather for both K and V. a single all-gather for both K and V.
swa_loc is the pre-translated full->SWA write target for hybrid SWA pools
(None for non-SWA pools); set_kv_buffer never translates internally.
""" """
cache_loc = ( cache_loc = (
forward_batch.out_cache_loc forward_batch.out_cache_loc
@@ -258,12 +265,21 @@ def _cp_allgather_and_save_kv_npu(
key_cache_full = kv_full[..., :k_feat_size].reshape(-1, *k_tail) key_cache_full = kv_full[..., :k_feat_size].reshape(-1, *k_tail)
value_cache_full = kv_full[..., k_feat_size:].reshape(-1, *v_tail) value_cache_full = kv_full[..., k_feat_size:].reshape(-1, *v_tail)
token_to_kv_pool.set_kv_buffer( if swa_loc is not None:
layer, token_to_kv_pool.set_kv_buffer(
cache_loc, layer,
key_cache_full, cache_loc,
value_cache_full, key_cache_full,
) value_cache_full,
swa_loc=swa_loc,
)
else:
token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
key_cache_full,
value_cache_full,
)
class AscendAttnBackend(AttentionBackend): class AscendAttnBackend(AttentionBackend):
@@ -329,6 +345,10 @@ class AscendAttnBackend(AttentionBackend):
self.full_to_swa_index_mapping = ( self.full_to_swa_index_mapping = (
model_runner.token_to_kv_pool.full_to_swa_index_mapping model_runner.token_to_kv_pool.full_to_swa_index_mapping
) )
self.use_sliding_window_kv_pool = (
isinstance(self.token_to_kv_pool, SWAKVPool)
and self.token_to_kv_pool.swa_layer_nums > 0
)
# head num padding # head num padding
self.padding_size_list = [1, 2, 4, 8, 16, 32, 64, 128] self.padding_size_list = [1, 2, 4, 8, 16, 32, 64, 128]
@@ -372,7 +392,10 @@ class AscendAttnBackend(AttentionBackend):
bs = forward_batch.batch_size bs = forward_batch.batch_size
if in_capture: if in_capture:
self._init_cuda_graph_metadata( self._init_cuda_graph_metadata(
bs, forward_batch.forward_mode, forward_batch.seq_lens bs,
forward_batch.forward_mode,
forward_batch.seq_lens,
forward_batch.out_cache_loc,
) )
self._apply_cuda_graph_metadata( self._apply_cuda_graph_metadata(
bs=bs, bs=bs,
@@ -385,6 +408,7 @@ class AscendAttnBackend(AttentionBackend):
), ),
forward_mode=forward_batch.forward_mode, forward_mode=forward_batch.forward_mode,
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
out_cache_loc=forward_batch.out_cache_loc,
) )
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
@@ -474,6 +498,13 @@ class AscendAttnBackend(AttentionBackend):
) )
) )
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
self.forward_metadata.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
self.graph_mode = False self.graph_mode = False
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
@@ -493,18 +524,29 @@ class AscendAttnBackend(AttentionBackend):
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
if self.use_sliding_window_kv_pool:
# refilled in place at replay; the captured graph reads this storage
self.swa_out_cache_loc_buf = torch.zeros(
max_num_tokens,
dtype=torch.int64,
device=self.device,
)
def _init_cuda_graph_metadata( def _init_cuda_graph_metadata(
self, self,
bs: int, bs: int,
forward_mode: ForwardMode, forward_mode: ForwardMode,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
out_cache_loc: Optional[torch.Tensor] = None,
) -> "ForwardMetadata": ) -> "ForwardMetadata":
"""Create and store the per-bs ForwardMetadata for CUDA graph capture.""" """Create and store the per-bs ForwardMetadata for CUDA graph capture."""
metadata = ForwardMetadata() metadata = ForwardMetadata()
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :] metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
if self.is_hybrid_swa: if self.is_hybrid_swa:
metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :] metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :]
if self.use_sliding_window_kv_pool and out_cache_loc is not None:
num_tokens = out_cache_loc.shape[0]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist() metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
metadata.seq_lens = seq_lens metadata.seq_lens = seq_lens
if ( if (
@@ -571,12 +613,21 @@ class AscendAttnBackend(AttentionBackend):
seq_lens_cpu: torch.Tensor, seq_lens_cpu: torch.Tensor,
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
out_cache_loc: Optional[torch.Tensor] = None,
): ):
"""Shared capture+replay body for the cuda-graph init path. """Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`. Public entry: :py:meth:`init_forward_metadata_out_graph`.
""" """
metadata = self.graph_metadata[bs] metadata = self.graph_metadata[bs]
# refill the captured SWA write-target buffer in place from the live loc
if self.use_sliding_window_kv_pool and out_cache_loc is not None:
n = out_cache_loc.shape[0]
self.swa_out_cache_loc_buf[n:].zero_()
self.swa_out_cache_loc_buf[:n].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(out_cache_loc)
)
max_len = seq_lens_cpu[:bs].max().item() max_len = seq_lens_cpu[:bs].max().item()
if forward_mode.is_target_verify(): if forward_mode.is_target_verify():
max_len += self.speculative_num_draft_tokens max_len += self.speculative_num_draft_tokens
@@ -1068,6 +1119,11 @@ class AscendAttnBackend(AttentionBackend):
v, v,
self.attn_cp_size, self.attn_cp_size,
self.token_to_kv_pool, self.token_to_kv_pool,
swa_loc=(
self.forward_metadata.swa_out_cache_loc
if self.use_sliding_window_kv_pool
else None
),
) )
else: else:
# support cross attention # support cross attention
@@ -1076,7 +1132,16 @@ class AscendAttnBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) if self.use_sliding_window_kv_pool and not layer.is_cross_attention:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
@@ -1587,9 +1652,18 @@ class AscendAttnBackend(AttentionBackend):
topk_indices: Optional[torch.Tensor] = None, topk_indices: Optional[torch.Tensor] = None,
): ):
if save_kv_cache: if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, forward_batch.out_cache_loc, k, v self.token_to_kv_pool.set_kv_buffer(
) layer,
forward_batch.out_cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
)
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
@@ -1651,6 +1725,14 @@ class AscendAttnBackend(AttentionBackend):
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
elif self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else: else:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
@@ -1862,6 +1944,14 @@ class AscendAttnBackend(AttentionBackend):
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
elif self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else: else:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
@@ -2195,7 +2285,16 @@ class AscendAttnBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) if self.use_sliding_window_kv_pool and not layer.is_cross_attention:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
num_tokens = q.shape[0] num_tokens = q.shape[0]
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
@@ -2487,9 +2586,18 @@ class AscendAttnBackend(AttentionBackend):
"3. When the environment variable ASCEND_USE_FIA is set to 0 and qk_head_dim exceeds 128 on Ascend NPU devices." "3. When the environment variable ASCEND_USE_FIA is set to 0 and qk_head_dim exceeds 128 on Ascend NPU devices."
) )
if save_kv_cache: if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, forward_batch.out_cache_loc, k, v self.token_to_kv_pool.set_kv_buffer(
) layer,
forward_batch.out_cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
)
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
num_block, block_size, _, _ = k_cache.shape num_block, block_size, _, _ = k_cache.shape
@@ -117,6 +117,8 @@ class ForwardMetadata:
max_extend_len: Optional[int] = None max_extend_len: Optional[int] = None
fp8_prefill_kv_indices: Optional[torch.Tensor] = None fp8_prefill_kv_indices: Optional[torch.Tensor] = None
swa_page_table: Optional[torch.Tensor] = None swa_page_table: Optional[torch.Tensor] = None
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: Optional[torch.Tensor] = None
global_workspace_buffer = None global_workspace_buffer = None
@@ -865,6 +867,20 @@ class AiterAttnBackend(AttentionBackend):
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
) )
# Refill the SWA write-target buffer from the live out_cache_loc and
# bind it onto the metadata before replay (_apply rebuilds it each call).
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
n = forward_batch.out_cache_loc.shape[0]
self.cuda_graph_swa_out_cache_loc[n:].zero_()
self.cuda_graph_swa_out_cache_loc[:n].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
self.forward_metadata.swa_out_cache_loc = self.cuda_graph_swa_out_cache_loc[
:n
]
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for aiter attention backend.""" """Init auxiliary variables for aiter attention backend."""
@@ -885,6 +901,11 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits = None num_kv_splits = None
swa_page_table = None swa_page_table = None
swa_out_cache_loc = None
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
swa_out_cache_loc = self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
max_kv_len = forward_batch.seq_lens_cpu.max().item() max_kv_len = forward_batch.seq_lens_cpu.max().item()
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
@@ -1005,6 +1026,7 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits=num_kv_splits, num_kv_splits=num_kv_splits,
run_graph=False, run_graph=False,
swa_page_table=swa_page_table, swa_page_table=swa_page_table,
swa_out_cache_loc=swa_out_cache_loc,
) )
elif forward_batch.forward_mode.is_draft_extend_v2(): elif forward_batch.forward_mode.is_draft_extend_v2():
@@ -1277,6 +1299,7 @@ class AiterAttnBackend(AttentionBackend):
max_kv_len, max_kv_len,
max_extend_len=max_q_len, max_extend_len=max_q_len,
swa_page_table=swa_page_table, swa_page_table=swa_page_table,
swa_out_cache_loc=swa_out_cache_loc,
) )
else: else:
qo_indptr = torch.arange( qo_indptr = torch.arange(
@@ -1425,6 +1448,7 @@ class AiterAttnBackend(AttentionBackend):
max(forward_batch.extend_seq_lens_cpu), max(forward_batch.extend_seq_lens_cpu),
forward_batch.seq_lens_cpu.max().item(), forward_batch.seq_lens_cpu.max().item(),
swa_page_table=swa_page_table, swa_page_table=swa_page_table,
swa_out_cache_loc=swa_out_cache_loc,
) )
def init_cuda_graph_state( def init_cuda_graph_state(
@@ -1539,6 +1563,13 @@ class AiterAttnBackend(AttentionBackend):
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
# SWA write-target buffer; refilled and bound onto forward_metadata
# in init_forward_metadata_out_graph before each replay.
self.cuda_graph_swa_out_cache_loc = torch.zeros(
(max_num_tokens,),
dtype=torch.int64,
device=self.device,
)
def _apply_cuda_graph_metadata( def _apply_cuda_graph_metadata(
self, self,
@@ -2020,9 +2051,20 @@ class AiterAttnBackend(AttentionBackend):
# launch_reshape_and_cache_flash; always route through # launch_reshape_and_cache_flash; always route through
# set_kv_buffer which dispatches to the SHUFFLE 5D writer. # set_kv_buffer which dispatches to the SHUFFLE 5D writer.
if self.kv_cache_is_vectorized_5d: if self.kv_cache_is_vectorized_5d:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, k_descale, v_descale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
k_descale,
v_descale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, k_descale, v_descale
)
# Only use SWA-specific kv cache write (reshape_and_cache_flash) when # Only use SWA-specific kv cache write (reshape_and_cache_flash) when
# both unified attention and sliding window kv pool are active. # both unified attention and sliding window kv pool are active.
# Non-SWA models (e.g. Qwen3-VL) enabled via SGLANG_USE_AITER_UNIFIED_ATTN # Non-SWA models (e.g. Qwen3-VL) enabled via SGLANG_USE_AITER_UNIFIED_ATTN
@@ -2059,9 +2101,20 @@ class AiterAttnBackend(AttentionBackend):
elif self.use_mla: elif self.use_mla:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
else: else:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, k_descale, v_descale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
k_descale,
v_descale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, k_descale, v_descale
)
if self.use_mla: if self.use_mla:
max_q_len = self.forward_metadata.max_q_len max_q_len = self.forward_metadata.max_q_len
@@ -2487,9 +2540,20 @@ class AiterAttnBackend(AttentionBackend):
if save_kv_cache: if save_kv_cache:
# SHUFFLE 5D pool path — see forward_extend for rationale. # SHUFFLE 5D pool path — see forward_extend for rationale.
if self.kv_cache_is_vectorized_5d: if self.kv_cache_is_vectorized_5d:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, forward_batch.out_cache_loc, k, v, k_descale, v_descale self.token_to_kv_pool.set_kv_buffer(
) layer,
forward_batch.out_cache_loc,
k,
v,
k_descale,
v_descale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v, k_descale, v_descale
)
# Only use SWA-specific kv cache write (reshape_and_cache_flash) when # Only use SWA-specific kv cache write (reshape_and_cache_flash) when
# both unified attention and sliding window kv pool are active. # both unified attention and sliding window kv pool are active.
# Non-SWA models (e.g. Qwen3-VL) enabled via SGLANG_USE_AITER_UNIFIED_ATTN # Non-SWA models (e.g. Qwen3-VL) enabled via SGLANG_USE_AITER_UNIFIED_ATTN
@@ -2531,6 +2595,14 @@ class AiterAttnBackend(AttentionBackend):
), ),
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
) )
elif self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else: else:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
@@ -61,6 +61,8 @@ class FlashAttentionMetadata:
page_table: torch.Tensor = None page_table: torch.Tensor = None
# Page table for Sliding Window Attention # Page table for Sliding Window Attention
swa_page_table: torch.Tensor = None swa_page_table: torch.Tensor = None
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: torch.Tensor = None
# Precomputed FA3 scheduler metadata (avoids per-layer prepare_varlen_num_blocks) # Precomputed FA3 scheduler metadata (avoids per-layer prepare_varlen_num_blocks)
scheduler_metadata: torch.Tensor = None scheduler_metadata: torch.Tensor = None
@@ -698,6 +700,12 @@ class FlashAttentionBackend(AttentionBackend):
metadata.page_table metadata.page_table
).to(torch.int32) ).to(torch.int32)
) )
if forward_batch.out_cache_loc is not None:
metadata.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
# Convert the page table to a strided format which is needed by FA3 API # Convert the page table to a strided format which is needed by FA3 API
if self.page_size > 1: if self.page_size > 1:
@@ -798,7 +806,26 @@ class FlashAttentionBackend(AttentionBackend):
# Dense-MHA CP: k, v are still rank-local; backend # Dense-MHA CP: k, v are still rank-local; backend
# all-gathers and writes to the per-rank pool. # all-gathers and writes to the per-rank pool.
cp_allgather_and_save_kv_cache( cp_allgather_and_save_kv_cache(
forward_batch, layer, k, v, self.attn_cp_size forward_batch,
layer,
k,
v,
self.attn_cp_size,
swa_loc=(
self.forward_metadata.swa_out_cache_loc
if self.use_sliding_window_kv_pool
else None
),
)
elif self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
) )
else: else:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
@@ -1242,7 +1269,17 @@ class FlashAttentionBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
elif not self.use_mla:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -1588,6 +1625,13 @@ class FlashAttentionBackend(AttentionBackend):
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
# SWA write-target buffer; metadata binds a [:num_tokens] view,
# refilled from the live out_cache_loc before each replay.
self.swa_out_cache_loc_buf = torch.zeros(
max_num_tokens,
dtype=torch.int64,
device=self.device,
)
# This is used by draft decode's first half of metadata when topk > 1 # This is used by draft decode's first half of metadata when topk > 1
if self.topk > 1: if self.topk > 1:
@@ -1849,6 +1893,9 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.decode_cuda_graph_metadata[ metadata.swa_page_table = self.decode_cuda_graph_metadata[
"swa_page_table" "swa_page_table"
][:bs, :] ][:bs, :]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[
:num_tokens
]
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
else: else:
# Draft Decode topk>1: two metadata objects # Draft Decode topk>1: two metadata objects
@@ -1905,6 +1952,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.decode_cuda_graph_metadata[ metadata.swa_page_table = self.decode_cuda_graph_metadata[
"swa_page_table" "swa_page_table"
][:bs, :] ][:bs, :]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
@@ -1924,6 +1972,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.target_verify_metadata[ metadata.swa_page_table = self.target_verify_metadata[
"swa_page_table" "swa_page_table"
][:bs, :] ][:bs, :]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
self.target_verify_metadata[bs] = metadata self.target_verify_metadata[bs] = metadata
else: else:
# Target Verify topk>1: two (or three with SWA) metadata objects # Target Verify topk>1: two (or three with SWA) metadata objects
@@ -1959,6 +2008,10 @@ class FlashAttentionBackend(AttentionBackend):
self.target_verify_metadata_topk_normal[bs] = metadata self.target_verify_metadata_topk_normal[bs] = metadata
self.target_verify_metadata_topk_expand[bs] = metadata_expand self.target_verify_metadata_topk_expand[bs] = metadata_expand
# topk>1 target-verify early-returns before _apply; bind the
# view here (buffer refilled at replay).
if self.use_sliding_window_kv_pool:
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
if self.has_swa: if self.has_swa:
metadata_swa = FlashAttentionMetadata() metadata_swa = FlashAttentionMetadata()
@@ -1995,6 +2048,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][ metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][
:bs, : :bs, :
] ]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
self.draft_extend_metadata[bs] = metadata self.draft_extend_metadata[bs] = metadata
if encoder_lens is not None: if encoder_lens is not None:
@@ -2037,6 +2091,15 @@ class FlashAttentionBackend(AttentionBackend):
metadata = None metadata = None
metadata_expand = None metadata_expand = None
# Refill the SWA write-target buffer (bound as a metadata view in
# _bind_metadata_buffers) from the live out_cache_loc before replay.
if self.use_sliding_window_kv_pool and out_cache_loc is not None:
n = out_cache_loc.shape[0]
self.swa_out_cache_loc_buf[n:].zero_()
self.swa_out_cache_loc_buf[:n].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(out_cache_loc)
)
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
if spec_info is not None: if spec_info is not None:
# Draft Decode # Draft Decode
@@ -27,6 +27,7 @@ from sglang.srt.layers.attention.utils import (
from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.cuda_graph_config import (
Backend, Backend,
Phase, Phase,
@@ -138,6 +139,8 @@ class MultiItemScoringParams:
@dataclass @dataclass
class DecodeMetadata: class DecodeMetadata:
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper] decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper]
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: Optional[torch.Tensor] = None
@dataclass @dataclass
@@ -146,6 +149,7 @@ class PrefillMetadata:
use_ragged: bool use_ragged: bool
extend_no_prefix: bool extend_no_prefix: bool
multi_item_params: Optional[MultiItemScoringParams] = None multi_item_params: Optional[MultiItemScoringParams] = None
swa_out_cache_loc: Optional[torch.Tensor] = None
# Reuse this workspace buffer across all flashinfer wrappers # Reuse this workspace buffer across all flashinfer wrappers
@@ -173,6 +177,7 @@ class FlashInferAttnBackend(AttentionBackend):
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.use_sliding_window_kv_pool = isinstance(self.token_to_kv_pool, SWAKVPool)
self.enable_mis = model_runner.server_args.enable_mis self.enable_mis = model_runner.server_args.enable_mis
# FIXME: remove dllm workarounds from flashinfer # FIXME: remove dllm workarounds from flashinfer
@@ -541,7 +546,28 @@ class FlashInferAttnBackend(AttentionBackend):
for w in self.decode_cuda_graph_metadata[bs]: for w in self.decode_cuda_graph_metadata[bs]:
w.begin_forward = partial(fast_decode_plan, w) w.begin_forward = partial(fast_decode_plan, w)
# Refill the SWA write-target buffer from the live out_cache_loc before
# replay (bound onto the metadata at capture below).
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
n = forward_batch.out_cache_loc.shape[0]
self.cuda_graph_swa_out_cache_loc[n:].zero_()
self.cuda_graph_swa_out_cache_loc[:n].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
if in_capture:
self.forward_metadata.swa_out_cache_loc = (
self.cuda_graph_swa_out_cache_loc[:n]
)
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
swa_out_cache_loc = None
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
swa_out_cache_loc = self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
self.indices_updater_decode.update( self.indices_updater_decode.update(
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
@@ -554,7 +580,9 @@ class FlashInferAttnBackend(AttentionBackend):
fixed_split_size=self.decode_split_tile_size, fixed_split_size=self.decode_split_tile_size,
disable_split_kv=False, disable_split_kv=False,
) )
self.forward_metadata = DecodeMetadata(self.decode_wrappers) self.forward_metadata = DecodeMetadata(
self.decode_wrappers, swa_out_cache_loc=swa_out_cache_loc
)
elif forward_batch.forward_mode.is_draft_extend(): elif forward_batch.forward_mode.is_draft_extend():
self.indices_updater_prefill.update( self.indices_updater_prefill.update(
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
@@ -568,7 +596,10 @@ class FlashInferAttnBackend(AttentionBackend):
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
) )
self.forward_metadata = PrefillMetadata( self.forward_metadata = PrefillMetadata(
self.prefill_wrappers_paged, False, False self.prefill_wrappers_paged,
False,
False,
swa_out_cache_loc=swa_out_cache_loc,
) )
elif forward_batch.forward_mode.is_target_verify(): elif forward_batch.forward_mode.is_target_verify():
self.indices_updater_prefill.update( self.indices_updater_prefill.update(
@@ -583,7 +614,10 @@ class FlashInferAttnBackend(AttentionBackend):
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
) )
self.forward_metadata = PrefillMetadata( self.forward_metadata = PrefillMetadata(
self.prefill_wrappers_verify, False, False self.prefill_wrappers_verify,
False,
False,
swa_out_cache_loc=swa_out_cache_loc,
) )
else: else:
prefix_lens = forward_batch.extend_prefix_lens prefix_lens = forward_batch.extend_prefix_lens
@@ -631,6 +665,7 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged, use_ragged,
extend_no_prefix, extend_no_prefix,
multi_item_params, multi_item_params,
swa_out_cache_loc=swa_out_cache_loc,
) )
def init_cuda_graph_state( def init_cuda_graph_state(
@@ -652,6 +687,14 @@ class FlashInferAttnBackend(AttentionBackend):
cuda_graph_kv_indices.clone() for _ in range(self.num_wrappers - 1) cuda_graph_kv_indices.clone() for _ in range(self.num_wrappers - 1)
] ]
# SWA write-target buffer; refilled and bound onto forward_metadata in
# init_forward_metadata_out_graph before each replay.
self.cuda_graph_swa_out_cache_loc = (
torch.zeros(max_num_tokens, dtype=torch.int64, device="cuda")
if self.use_sliding_window_kv_pool
else None
)
# Ensure tensors are properly allocated # Ensure tensors are properly allocated
for i in range(self.num_wrappers): for i in range(self.num_wrappers):
# Force allocation by performing a small operation # Force allocation by performing a small operation
@@ -771,9 +814,20 @@ class FlashInferAttnBackend(AttentionBackend):
if k is not None: if k is not None:
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, layer.k_scale, layer.v_scale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
causal = ( causal = (
not layer.is_cross_attention not layer.is_cross_attention
@@ -863,9 +917,20 @@ class FlashInferAttnBackend(AttentionBackend):
o, _ = _safe_merge_state(o1, s1, o2, s2) o, _ = _safe_merge_state(o1, s1, o2, s2)
if save_kv_cache: if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, layer.k_scale, layer.v_scale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
return o.view(-1, layer.tp_q_head_num * layer.head_dim) return o.view(-1, layer.tp_q_head_num * layer.head_dim)
@@ -891,9 +956,20 @@ class FlashInferAttnBackend(AttentionBackend):
if k is not None: if k is not None:
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, layer.k_scale, layer.v_scale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
# Call the wrapped function # Call the wrapped function
o = decode_wrapper.forward( o = decode_wrapper.forward(
@@ -5,6 +5,7 @@ from typing import TYPE_CHECKING
import torch import torch
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -24,6 +25,14 @@ class IntelAMXAttnBackend(AttentionBackend):
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
# full->SWA translated out_cache_loc, computed once per forward (the only
# set_kv_buffer is in eager forward_extend; decode writes KV in-kernel).
self.use_sliding_window_kv_pool = (
isinstance(self.token_to_kv_pool, SWAKVPool)
and self.token_to_kv_pool.swa_layer_nums > 0
)
self.swa_out_cache_loc = None
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // model_runner.tp_size model_runner.model_config.num_attention_heads // model_runner.tp_size
) )
@@ -61,6 +70,15 @@ class IntelAMXAttnBackend(AttentionBackend):
max_extend_len = torch.max(forward_batch.extend_seq_lens).item() max_extend_len = torch.max(forward_batch.extend_seq_lens).item()
self.forward_metadata = (attn_logits, max_extend_len) self.forward_metadata = (attn_logits, max_extend_len)
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
self.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
else:
self.swa_out_cache_loc = None
def get_cpu_graph_seq_len_fill_value(self): def get_cpu_graph_seq_len_fill_value(self):
return 1 return 1
@@ -110,7 +128,12 @@ class IntelAMXAttnBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if save_kv_cache and k is not None and v is not None: if save_kv_cache and k is not None and v is not None:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) if self.use_sliding_window_kv_pool and not layer.is_cross_attention:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, swa_loc=self.swa_out_cache_loc
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
_, max_extend_len = self.forward_metadata _, max_extend_len = self.forward_metadata
self.extend_attention_fwd( self.extend_attention_fwd(
@@ -7,6 +7,7 @@ from torch.nn.functional import scaled_dot_product_attention
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -23,6 +24,12 @@ class TorchNativeAttnBackend(AttentionBackend):
# corresponding ForwardBatch fields. # corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.use_sliding_window_kv_pool = (
isinstance(self.token_to_kv_pool, SWAKVPool)
and self.token_to_kv_pool.swa_layer_nums > 0
)
# full->SWA translated out_cache_loc, computed once per forward
self.swa_out_cache_loc = None
@staticmethod @staticmethod
def _make_sliding_window_mask( def _make_sliding_window_mask(
@@ -41,7 +48,14 @@ class TorchNativeAttnBackend(AttentionBackend):
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass.""" """Init the metadata for a forward pass."""
pass if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
self.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
else:
self.swa_out_cache_loc = None
def _run_sdpa_forward_extend( def _run_sdpa_forward_extend(
self, self,
@@ -281,7 +295,12 @@ class TorchNativeAttnBackend(AttentionBackend):
cache_loc = forward_batch.out_cache_loc cache_loc = forward_batch.out_cache_loc
if save_kv_cache and k is not None and v is not None: if save_kv_cache and k is not None and v is not None:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, swa_loc=self.swa_out_cache_loc
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
use_gqa = layer.tp_q_head_num != layer.tp_k_head_num use_gqa = layer.tp_q_head_num != layer.tp_k_head_num
@@ -347,7 +366,12 @@ class TorchNativeAttnBackend(AttentionBackend):
cache_loc = forward_batch.out_cache_loc cache_loc = forward_batch.out_cache_loc
if save_kv_cache and k is not None and v is not None: if save_kv_cache and k is not None and v is not None:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, swa_loc=self.swa_out_cache_loc
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
use_gqa = layer.tp_q_head_num != layer.tp_k_head_num use_gqa = layer.tp_q_head_num != layer.tp_k_head_num
@@ -82,6 +82,8 @@ class ForwardMetadata:
window_kv_offsets: torch.Tensor window_kv_offsets: torch.Tensor
# Separate attn_logits for SWA layers when v_head_dim differs # Separate attn_logits for SWA layers when v_head_dim differs
swa_attn_logits: Optional[torch.Tensor] = None swa_attn_logits: Optional[torch.Tensor] = None
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: Optional[torch.Tensor] = None
class TritonAttnBackend(AttentionBackend): class TritonAttnBackend(AttentionBackend):
@@ -124,6 +126,7 @@ class TritonAttnBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator self.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator
self.use_sliding_window_kv_pool = isinstance(self.token_to_kv_pool, SWAKVPool)
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
@@ -521,8 +524,9 @@ class TritonAttnBackend(AttentionBackend):
forward_mode=forward_mode, forward_mode=forward_mode,
spec_info=spec_info, spec_info=spec_info,
) )
swa_out_cache_loc = self._fill_cuda_graph_swa_out_cache_loc(forward_batch)
self.forward_metadata = self._build_cuda_graph_forward_metadata( self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info bs, forward_mode, spec_info, swa_out_cache_loc
) )
else: else:
self._apply_cuda_graph_metadata( self._apply_cuda_graph_metadata(
@@ -532,6 +536,29 @@ class TritonAttnBackend(AttentionBackend):
forward_mode=forward_mode, forward_mode=forward_mode,
spec_info=spec_info, spec_info=spec_info,
) )
# Metadata view is reused from capture; just refill the buffer.
self._fill_cuda_graph_swa_out_cache_loc(forward_batch)
def _fill_cuda_graph_swa_out_cache_loc(
self, forward_batch: ForwardBatch
) -> Optional[torch.Tensor]:
"""Refill the SWA write-target buffer from the live out_cache_loc and
return the [:n] view (None for non-SWA / multi-step draft), so the
captured store reads fresh slots on replay."""
if not self.use_sliding_window_kv_pool:
return None
out_cache_loc = forward_batch.out_cache_loc
if (
out_cache_loc is None
or out_cache_loc.shape[0] > self.cuda_graph_swa_out_cache_loc.shape[0]
):
return None
n = out_cache_loc.shape[0]
self.cuda_graph_swa_out_cache_loc[n:].zero_()
self.cuda_graph_swa_out_cache_loc[:n].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(out_cache_loc)
)
return self.cuda_graph_swa_out_cache_loc[:n]
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend.""" """Init auxiliary variables for triton attention backend."""
@@ -742,6 +769,12 @@ class TritonAttnBackend(AttentionBackend):
max_extend_len = int(forward_batch.extend_seq_lens.max()) max_extend_len = int(forward_batch.extend_seq_lens.max())
num_kv_splits = None num_kv_splits = None
swa_out_cache_loc = None
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
swa_out_cache_loc = self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
self.forward_metadata = ForwardMetadata( self.forward_metadata = ForwardMetadata(
attn_logits, attn_logits,
attn_lse, attn_lse,
@@ -757,6 +790,7 @@ class TritonAttnBackend(AttentionBackend):
window_num_kv_splits, window_num_kv_splits,
window_kv_offsets, window_kv_offsets,
swa_attn_logits=swa_attn_logits, swa_attn_logits=swa_attn_logits,
swa_out_cache_loc=swa_out_cache_loc,
) )
def init_cuda_graph_state( def init_cuda_graph_state(
@@ -839,11 +873,20 @@ class TritonAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
if self.use_sliding_window_kv_pool:
# SWA write-target buffer; refilled at replay from out_cache_loc.
self.cuda_graph_swa_out_cache_loc = torch.zeros(
(max_num_tokens,),
dtype=torch.int64,
device=self.device,
)
def _build_cuda_graph_forward_metadata( def _build_cuda_graph_forward_metadata(
self, self,
bs: int, bs: int,
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
swa_out_cache_loc: Optional[torch.Tensor] = None,
) -> ForwardMetadata: ) -> ForwardMetadata:
"""Construct ForwardMetadata from the current cuda-graph buffer state. """Construct ForwardMetadata from the current cuda-graph buffer state.
@@ -851,7 +894,8 @@ class TritonAttnBackend(AttentionBackend):
(either via replay or directly). All fields reference the same (either via replay or directly). All fields reference the same
self.cuda_graph_* tensors that the captured graph kernels will self.cuda_graph_* tensors that the captured graph kernels will
read — the Python object is rebuilt each capture, but the underlying read — the Python object is rebuilt each capture, but the underlying
GPU memory addresses are stable. GPU memory addresses are stable. ``swa_out_cache_loc`` is the
pre-allocated SWA write-target buffer view (or None for non-SWA).
""" """
swa = self.sliding_window_size is not None and self.sliding_window_size > 0 swa = self.sliding_window_size is not None and self.sliding_window_size > 0
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
@@ -872,6 +916,7 @@ class TritonAttnBackend(AttentionBackend):
), ),
window_kv_offsets=None, window_kv_offsets=None,
swa_attn_logits=self.cuda_graph_swa_attn_logits, swa_attn_logits=self.cuda_graph_swa_attn_logits,
swa_out_cache_loc=swa_out_cache_loc,
) )
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
custom_mask = ( custom_mask = (
@@ -896,6 +941,7 @@ class TritonAttnBackend(AttentionBackend):
self.cuda_graph_window_num_kv_splits if swa else None self.cuda_graph_window_num_kv_splits if swa else None
), ),
window_kv_offsets=self.cuda_graph_window_kv_offsets if swa else None, window_kv_offsets=self.cuda_graph_window_kv_offsets if swa else None,
swa_out_cache_loc=swa_out_cache_loc,
) )
elif forward_mode.is_draft_extend(include_v2=True): elif forward_mode.is_draft_extend(include_v2=True):
return ForwardMetadata( return ForwardMetadata(
@@ -920,6 +966,7 @@ class TritonAttnBackend(AttentionBackend):
window_kv_indices=None, window_kv_indices=None,
window_num_kv_splits=None, window_num_kv_splits=None,
window_kv_offsets=None, window_kv_offsets=None,
swa_out_cache_loc=swa_out_cache_loc,
) )
else: else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
@@ -1006,7 +1053,27 @@ class TritonAttnBackend(AttentionBackend):
else: else:
# Save KV cache first (must do this before unified kernel) # Save KV cache first (must do this before unified kernel)
if save_kv_cache: if save_kv_cache:
if layer.k_scale is None: if self.use_sliding_window_kv_pool:
# SWA pool (never MLA); clone k,v when scaling, as below.
if layer.k_scale is None:
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
k,
v,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
k.clone(),
v.clone(),
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
elif layer.k_scale is None:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
@@ -1270,6 +1337,16 @@ class TritonAttnBackend(AttentionBackend):
k, k,
v, v,
) )
elif self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
forward_batch.out_cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else: else:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
@@ -60,6 +60,8 @@ class TRTLLMMHAMetadata:
page_table: torch.Tensor = None page_table: torch.Tensor = None
# Page table for SWA layers (translated from full pool indices to SWA pool indices) # Page table for SWA layers (translated from full pool indices to SWA pool indices)
swa_page_table: torch.Tensor = None swa_page_table: torch.Tensor = None
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: torch.Tensor = None
class TRTLLMHAAttnBackend(FlashInferAttnBackend): class TRTLLMHAAttnBackend(FlashInferAttnBackend):
@@ -252,6 +254,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
), ),
} }
# SWA write-target buffer; bound as a [:num_tokens] view in
# _build_cuda_graph_metadata, refilled before each replay in
# init_forward_metadata_out_graph.
self.cuda_graph_swa_out_cache_loc = (
torch.zeros(max_num_tokens, dtype=torch.int64, device=self.device)
if self.use_sliding_window_kv_pool
else None
)
if ( if (
self.speculative_num_draft_tokens is not None self.speculative_num_draft_tokens is not None
and self.speculative_num_draft_tokens > 0 and self.speculative_num_draft_tokens > 0
@@ -414,6 +425,10 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
) )
self.draft_extend_metadata[bs] = metadata self.draft_extend_metadata[bs] = metadata
# Bind the SWA write-target buffer slice (refilled at replay).
if self.use_sliding_window_kv_pool:
metadata.swa_out_cache_loc = self.cuda_graph_swa_out_cache_loc[:num_tokens]
return metadata return metadata
def _apply_cuda_graph_metadata( def _apply_cuda_graph_metadata(
@@ -620,6 +635,17 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
seq_lens_cpu=forward_batch.seq_lens_cpu, seq_lens_cpu=forward_batch.seq_lens_cpu,
) )
# Refill the SWA write-target buffer from the live out_cache_loc before
# replay (the per-bs metadata holds a view bound in _build).
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
n = forward_batch.out_cache_loc.shape[0]
self.cuda_graph_swa_out_cache_loc[n:].zero_()
self.cuda_graph_swa_out_cache_loc[:n].copy_(
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize the metadata for a forward pass.""" """Initialize the metadata for a forward pass."""
@@ -718,6 +744,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# Compute SWA page table (None for non-SWA models) # Compute SWA page table (None for non-SWA models)
metadata.swa_page_table = self._maybe_translate_swa(metadata.page_table) metadata.swa_page_table = self._maybe_translate_swa(metadata.page_table)
# int64 scatter index (unlike the int32 read page table above).
if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None:
metadata.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
# Convert the page tables to a strided format # Convert the page tables to a strided format
if self.page_size > 1: if self.page_size > 1:
self.strided_indices = torch.arange( self.strided_indices = torch.arange(
@@ -762,9 +796,20 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
else: else:
# Use original set_kv_buffer path # Use original set_kv_buffer path
if save_kv_cache and k is not None: if save_kv_cache and k is not None:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, layer.k_scale, layer.v_scale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
# For XQA, q_dtype should be bf16 # For XQA, q_dtype should be bf16
if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl): if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl):
@@ -848,9 +893,20 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
else: else:
# Use original set_kv_buffer path # Use original set_kv_buffer path
if save_kv_cache and k is not None: if save_kv_cache and k is not None:
self.token_to_kv_pool.set_kv_buffer( if self.use_sliding_window_kv_pool:
layer, cache_loc, k, v, layer.k_scale, layer.v_scale self.token_to_kv_pool.set_kv_buffer(
) layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
if self.data_type == torch.float8_e4m3fn: if self.data_type == torch.float8_e4m3fn:
q = q.to(torch.float8_e4m3fn) q = q.to(torch.float8_e4m3fn)
@@ -390,6 +390,12 @@ class XPUAttentionBackend(AttentionBackend):
metadata.page_table metadata.page_table
).to(torch.int32) ).to(torch.int32)
) )
if forward_batch.out_cache_loc is not None:
metadata.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
if self.use_mla: if self.use_mla:
workspace_size = flash_mla_get_workspace_size( workspace_size = flash_mla_get_workspace_size(
@@ -416,6 +422,12 @@ class XPUAttentionBackend(AttentionBackend):
metadata.page_table metadata.page_table
).to(torch.int32) ).to(torch.int32)
) )
if forward_batch.out_cache_loc is not None:
metadata.swa_out_cache_loc = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
# Convert the page table to a strided format which is needed by FA3 API # Convert the page table to a strided format which is needed by FA3 API
if self.page_size > 1: if self.page_size > 1:
@@ -464,7 +476,17 @@ class XPUAttentionBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
elif not self.use_mla:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -766,7 +788,17 @@ class XPUAttentionBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
k,
v,
layer.k_scale,
layer.v_scale,
swa_loc=self.forward_metadata.swa_out_cache_loc,
)
elif not self.use_mla:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
+23 -9
View File
@@ -415,10 +415,12 @@ def cp_all_gather_rerange_kv_cache(input_tensor, cp_size, forward_batch, stream)
return output_tensor return output_tensor
def cp_allgather_and_save_kv_cache(forward_batch, layer, k, v, cp_size): def cp_allgather_and_save_kv_cache(forward_batch, layer, k, v, cp_size, swa_loc=None):
""" """
Allgather KV cache from all CP ranks and write the full result Allgather KV cache from all CP ranks and write the full result
into each rank's local memory pool. into each rank's local memory pool.
swa_loc is the pre-translated full->SWA write target for hybrid SWA pools.
""" """
cache_loc = ( cache_loc = (
forward_batch.out_cache_loc forward_batch.out_cache_loc
@@ -436,14 +438,26 @@ def cp_allgather_and_save_kv_cache(forward_batch, layer, k, v, cp_size):
v, cp_size, forward_batch, torch.cuda.current_stream() v, cp_size, forward_batch, torch.cuda.current_stream()
) )
get_token_to_kv_pool().set_kv_buffer( pool = get_token_to_kv_pool()
layer, if swa_loc is not None:
cache_loc, pool.set_kv_buffer(
key_cache_full, layer,
value_cache_full, cache_loc,
layer.k_scale, key_cache_full,
layer.v_scale, value_cache_full,
) layer.k_scale,
layer.v_scale,
swa_loc=swa_loc,
)
else:
pool.set_kv_buffer(
layer,
cache_loc,
key_cache_full,
value_cache_full,
layer.k_scale,
layer.v_scale,
)
def cp_attn_forward_extend( def cp_attn_forward_extend(
@@ -157,15 +157,18 @@ class SWAKVPool(BaseSWAKVPool):
cache_v: torch.Tensor, cache_v: torch.Tensor,
k_scale: float = 1.0, k_scale: float = 1.0,
v_scale: float = 1.0, v_scale: float = 1.0,
swa_loc: Optional[torch.Tensor] = None,
): ):
layer_id = layer.layer_id layer_id = layer.layer_id
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id] layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
if is_swa_layer: if is_swa_layer:
loc = self.translate_loc_from_full_to_swa(loc) # swa_loc is the full->SWA translation, computed once per forward by
# the attention backend; set_kv_buffer never translates internally.
assert swa_loc is not None
self.swa_kv_pool.set_kv_buffer( self.swa_kv_pool.set_kv_buffer(
None, None,
loc, swa_loc,
cache_k, cache_k,
cache_v, cache_v,
k_scale, k_scale,
@@ -576,7 +576,9 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
+ (bs - raw_bs) * self.seq_len_fill_value, + (bs - raw_bs) * self.seq_len_fill_value,
seq_lens_cpu=buffers.seq_lens_cpu, seq_lens_cpu=buffers.seq_lens_cpu,
encoder_lens=None, encoder_lens=None,
out_cache_loc=forward_batch.out_cache_loc, # per-step write target (advanced in-graph by assign_new_state);
# forward_batch.out_cache_loc is frozen at step 0.
out_cache_loc=buffers.out_cache_loc[:num_tokens],
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
) )
self.eagle_worker.draft_extend_attn_backend_list[ self.eagle_worker.draft_extend_attn_backend_list[
@@ -0,0 +1,82 @@
"""Unit coverage for SWAKVPool.set_kv_buffer with a pre-translated swa_loc.
The attention backend translates out_cache_loc once per forward and passes it
in via ``swa_loc`` (cached on its forward metadata); set_kv_buffer uses it
directly for SWA layers and asserts it is provided. The per-backend cuda-graph
buffer plumbing is covered by the backend SWA integration tests.
"""
import sys
import unittest
from pathlib import Path
from types import SimpleNamespace
import torch
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestSWAKVPoolSetKVBuffer(CustomTestCase):
"""set_kv_buffer: SWA layers require a pre-translated swa_loc; full layers
use loc unchanged."""
def _pool_and_record(self):
pool = object.__new__(SWAKVPool)
# layer 0 -> full pool, layer 1 -> swa pool
pool.layers_mapping = {0: (0, False), 1: (0, True)}
recorded = {}
def _swa_set(layer, loc, k, v, k_scale, v_scale, layer_id_override):
recorded["swa_loc"] = loc
def _full_set(layer, loc, k, v, k_scale, v_scale, layer_id_override):
recorded["full_loc"] = loc
pool.swa_kv_pool = SimpleNamespace(set_kv_buffer=_swa_set)
pool.full_kv_pool = SimpleNamespace(set_kv_buffer=_full_set)
return pool, recorded
def test_swa_layer_uses_swa_loc_directly(self):
pool, recorded = self._pool_and_record()
swa_loc = torch.tensor([7, 8])
pool.set_kv_buffer(
SimpleNamespace(layer_id=1),
torch.tensor([3, 4]),
None,
None,
swa_loc=swa_loc,
)
self.assertIs(recorded["swa_loc"], swa_loc)
def test_swa_layer_requires_swa_loc(self):
# set_kv_buffer never translates internally; SWA layers must be given a
# pre-translated swa_loc.
pool, _ = self._pool_and_record()
with self.assertRaises(AssertionError):
pool.set_kv_buffer(
SimpleNamespace(layer_id=1), torch.tensor([3, 4]), None, None
)
def test_full_layer_ignores_swa_loc(self):
pool, recorded = self._pool_and_record()
loc = torch.tensor([3, 4])
# Full layer: swa_loc supplied but ignored; loc is used.
pool.set_kv_buffer(
SimpleNamespace(layer_id=0),
loc,
None,
None,
swa_loc=torch.tensor([99, 99]),
)
self.assertIs(recorded["full_loc"], loc)
if __name__ == "__main__":
unittest.main()