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 (
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -263,19 +264,14 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
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(
|
||||
layer,
|
||||
cache_loc,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
elif not self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
@@ -673,19 +669,14 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
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(
|
||||
layer,
|
||||
cache_loc,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
elif not self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
else:
|
||||
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.radix_attention import AttentionType
|
||||
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
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)
|
||||
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(
|
||||
layer,
|
||||
cache_loc,
|
||||
key_cache_full,
|
||||
value_cache_full,
|
||||
swa_loc=swa_loc,
|
||||
)
|
||||
else:
|
||||
token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
key_cache_full,
|
||||
value_cache_full,
|
||||
)
|
||||
token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, swa_loc),
|
||||
key_cache_full,
|
||||
value_cache_full,
|
||||
)
|
||||
|
||||
|
||||
class AscendAttnBackend(AttentionBackend):
|
||||
@@ -1119,11 +1111,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
v,
|
||||
self.attn_cp_size,
|
||||
self.token_to_kv_pool,
|
||||
swa_loc=(
|
||||
self.forward_metadata.swa_out_cache_loc
|
||||
if self.use_sliding_window_kv_pool
|
||||
else None
|
||||
),
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
# support cross attention
|
||||
@@ -1132,16 +1120,14 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
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)
|
||||
swa_loc = (
|
||||
self.forward_metadata.swa_out_cache_loc
|
||||
if not layer.is_cross_attention
|
||||
else None
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, KVWriteLoc(cache_loc, swa_loc), k, v
|
||||
)
|
||||
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
@@ -1652,18 +1638,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
topk_indices: Optional[torch.Tensor] = None,
|
||||
):
|
||||
if save_kv_cache:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(
|
||||
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
|
||||
)
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
)
|
||||
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
@@ -1725,17 +1708,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, forward_batch.out_cache_loc, k, k_rope
|
||||
)
|
||||
elif self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k,
|
||||
v,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, forward_batch.out_cache_loc, k, v
|
||||
layer,
|
||||
KVWriteLoc(
|
||||
forward_batch.out_cache_loc,
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
)
|
||||
|
||||
if not self.use_mla:
|
||||
@@ -1944,17 +1925,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, forward_batch.out_cache_loc, k, k_rope
|
||||
)
|
||||
elif self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k,
|
||||
v,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, forward_batch.out_cache_loc, k, v
|
||||
layer,
|
||||
KVWriteLoc(
|
||||
forward_batch.out_cache_loc,
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
)
|
||||
|
||||
if sinks is not None:
|
||||
@@ -2285,16 +2264,17 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
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)
|
||||
# swa_out_cache_loc is the full->SWA write target, derived from
|
||||
# out_cache_loc; it must not be applied to cross-attention writes
|
||||
# (which target encoder_out_cache_loc) and is None for non-SWA pools.
|
||||
swa_loc = (
|
||||
self.forward_metadata.swa_out_cache_loc
|
||||
if not layer.is_cross_attention
|
||||
else None
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, KVWriteLoc(cache_loc, swa_loc), k, v
|
||||
)
|
||||
num_tokens = q.shape[0]
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
@@ -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."
|
||||
)
|
||||
if save_kv_cache:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(
|
||||
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
|
||||
)
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
num_block, block_size, _, _ = k_cache.shape
|
||||
|
||||
@@ -7,6 +7,7 @@ from sglang.srt.mem_cache.memory_pool import (
|
||||
MHATokenToKVPool,
|
||||
MLATokenToKVPool,
|
||||
get_tensor_size_bytes,
|
||||
unwrap_write_loc,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
from sglang.srt.utils.common import is_npu
|
||||
@@ -170,13 +171,14 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: "RadixAttention",
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
layer_id_override: Optional[int] = None,
|
||||
):
|
||||
loc, _ = unwrap_write_loc(loc_info)
|
||||
if layer_id_override is not None:
|
||||
layer_id = layer_id_override
|
||||
else:
|
||||
@@ -434,10 +436,11 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: "RadixAttention",
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
):
|
||||
loc, _ = unwrap_write_loc(loc_info)
|
||||
layer_id = layer.layer_id
|
||||
if cache_k.dtype != self.dtype:
|
||||
cache_k = cache_k.to(self.dtype)
|
||||
|
||||
@@ -69,6 +69,7 @@ from sglang.srt.layers.attention.utils import (
|
||||
pad_sequence_with_mask,
|
||||
)
|
||||
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.utils import get_bool_env_var
|
||||
|
||||
@@ -2051,20 +2052,14 @@ class AiterAttnBackend(AttentionBackend):
|
||||
# launch_reshape_and_cache_flash; always route through
|
||||
# set_kv_buffer which dispatches to the SHUFFLE 5D writer.
|
||||
if self.kv_cache_is_vectorized_5d:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
k_descale,
|
||||
v_descale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, k_descale, v_descale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
k_descale,
|
||||
v_descale,
|
||||
)
|
||||
# Only use SWA-specific kv cache write (reshape_and_cache_flash) when
|
||||
# both unified attention and sliding window kv pool are active.
|
||||
# Non-SWA models (e.g. Qwen3-VL) enabled via SGLANG_USE_AITER_UNIFIED_ATTN
|
||||
@@ -2101,20 +2096,14 @@ class AiterAttnBackend(AttentionBackend):
|
||||
elif self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
||||
else:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
k_descale,
|
||||
v_descale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, k_descale, v_descale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
k_descale,
|
||||
v_descale,
|
||||
)
|
||||
|
||||
if self.use_mla:
|
||||
max_q_len = self.forward_metadata.max_q_len
|
||||
@@ -2540,20 +2529,17 @@ class AiterAttnBackend(AttentionBackend):
|
||||
if save_kv_cache:
|
||||
# SHUFFLE 5D pool path — see forward_extend for rationale.
|
||||
if self.kv_cache_is_vectorized_5d:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(
|
||||
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
|
||||
)
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
k_descale,
|
||||
v_descale,
|
||||
)
|
||||
# Only use SWA-specific kv cache write (reshape_and_cache_flash) when
|
||||
# both unified attention and sliding window kv pool are active.
|
||||
# Non-SWA models (e.g. Qwen3-VL) enabled via SGLANG_USE_AITER_UNIFIED_ATTN
|
||||
@@ -2595,17 +2581,15 @@ class AiterAttnBackend(AttentionBackend):
|
||||
),
|
||||
forward_batch.out_cache_loc,
|
||||
)
|
||||
elif self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k,
|
||||
v,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, forward_batch.out_cache_loc, k, v
|
||||
layer,
|
||||
KVWriteLoc(
|
||||
forward_batch.out_cache_loc,
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
)
|
||||
|
||||
if self.use_mla:
|
||||
|
||||
@@ -18,6 +18,7 @@ from sglang.srt.layers.utils.cp_utils import (
|
||||
cp_allgather_and_save_kv_cache,
|
||||
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
@@ -817,19 +818,14 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else None
|
||||
),
|
||||
)
|
||||
elif self.use_sliding_window_kv_pool:
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
|
||||
# Use precomputed metadata across all layers
|
||||
@@ -1269,19 +1265,14 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
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(
|
||||
layer,
|
||||
cache_loc,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
elif not self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
else:
|
||||
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.radix_attention import AttentionType
|
||||
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.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
@@ -814,20 +815,14 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
if save_kv_cache:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
|
||||
causal = (
|
||||
not layer.is_cross_attention
|
||||
@@ -917,20 +912,14 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
o, _ = _safe_merge_state(o1, s1, o2, s2)
|
||||
|
||||
if save_kv_cache:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
|
||||
return o.view(-1, layer.tp_q_head_num * layer.head_dim)
|
||||
|
||||
@@ -956,20 +945,14 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
if save_kv_cache:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
|
||||
# Call the wrapped function
|
||||
o = decode_wrapper.forward(
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import TYPE_CHECKING
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
@@ -128,12 +129,12 @@ class IntelAMXAttnBackend(AttentionBackend):
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
if save_kv_cache and k is not None and v is not None:
|
||||
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)
|
||||
# Cross-attention never writes to the SWA pool, so only thread the
|
||||
# full->SWA location for non-cross-attention layers.
|
||||
swa_loc = None if layer.is_cross_attention else self.swa_out_cache_loc
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, KVWriteLoc(cache_loc, swa_loc), k, v
|
||||
)
|
||||
|
||||
_, max_extend_len = self.forward_metadata
|
||||
self.extend_attention_fwd(
|
||||
|
||||
@@ -7,6 +7,7 @@ from torch.nn.functional import scaled_dot_product_attention
|
||||
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.radix_attention import AttentionType
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
@@ -295,12 +296,9 @@ class TorchNativeAttnBackend(AttentionBackend):
|
||||
cache_loc = forward_batch.out_cache_loc
|
||||
|
||||
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(
|
||||
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)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, KVWriteLoc(cache_loc, self.swa_out_cache_loc), k, v
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
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(
|
||||
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)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, KVWriteLoc(cache_loc, self.swa_out_cache_loc), k, v
|
||||
)
|
||||
|
||||
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.dp_attention import get_attention_tp_size
|
||||
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.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
@@ -1053,30 +1054,14 @@ class TritonAttnBackend(AttentionBackend):
|
||||
else:
|
||||
# Save KV cache first (must do this before unified kernel)
|
||||
if save_kv_cache:
|
||||
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:
|
||||
loc_info = KVWriteLoc(
|
||||
forward_batch.out_cache_loc,
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
if layer.k_scale is None:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
loc_info,
|
||||
k,
|
||||
v,
|
||||
)
|
||||
@@ -1087,14 +1072,14 @@ class TritonAttnBackend(AttentionBackend):
|
||||
k_scaled = k.clone().div_(layer.k_scale)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
loc_info,
|
||||
k_scaled,
|
||||
v,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
loc_info,
|
||||
k.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer
|
||||
v.clone(),
|
||||
layer.k_scale,
|
||||
@@ -1337,20 +1322,13 @@ class TritonAttnBackend(AttentionBackend):
|
||||
k,
|
||||
v,
|
||||
)
|
||||
elif self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
KVWriteLoc(
|
||||
forward_batch.out_cache_loc,
|
||||
self.forward_metadata.swa_out_cache_loc,
|
||||
),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
|
||||
@@ -20,6 +20,7 @@ from sglang.srt.layers.attention.triton_ops.trtllm_fp8_kv_kernel import (
|
||||
fused_fp8_set_kv_buffer,
|
||||
)
|
||||
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.utils import is_flashinfer_available
|
||||
@@ -796,20 +797,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
else:
|
||||
# Use original set_kv_buffer path
|
||||
if save_kv_cache and k is not None:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
|
||||
# For XQA, q_dtype should be bf16
|
||||
if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl):
|
||||
@@ -893,20 +888,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
else:
|
||||
# Use original set_kv_buffer path
|
||||
if save_kv_cache and k is not None:
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
|
||||
if self.data_type == 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,
|
||||
)
|
||||
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
|
||||
@@ -476,19 +477,14 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
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(
|
||||
layer,
|
||||
cache_loc,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
elif not self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
@@ -788,19 +784,14 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
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(
|
||||
layer,
|
||||
cache_loc,
|
||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||
k,
|
||||
v,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||
)
|
||||
elif not self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||
)
|
||||
else:
|
||||
k_rope_val = (
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
is_allocation_symmetric,
|
||||
)
|
||||
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.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()
|
||||
)
|
||||
|
||||
pool = get_token_to_kv_pool()
|
||||
if swa_loc is not None:
|
||||
pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
key_cache_full,
|
||||
value_cache_full,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
swa_loc=swa_loc,
|
||||
)
|
||||
else:
|
||||
pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
key_cache_full,
|
||||
value_cache_full,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
get_token_to_kv_pool().set_kv_buffer(
|
||||
layer,
|
||||
KVWriteLoc(cache_loc, swa_loc),
|
||||
key_cache_full,
|
||||
value_cache_full,
|
||||
layer.k_scale,
|
||||
layer.v_scale,
|
||||
)
|
||||
|
||||
|
||||
def cp_attn_forward_extend(
|
||||
|
||||
@@ -765,6 +765,26 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
||||
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):
|
||||
@abc.abstractmethod
|
||||
def __init__(
|
||||
@@ -1199,13 +1219,14 @@ class MHATokenToKVPool(KVCache):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = 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
|
||||
# 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)")
|
||||
@@ -1523,13 +1544,14 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = 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)")
|
||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||
|
||||
@@ -1919,10 +1941,11 @@ class MLATokenToKVPool(KVCache):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: 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)")
|
||||
layer_id = layer.layer_id
|
||||
assert not self.dsa_kv_cache_store_fp8
|
||||
@@ -2113,10 +2136,12 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: 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)")
|
||||
layer_id = layer.layer_id
|
||||
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.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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -152,14 +156,14 @@ class SWAKVPool(BaseSWAKVPool):
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
loc_info,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
k_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_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||
if is_swa_layer:
|
||||
|
||||
Reference in New Issue
Block a user