Bundle set_kv_buffer write targets into KVWriteLoc (loc + swa_loc) (#27695)
This commit is contained in:
@@ -22,6 +22,7 @@ from sglang.srt.layers.radix_attention import AttentionType
|
|||||||
from sglang.srt.layers.utils.cp_utils import (
|
from sglang.srt.layers.utils.cp_utils import (
|
||||||
cp_allgather_and_save_kv_cache,
|
cp_allgather_and_save_kv_cache,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -263,19 +264,14 @@ 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 self.use_sliding_window_kv_pool:
|
if not self.use_mla:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
cache_loc,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
layer.v_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
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
@@ -673,19 +669,14 @@ 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 self.use_sliding_window_kv_pool:
|
if not self.use_mla:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
cache_loc,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
layer.v_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
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
|
|||||||
@@ -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.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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
|
||||||
@@ -265,21 +266,12 @@ 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)
|
||||||
|
|
||||||
if swa_loc is not None:
|
token_to_kv_pool.set_kv_buffer(
|
||||||
token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, swa_loc),
|
||||||
cache_loc,
|
key_cache_full,
|
||||||
key_cache_full,
|
value_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):
|
||||||
@@ -1119,11 +1111,7 @@ 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=(
|
swa_loc=self.forward_metadata.swa_out_cache_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
|
||||||
@@ -1132,16 +1120,14 @@ 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
|
||||||
)
|
)
|
||||||
if self.use_sliding_window_kv_pool and not layer.is_cross_attention:
|
swa_loc = (
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.forward_metadata.swa_out_cache_loc
|
||||||
layer,
|
if not layer.is_cross_attention
|
||||||
cache_loc,
|
else None
|
||||||
k,
|
)
|
||||||
v,
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
layer, KVWriteLoc(cache_loc, swa_loc), k, v
|
||||||
)
|
)
|
||||||
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)
|
||||||
@@ -1652,18 +1638,15 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
topk_indices: Optional[torch.Tensor] = None,
|
topk_indices: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
if save_kv_cache:
|
if save_kv_cache:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(
|
||||||
forward_batch.out_cache_loc,
|
forward_batch.out_cache_loc,
|
||||||
k,
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
v,
|
),
|
||||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
k,
|
||||||
)
|
v,
|
||||||
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)
|
||||||
@@ -1725,17 +1708,15 @@ 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,
|
||||||
|
KVWriteLoc(
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
|
),
|
||||||
|
k,
|
||||||
|
v,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.use_mla:
|
if not self.use_mla:
|
||||||
@@ -1944,17 +1925,15 @@ 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,
|
||||||
|
KVWriteLoc(
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
|
),
|
||||||
|
k,
|
||||||
|
v,
|
||||||
)
|
)
|
||||||
|
|
||||||
if sinks is not None:
|
if sinks is not None:
|
||||||
@@ -2285,16 +2264,17 @@ 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
|
||||||
)
|
)
|
||||||
if self.use_sliding_window_kv_pool and not layer.is_cross_attention:
|
# swa_out_cache_loc is the full->SWA write target, derived from
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
# out_cache_loc; it must not be applied to cross-attention writes
|
||||||
layer,
|
# (which target encoder_out_cache_loc) and is None for non-SWA pools.
|
||||||
cache_loc,
|
swa_loc = (
|
||||||
k,
|
self.forward_metadata.swa_out_cache_loc
|
||||||
v,
|
if not layer.is_cross_attention
|
||||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
else None
|
||||||
)
|
)
|
||||||
else:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
layer, KVWriteLoc(cache_loc, swa_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)
|
||||||
@@ -2586,18 +2566,15 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(
|
||||||
forward_batch.out_cache_loc,
|
forward_batch.out_cache_loc,
|
||||||
k,
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
v,
|
),
|
||||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
k,
|
||||||
)
|
v,
|
||||||
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
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from sglang.srt.mem_cache.memory_pool import (
|
|||||||
MHATokenToKVPool,
|
MHATokenToKVPool,
|
||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
get_tensor_size_bytes,
|
get_tensor_size_bytes,
|
||||||
|
unwrap_write_loc,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var
|
from sglang.srt.utils import get_bool_env_var
|
||||||
from sglang.srt.utils.common import is_npu
|
from sglang.srt.utils.common import is_npu
|
||||||
@@ -170,13 +171,14 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: "RadixAttention",
|
layer: "RadixAttention",
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
k_scale: Optional[float] = None,
|
k_scale: Optional[float] = None,
|
||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
layer_id_override: Optional[int] = None,
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
loc, _ = unwrap_write_loc(loc_info)
|
||||||
if layer_id_override is not None:
|
if layer_id_override is not None:
|
||||||
layer_id = layer_id_override
|
layer_id = layer_id_override
|
||||||
else:
|
else:
|
||||||
@@ -434,10 +436,11 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: "RadixAttention",
|
layer: "RadixAttention",
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
loc, _ = unwrap_write_loc(loc_info)
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
if cache_k.dtype != self.dtype:
|
if cache_k.dtype != self.dtype:
|
||||||
cache_k = cache_k.to(self.dtype)
|
cache_k = cache_k.to(self.dtype)
|
||||||
|
|||||||
@@ -69,6 +69,7 @@ from sglang.srt.layers.attention.utils import (
|
|||||||
pad_sequence_with_mask,
|
pad_sequence_with_mask,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.utils import get_bool_env_var
|
from sglang.srt.utils import get_bool_env_var
|
||||||
|
|
||||||
@@ -2051,20 +2052,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
k_descale,
|
||||||
k_descale,
|
v_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
|
||||||
@@ -2101,20 +2096,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
k_descale,
|
||||||
k_descale,
|
v_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
|
||||||
@@ -2540,20 +2529,17 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(
|
||||||
forward_batch.out_cache_loc,
|
forward_batch.out_cache_loc,
|
||||||
k,
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
v,
|
),
|
||||||
k_descale,
|
k,
|
||||||
v_descale,
|
v,
|
||||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
k_descale,
|
||||||
)
|
v_descale,
|
||||||
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
|
||||||
@@ -2595,17 +2581,15 @@ 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,
|
||||||
|
KVWriteLoc(
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
|
),
|
||||||
|
k,
|
||||||
|
v,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from sglang.srt.layers.utils.cp_utils import (
|
|||||||
cp_allgather_and_save_kv_cache,
|
cp_allgather_and_save_kv_cache,
|
||||||
cp_attn_forward_extend,
|
cp_attn_forward_extend,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
@@ -817,19 +818,14 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
elif self.use_sliding_window_kv_pool:
|
else:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
cache_loc,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
layer.v_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
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Use precomputed metadata across all layers
|
# Use precomputed metadata across all layers
|
||||||
@@ -1269,19 +1265,14 @@ 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 self.use_sliding_window_kv_pool:
|
if not self.use_mla:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
cache_loc,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
layer.v_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
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
|
|||||||
@@ -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.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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,
|
||||||
@@ -814,20 +815,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
layer.k_scale,
|
||||||
layer.k_scale,
|
layer.v_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
|
||||||
@@ -917,20 +912,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
layer.k_scale,
|
||||||
layer.k_scale,
|
layer.v_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)
|
||||||
|
|
||||||
@@ -956,20 +945,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
layer.k_scale,
|
||||||
layer.k_scale,
|
layer.v_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.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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
|
||||||
|
|
||||||
@@ -128,12 +129,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:
|
||||||
if self.use_sliding_window_kv_pool and not layer.is_cross_attention:
|
# Cross-attention never writes to the SWA pool, so only thread the
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
# full->SWA location for non-cross-attention layers.
|
||||||
layer, cache_loc, k, v, swa_loc=self.swa_out_cache_loc
|
swa_loc = None if layer.is_cross_attention else self.swa_out_cache_loc
|
||||||
)
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
else:
|
layer, KVWriteLoc(cache_loc, swa_loc), k, v
|
||||||
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.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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
|
||||||
|
|
||||||
@@ -295,12 +296,9 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer, KVWriteLoc(cache_loc, self.swa_out_cache_loc), k, v
|
||||||
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
|
||||||
|
|
||||||
@@ -366,12 +364,9 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer, KVWriteLoc(cache_loc, self.swa_out_cache_loc), k, v
|
||||||
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
|
||||||
|
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
|||||||
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
|
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
|
||||||
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.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
@@ -1053,30 +1054,14 @@ 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 self.use_sliding_window_kv_pool:
|
loc_info = KVWriteLoc(
|
||||||
# SWA pool (never MLA); clone k,v when scaling, as below.
|
forward_batch.out_cache_loc,
|
||||||
if layer.k_scale is None:
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
)
|
||||||
layer,
|
if layer.k_scale is None:
|
||||||
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,
|
loc_info,
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
)
|
)
|
||||||
@@ -1087,14 +1072,14 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
k_scaled = k.clone().div_(layer.k_scale)
|
k_scaled = k.clone().div_(layer.k_scale)
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
forward_batch.out_cache_loc,
|
loc_info,
|
||||||
k_scaled,
|
k_scaled,
|
||||||
v,
|
v,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
forward_batch.out_cache_loc,
|
loc_info,
|
||||||
k.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer
|
k.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer
|
||||||
v.clone(),
|
v.clone(),
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
@@ -1337,20 +1322,13 @@ 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,
|
||||||
forward_batch.out_cache_loc,
|
KVWriteLoc(
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
self.forward_metadata.swa_out_cache_loc,
|
||||||
|
),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from sglang.srt.layers.attention.triton_ops.trtllm_fp8_kv_kernel import (
|
|||||||
fused_fp8_set_kv_buffer,
|
fused_fp8_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.utils import canonicalize_stride
|
from sglang.srt.layers.attention.utils import canonicalize_stride
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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.utils import is_flashinfer_available
|
from sglang.srt.utils import is_flashinfer_available
|
||||||
@@ -796,20 +797,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
layer.k_scale,
|
||||||
layer.k_scale,
|
layer.v_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):
|
||||||
@@ -893,20 +888,14 @@ 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:
|
||||||
if self.use_sliding_window_kv_pool:
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
layer,
|
||||||
layer,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
cache_loc,
|
k,
|
||||||
k,
|
v,
|
||||||
v,
|
layer.k_scale,
|
||||||
layer.k_scale,
|
layer.v_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)
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from sglang.srt.layers.attention.flashattention_backend import (
|
|||||||
prepare_swa_spec_page_table_triton,
|
prepare_swa_spec_page_table_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import get_global_server_args
|
from sglang.srt.managers.schedule_batch import get_global_server_args
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
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
|
||||||
|
|
||||||
@@ -476,19 +477,14 @@ 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 self.use_sliding_window_kv_pool:
|
if not self.use_mla:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
cache_loc,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
layer.v_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
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
@@ -788,19 +784,14 @@ 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 self.use_sliding_window_kv_pool:
|
if not self.use_mla:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
cache_loc,
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
k,
|
k,
|
||||||
v,
|
v,
|
||||||
layer.k_scale,
|
layer.k_scale,
|
||||||
layer.v_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
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
k_rope_val = (
|
k_rope_val = (
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
is_allocation_symmetric,
|
is_allocation_symmetric,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe import get_moe_a2a_backend
|
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
|
|
||||||
@@ -438,26 +439,14 @@ def cp_allgather_and_save_kv_cache(forward_batch, layer, k, v, cp_size, swa_loc=
|
|||||||
v, cp_size, forward_batch, torch.cuda.current_stream()
|
v, cp_size, forward_batch, torch.cuda.current_stream()
|
||||||
)
|
)
|
||||||
|
|
||||||
pool = get_token_to_kv_pool()
|
get_token_to_kv_pool().set_kv_buffer(
|
||||||
if swa_loc is not None:
|
layer,
|
||||||
pool.set_kv_buffer(
|
KVWriteLoc(cache_loc, swa_loc),
|
||||||
layer,
|
key_cache_full,
|
||||||
cache_loc,
|
value_cache_full,
|
||||||
key_cache_full,
|
layer.k_scale,
|
||||||
value_cache_full,
|
layer.v_scale,
|
||||||
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(
|
||||||
|
|||||||
@@ -765,6 +765,26 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_()
|
self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class KVWriteLoc:
|
||||||
|
"""Write target(s) for ``KVCache.set_kv_buffer``.
|
||||||
|
|
||||||
|
``loc`` is the full-pool write location; ``swa_loc`` is the pre-translated
|
||||||
|
full->SWA location for hybrid SWA pools (``None`` otherwise). Bundling them
|
||||||
|
lets a backend issue one ``set_kv_buffer`` call regardless of pool type.
|
||||||
|
"""
|
||||||
|
|
||||||
|
loc: torch.Tensor
|
||||||
|
swa_loc: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
|
def unwrap_write_loc(loc_info):
|
||||||
|
"""Return ``(loc, swa_loc)`` from a ``KVWriteLoc`` or a bare loc tensor."""
|
||||||
|
if isinstance(loc_info, KVWriteLoc):
|
||||||
|
return loc_info.loc, loc_info.swa_loc
|
||||||
|
return loc_info, None
|
||||||
|
|
||||||
|
|
||||||
class KVCache(abc.ABC):
|
class KVCache(abc.ABC):
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -1199,13 +1219,14 @@ class MHATokenToKVPool(KVCache):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
k_scale: Optional[float] = None,
|
k_scale: Optional[float] = None,
|
||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
layer_id_override: Optional[int] = None,
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
loc, _ = unwrap_write_loc(loc_info)
|
||||||
# Catch stale slot ids here instead of as illegal-addr / silent KV
|
# Catch stale slot ids here instead of as illegal-addr / silent KV
|
||||||
# corruption in the store_kvcache write (gated on SGLANG_ENABLE_ASYNC_ASSERT).
|
# corruption in the store_kvcache write (gated on SGLANG_ENABLE_ASYNC_ASSERT).
|
||||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA)")
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA)")
|
||||||
@@ -1523,13 +1544,14 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
k_scale: Optional[float] = None,
|
k_scale: Optional[float] = None,
|
||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
layer_id_override: Optional[int] = None,
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
loc, _ = unwrap_write_loc(loc_info)
|
||||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA-FP4)")
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA-FP4)")
|
||||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
|
|
||||||
@@ -1919,10 +1941,11 @@ class MLATokenToKVPool(KVCache):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
loc, _ = unwrap_write_loc(loc_info)
|
||||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
assert not self.dsa_kv_cache_store_fp8
|
assert not self.dsa_kv_cache_store_fp8
|
||||||
@@ -2113,10 +2136,12 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
# loc_info may be a KVWriteLoc; MLA pools have no SWA target.
|
||||||
|
loc, _ = unwrap_write_loc(loc_info)
|
||||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA-FP4)")
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA-FP4)")
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
assert not self.dsa_kv_cache_store_fp8
|
assert not self.dsa_kv_cache_store_fp8
|
||||||
|
|||||||
@@ -5,7 +5,11 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache, MHATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
|
KVCache,
|
||||||
|
MHATokenToKVPool,
|
||||||
|
unwrap_write_loc,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
|
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -152,14 +156,14 @@ class SWAKVPool(BaseSWAKVPool):
|
|||||||
def set_kv_buffer(
|
def set_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
loc: torch.Tensor,
|
loc_info,
|
||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
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,
|
|
||||||
):
|
):
|
||||||
|
# loc_info bundles the full loc and the pre-translated SWA loc.
|
||||||
|
loc, swa_loc = unwrap_write_loc(loc_info)
|
||||||
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:
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
"""Unit coverage for SWAKVPool.set_kv_buffer with a pre-translated swa_loc.
|
"""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
|
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
|
in via a ``KVWriteLoc`` (loc + swa_loc) on the loc_info argument; set_kv_buffer
|
||||||
directly for SWA layers and asserts it is provided. The per-backend cuda-graph
|
uses swa_loc directly for SWA layers and asserts it is provided. The per-backend
|
||||||
buffer plumbing is covered by the backend SWA integration tests.
|
cuda-graph buffer plumbing is covered by the backend SWA integration tests.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import sys
|
import sys
|
||||||
@@ -13,6 +13,7 @@ from types import SimpleNamespace
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -48,21 +49,27 @@ class TestSWAKVPoolSetKVBuffer(CustomTestCase):
|
|||||||
swa_loc = torch.tensor([7, 8])
|
swa_loc = torch.tensor([7, 8])
|
||||||
pool.set_kv_buffer(
|
pool.set_kv_buffer(
|
||||||
SimpleNamespace(layer_id=1),
|
SimpleNamespace(layer_id=1),
|
||||||
torch.tensor([3, 4]),
|
KVWriteLoc(torch.tensor([3, 4]), swa_loc),
|
||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
swa_loc=swa_loc,
|
|
||||||
)
|
)
|
||||||
self.assertIs(recorded["swa_loc"], swa_loc)
|
self.assertIs(recorded["swa_loc"], swa_loc)
|
||||||
|
|
||||||
def test_swa_layer_requires_swa_loc(self):
|
def test_swa_layer_requires_swa_loc(self):
|
||||||
# set_kv_buffer never translates internally; SWA layers must be given a
|
# set_kv_buffer never translates internally; SWA layers must be given a
|
||||||
# pre-translated swa_loc.
|
# pre-translated swa_loc (loc_info without swa_loc, or a bare loc).
|
||||||
pool, _ = self._pool_and_record()
|
pool, _ = self._pool_and_record()
|
||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(AssertionError):
|
||||||
pool.set_kv_buffer(
|
pool.set_kv_buffer(
|
||||||
SimpleNamespace(layer_id=1), torch.tensor([3, 4]), None, None
|
SimpleNamespace(layer_id=1), torch.tensor([3, 4]), None, None
|
||||||
)
|
)
|
||||||
|
with self.assertRaises(AssertionError):
|
||||||
|
pool.set_kv_buffer(
|
||||||
|
SimpleNamespace(layer_id=1),
|
||||||
|
KVWriteLoc(torch.tensor([3, 4])),
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
def test_full_layer_ignores_swa_loc(self):
|
def test_full_layer_ignores_swa_loc(self):
|
||||||
pool, recorded = self._pool_and_record()
|
pool, recorded = self._pool_and_record()
|
||||||
@@ -70,10 +77,9 @@ class TestSWAKVPoolSetKVBuffer(CustomTestCase):
|
|||||||
# Full layer: swa_loc supplied but ignored; loc is used.
|
# Full layer: swa_loc supplied but ignored; loc is used.
|
||||||
pool.set_kv_buffer(
|
pool.set_kv_buffer(
|
||||||
SimpleNamespace(layer_id=0),
|
SimpleNamespace(layer_id=0),
|
||||||
loc,
|
KVWriteLoc(loc, torch.tensor([99, 99])),
|
||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
swa_loc=torch.tensor([99, 99]),
|
|
||||||
)
|
)
|
||||||
self.assertIs(recorded["full_loc"], loc)
|
self.assertIs(recorded["full_loc"], loc)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user