[SWA] Cache full→SWA out_cache_loc per forward across attention backends (#27617)
This commit is contained in:
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user