[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
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(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
@@ -276,7 +286,16 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
)
if is_cp_mode:
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
@@ -654,7 +673,17 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
if not layer.is_cross_attention
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(
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.radix_attention import AttentionType
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.speculative.spec_info import SpecInput
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
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
extend_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(
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.
@@ -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
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 = (
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)
value_cache_full = kv_full[..., k_feat_size:].reshape(-1, *v_tail)
token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
key_cache_full,
value_cache_full,
)
if swa_loc is not None:
token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
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):
@@ -329,6 +345,10 @@ class AscendAttnBackend(AttentionBackend):
self.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
self.padding_size_list = [1, 2, 4, 8, 16, 32, 64, 128]
@@ -372,7 +392,10 @@ class AscendAttnBackend(AttentionBackend):
bs = forward_batch.batch_size
if in_capture:
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(
bs=bs,
@@ -385,6 +408,7 @@ class AscendAttnBackend(AttentionBackend):
),
forward_mode=forward_batch.forward_mode,
spec_info=forward_batch.spec_info,
out_cache_loc=forward_batch.out_cache_loc,
)
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
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
@@ -493,18 +524,29 @@ class AscendAttnBackend(AttentionBackend):
dtype=torch.int32,
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(
self,
bs: int,
forward_mode: ForwardMode,
seq_lens: torch.Tensor,
out_cache_loc: Optional[torch.Tensor] = None,
) -> "ForwardMetadata":
"""Create and store the per-bs ForwardMetadata for CUDA graph capture."""
metadata = ForwardMetadata()
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
if self.is_hybrid_swa:
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 = seq_lens
if (
@@ -571,12 +613,21 @@ class AscendAttnBackend(AttentionBackend):
seq_lens_cpu: torch.Tensor,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
out_cache_loc: Optional[torch.Tensor] = None,
):
"""Shared capture+replay body for the cuda-graph init path.
Public entry: :py:meth:`init_forward_metadata_out_graph`.
"""
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()
if forward_mode.is_target_verify():
max_len += self.speculative_num_draft_tokens
@@ -1068,6 +1119,11 @@ class AscendAttnBackend(AttentionBackend):
v,
self.attn_cp_size,
self.token_to_kv_pool,
swa_loc=(
self.forward_metadata.swa_out_cache_loc
if self.use_sliding_window_kv_pool
else None
),
)
else:
# support cross attention
@@ -1076,7 +1132,16 @@ class AscendAttnBackend(AttentionBackend):
if not layer.is_cross_attention
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)
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,
):
if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
)
if 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:
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)
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(
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:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
@@ -1862,6 +1944,14 @@ class AscendAttnBackend(AttentionBackend):
self.token_to_kv_pool.set_kv_buffer(
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:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
@@ -2195,7 +2285,16 @@ class AscendAttnBackend(AttentionBackend):
if not layer.is_cross_attention
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]
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)
@@ -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."
)
if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
)
if 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:
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)
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
num_block, block_size, _, _ = k_cache.shape
@@ -117,6 +117,8 @@ class ForwardMetadata:
max_extend_len: Optional[int] = None
fp8_prefill_kv_indices: 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
@@ -865,6 +867,20 @@ class AiterAttnBackend(AttentionBackend):
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):
"""Init auxiliary variables for aiter attention backend."""
@@ -885,6 +901,11 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits = 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()
if forward_batch.forward_mode.is_decode_or_idle():
@@ -1005,6 +1026,7 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits=num_kv_splits,
run_graph=False,
swa_page_table=swa_page_table,
swa_out_cache_loc=swa_out_cache_loc,
)
elif forward_batch.forward_mode.is_draft_extend_v2():
@@ -1277,6 +1299,7 @@ class AiterAttnBackend(AttentionBackend):
max_kv_len,
max_extend_len=max_q_len,
swa_page_table=swa_page_table,
swa_out_cache_loc=swa_out_cache_loc,
)
else:
qo_indptr = torch.arange(
@@ -1425,6 +1448,7 @@ class AiterAttnBackend(AttentionBackend):
max(forward_batch.extend_seq_lens_cpu),
forward_batch.seq_lens_cpu.max().item(),
swa_page_table=swa_page_table,
swa_out_cache_loc=swa_out_cache_loc,
)
def init_cuda_graph_state(
@@ -1539,6 +1563,13 @@ class AiterAttnBackend(AttentionBackend):
dtype=torch.int32,
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(
self,
@@ -2020,9 +2051,20 @@ class AiterAttnBackend(AttentionBackend):
# launch_reshape_and_cache_flash; always route through
# set_kv_buffer which dispatches to the SHUFFLE 5D writer.
if self.kv_cache_is_vectorized_5d:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, k_descale, v_descale
)
if self.use_sliding_window_kv_pool:
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
# both unified attention and sliding window kv pool are active.
# 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:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, k_descale, v_descale
)
if self.use_sliding_window_kv_pool:
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:
max_q_len = self.forward_metadata.max_q_len
@@ -2487,9 +2540,20 @@ class AiterAttnBackend(AttentionBackend):
if save_kv_cache:
# SHUFFLE 5D pool path — see forward_extend for rationale.
if self.kv_cache_is_vectorized_5d:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v, k_descale, v_descale
)
if self.use_sliding_window_kv_pool:
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
# both unified attention and sliding window kv pool are active.
# 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,
)
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:
self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v
@@ -61,6 +61,8 @@ class FlashAttentionMetadata:
page_table: torch.Tensor = None
# Page table for Sliding Window Attention
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)
scheduler_metadata: torch.Tensor = None
@@ -698,6 +700,12 @@ class FlashAttentionBackend(AttentionBackend):
metadata.page_table
).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
if self.page_size > 1:
@@ -798,7 +806,26 @@ class FlashAttentionBackend(AttentionBackend):
# Dense-MHA CP: k, v are still rank-local; backend
# all-gathers and writes to the per-rank pool.
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:
self.token_to_kv_pool.set_kv_buffer(
@@ -1242,7 +1269,17 @@ class FlashAttentionBackend(AttentionBackend):
if not layer.is_cross_attention
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(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
@@ -1588,6 +1625,13 @@ class FlashAttentionBackend(AttentionBackend):
dtype=torch.int32,
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
if self.topk > 1:
@@ -1849,6 +1893,9 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.decode_cuda_graph_metadata[
"swa_page_table"
][:bs, :]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[
:num_tokens
]
self.decode_cuda_graph_metadata[bs] = metadata
else:
# Draft Decode topk>1: two metadata objects
@@ -1905,6 +1952,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.decode_cuda_graph_metadata[
"swa_page_table"
][:bs, :]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
self.decode_cuda_graph_metadata[bs] = metadata
elif forward_mode.is_target_verify():
@@ -1924,6 +1972,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.target_verify_metadata[
"swa_page_table"
][:bs, :]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
self.target_verify_metadata[bs] = metadata
else:
# 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_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:
metadata_swa = FlashAttentionMetadata()
@@ -1995,6 +2048,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][
:bs, :
]
metadata.swa_out_cache_loc = self.swa_out_cache_loc_buf[:num_tokens]
self.draft_extend_metadata[bs] = metadata
if encoder_lens is not None:
@@ -2037,6 +2091,15 @@ class FlashAttentionBackend(AttentionBackend):
metadata = 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 spec_info is not None:
# 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.radix_attention import AttentionType
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 (
Backend,
Phase,
@@ -138,6 +139,8 @@ class MultiItemScoringParams:
@dataclass
class DecodeMetadata:
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper]
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: Optional[torch.Tensor] = None
@dataclass
@@ -146,6 +149,7 @@ class PrefillMetadata:
use_ragged: bool
extend_no_prefix: bool
multi_item_params: Optional[MultiItemScoringParams] = None
swa_out_cache_loc: Optional[torch.Tensor] = None
# 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.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
# FIXME: remove dllm workarounds from flashinfer
@@ -541,7 +546,28 @@ class FlashInferAttnBackend(AttentionBackend):
for w in self.decode_cuda_graph_metadata[bs]:
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):
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():
self.indices_updater_decode.update(
forward_batch.req_pool_indices,
@@ -554,7 +580,9 @@ class FlashInferAttnBackend(AttentionBackend):
fixed_split_size=self.decode_split_tile_size,
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():
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
@@ -568,7 +596,10 @@ class FlashInferAttnBackend(AttentionBackend):
spec_info=forward_batch.spec_info,
)
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():
self.indices_updater_prefill.update(
@@ -583,7 +614,10 @@ class FlashInferAttnBackend(AttentionBackend):
spec_info=forward_batch.spec_info,
)
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:
prefix_lens = forward_batch.extend_prefix_lens
@@ -631,6 +665,7 @@ class FlashInferAttnBackend(AttentionBackend):
use_ragged,
extend_no_prefix,
multi_item_params,
swa_out_cache_loc=swa_out_cache_loc,
)
def init_cuda_graph_state(
@@ -652,6 +687,14 @@ class FlashInferAttnBackend(AttentionBackend):
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
for i in range(self.num_wrappers):
# Force allocation by performing a small operation
@@ -771,9 +814,20 @@ class FlashInferAttnBackend(AttentionBackend):
if k is not None:
assert v is not None
if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
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,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
causal = (
not layer.is_cross_attention
@@ -863,9 +917,20 @@ class FlashInferAttnBackend(AttentionBackend):
o, _ = _safe_merge_state(o1, s1, o2, s2)
if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
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,
)
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)
@@ -891,9 +956,20 @@ class FlashInferAttnBackend(AttentionBackend):
if k is not None:
assert v is not None
if save_kv_cache:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
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,
)
else:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
# Call the wrapped function
o = decode_wrapper.forward(
@@ -5,6 +5,7 @@ from typing import TYPE_CHECKING
import torch
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
if TYPE_CHECKING:
@@ -24,6 +25,14 @@ class IntelAMXAttnBackend(AttentionBackend):
self.req_to_token_pool = model_runner.req_to_token_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 = (
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()
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):
return 1
@@ -110,7 +128,12 @@ class IntelAMXAttnBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc
)
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
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.radix_attention import AttentionType
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
if TYPE_CHECKING:
@@ -23,6 +24,12 @@ class TorchNativeAttnBackend(AttentionBackend):
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_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
def _make_sliding_window_mask(
@@ -41,7 +48,14 @@ class TorchNativeAttnBackend(AttentionBackend):
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""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(
self,
@@ -281,7 +295,12 @@ class TorchNativeAttnBackend(AttentionBackend):
cache_loc = forward_batch.out_cache_loc
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
@@ -347,7 +366,12 @@ class TorchNativeAttnBackend(AttentionBackend):
cache_loc = forward_batch.out_cache_loc
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
@@ -82,6 +82,8 @@ class ForwardMetadata:
window_kv_offsets: torch.Tensor
# Separate attn_logits for SWA layers when v_head_dim differs
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):
@@ -124,6 +126,7 @@ class TritonAttnBackend(AttentionBackend):
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.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.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
@@ -521,8 +524,9 @@ class TritonAttnBackend(AttentionBackend):
forward_mode=forward_mode,
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(
bs, forward_mode, spec_info
bs, forward_mode, spec_info, swa_out_cache_loc
)
else:
self._apply_cuda_graph_metadata(
@@ -532,6 +536,29 @@ class TritonAttnBackend(AttentionBackend):
forward_mode=forward_mode,
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):
"""Init auxiliary variables for triton attention backend."""
@@ -742,6 +769,12 @@ class TritonAttnBackend(AttentionBackend):
max_extend_len = int(forward_batch.extend_seq_lens.max())
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(
attn_logits,
attn_lse,
@@ -757,6 +790,7 @@ class TritonAttnBackend(AttentionBackend):
window_num_kv_splits,
window_kv_offsets,
swa_attn_logits=swa_attn_logits,
swa_out_cache_loc=swa_out_cache_loc,
)
def init_cuda_graph_state(
@@ -839,11 +873,20 @@ class TritonAttnBackend(AttentionBackend):
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(
self,
bs: int,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
swa_out_cache_loc: Optional[torch.Tensor] = None,
) -> ForwardMetadata:
"""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
self.cuda_graph_* tensors that the captured graph kernels will
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
if forward_mode.is_decode_or_idle():
@@ -872,6 +916,7 @@ class TritonAttnBackend(AttentionBackend):
),
window_kv_offsets=None,
swa_attn_logits=self.cuda_graph_swa_attn_logits,
swa_out_cache_loc=swa_out_cache_loc,
)
elif forward_mode.is_target_verify():
custom_mask = (
@@ -896,6 +941,7 @@ class TritonAttnBackend(AttentionBackend):
self.cuda_graph_window_num_kv_splits 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):
return ForwardMetadata(
@@ -920,6 +966,7 @@ class TritonAttnBackend(AttentionBackend):
window_kv_indices=None,
window_num_kv_splits=None,
window_kv_offsets=None,
swa_out_cache_loc=swa_out_cache_loc,
)
else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
@@ -1006,7 +1053,27 @@ class TritonAttnBackend(AttentionBackend):
else:
# Save KV cache first (must do this before unified kernel)
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(
layer,
forward_batch.out_cache_loc,
@@ -1270,6 +1337,16 @@ class TritonAttnBackend(AttentionBackend):
k,
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:
self.token_to_kv_pool.set_kv_buffer(
layer,
@@ -60,6 +60,8 @@ class TRTLLMMHAMetadata:
page_table: torch.Tensor = None
# Page table for SWA layers (translated from full pool indices to SWA pool indices)
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):
@@ -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 (
self.speculative_num_draft_tokens is not None
and self.speculative_num_draft_tokens > 0
@@ -414,6 +425,10 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
)
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
def _apply_cuda_graph_metadata(
@@ -620,6 +635,17 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
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):
"""Initialize the metadata for a forward pass."""
@@ -718,6 +744,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
# Compute SWA page table (None for non-SWA models)
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
if self.page_size > 1:
self.strided_indices = torch.arange(
@@ -762,9 +796,20 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
else:
# Use original set_kv_buffer path
if save_kv_cache and k is not None:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
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,
)
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
if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl):
@@ -848,9 +893,20 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
else:
# Use original set_kv_buffer path
if save_kv_cache and k is not None:
self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
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,
)
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:
q = q.to(torch.float8_e4m3fn)
@@ -390,6 +390,12 @@ class XPUAttentionBackend(AttentionBackend):
metadata.page_table
).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:
workspace_size = flash_mla_get_workspace_size(
@@ -416,6 +422,12 @@ class XPUAttentionBackend(AttentionBackend):
metadata.page_table
).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
if self.page_size > 1:
@@ -464,7 +476,17 @@ class XPUAttentionBackend(AttentionBackend):
if not layer.is_cross_attention
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(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
)
@@ -766,7 +788,17 @@ class XPUAttentionBackend(AttentionBackend):
if not layer.is_cross_attention
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(
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
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
into each rank's local memory pool.
swa_loc is the pre-translated full->SWA write target for hybrid SWA pools.
"""
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()
)
get_token_to_kv_pool().set_kv_buffer(
layer,
cache_loc,
key_cache_full,
value_cache_full,
layer.k_scale,
layer.v_scale,
)
pool = get_token_to_kv_pool()
if swa_loc is not None:
pool.set_kv_buffer(
layer,
cache_loc,
key_cache_full,
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(
@@ -157,15 +157,18 @@ class SWAKVPool(BaseSWAKVPool):
cache_v: torch.Tensor,
k_scale: float = 1.0,
v_scale: float = 1.0,
swa_loc: Optional[torch.Tensor] = None,
):
layer_id = layer.layer_id
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
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(
None,
loc,
swa_loc,
cache_k,
cache_v,
k_scale,
@@ -576,7 +576,9 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
+ (bs - raw_bs) * self.seq_len_fill_value,
seq_lens_cpu=buffers.seq_lens_cpu,
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,
)
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()