Increase the KV cache pool when using indexShare by 15% (#30310)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
c7ca332fb0
commit
455ab36eeb
@@ -3137,6 +3137,7 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
start_layer: Optional[int] = None,
|
||||
end_layer: Optional[int] = None,
|
||||
index_buf_size: Optional[int] = None,
|
||||
skip_topk_layers: Optional[List[bool]] = None,
|
||||
):
|
||||
override_dim = (
|
||||
kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None
|
||||
@@ -3164,6 +3165,13 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
# num head == 1 and head dim == 128 for index_k in DSA
|
||||
assert index_head_dim == 128
|
||||
|
||||
self.skip_topk_layers = (
|
||||
list(skip_topk_layers)
|
||||
if skip_topk_layers is not None
|
||||
else [False] * layer_num
|
||||
)
|
||||
assert len(self.skip_topk_layers) == layer_num
|
||||
|
||||
if _is_hip:
|
||||
if aiter_can_use_preshuffle_paged_mqa():
|
||||
assert (
|
||||
@@ -3180,6 +3188,10 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
if self.custom_mem_pool
|
||||
else nullcontext()
|
||||
):
|
||||
cols = self.page_size * (
|
||||
index_head_dim + index_head_dim // self.quant_block_size * 4
|
||||
)
|
||||
num_pages = (index_buf_size + page_size + 1) // self.page_size
|
||||
self.index_k_with_scale_buffer = [
|
||||
torch.zeros(
|
||||
# Layout:
|
||||
@@ -3188,17 +3200,11 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
# data: for page i,
|
||||
# * buf[i, :page_size * head_dim] for fp8 data
|
||||
# * buf[i, page_size * head_dim:].view(float32) for scale
|
||||
(
|
||||
(index_buf_size + page_size + 1) // self.page_size,
|
||||
self.page_size
|
||||
* (
|
||||
index_head_dim + index_head_dim // self.quant_block_size * 4
|
||||
),
|
||||
),
|
||||
(0 if self.skip_topk_layers[i] else num_pages, cols),
|
||||
dtype=self.index_k_with_scale_buffer_dtype,
|
||||
device=device,
|
||||
)
|
||||
for _ in range(layer_num)
|
||||
for i in range(layer_num)
|
||||
]
|
||||
self._finalize_allocation_log(size)
|
||||
|
||||
@@ -3215,7 +3221,9 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
|
||||
tgt_loc_flat = tgt_loc.view(-1).long()
|
||||
src_loc_flat = src_loc.view(-1).long()
|
||||
for index_k in self.index_k_with_scale_buffer:
|
||||
for i, index_k in enumerate(self.index_k_with_scale_buffer):
|
||||
if self.skip_topk_layers[i]:
|
||||
continue
|
||||
index_k[tgt_loc_flat] = index_k[src_loc_flat]
|
||||
|
||||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
@@ -3306,6 +3314,9 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
page_chunk_size = max(1, chunk_size // self.page_size)
|
||||
for layer_id in range(self.layer_num):
|
||||
if self.skip_topk_layers[layer_id]:
|
||||
index_k_cpu.append([])
|
||||
continue
|
||||
index_k_cpu.append([])
|
||||
for i in range(0, len(page_indices), page_chunk_size):
|
||||
chunk_page_indices = page_indices[i : i + page_chunk_size]
|
||||
@@ -3328,6 +3339,8 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
page_chunk_size = max(1, chunk_size // self.page_size)
|
||||
for layer_id in range(self.layer_num):
|
||||
if self.skip_topk_layers[layer_id]:
|
||||
continue
|
||||
for i in range(0, len(page_indices), page_chunk_size):
|
||||
chunk_page_indices = page_indices[i : i + page_chunk_size]
|
||||
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
|
||||
@@ -3346,7 +3359,12 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num)
|
||||
]
|
||||
item_lens = [
|
||||
self.index_k_with_scale_buffer[i][0].nbytes for i in range(self.layer_num)
|
||||
(
|
||||
0
|
||||
if self.skip_topk_layers[i]
|
||||
else self.index_k_with_scale_buffer[i][0].nbytes
|
||||
)
|
||||
for i in range(self.layer_num)
|
||||
]
|
||||
return data_ptrs, data_lens, item_lens
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import (
|
||||
dsa_layer_skips_topk,
|
||||
get_dsa_index_head_dim,
|
||||
get_minimax_sparse_attention_config,
|
||||
get_minimax_sparse_disable_value_layer_ids,
|
||||
@@ -947,6 +948,11 @@ class ModelRunnerKVCacheMixin:
|
||||
pool_kwargs["host_to_device_ratio"] = parse_hisparse_config(
|
||||
self.server_args
|
||||
).host_to_device_ratio
|
||||
elif not self.is_draft_worker:
|
||||
pool_kwargs["skip_topk_layers"] = [
|
||||
dsa_layer_skips_topk(self.model_config.hf_config, layer_id)
|
||||
for layer_id in range(self.start_layer, self.end_layer)
|
||||
]
|
||||
self.token_to_kv_pool = PoolCls(
|
||||
self.max_total_num_tokens,
|
||||
page_size=self.page_size,
|
||||
|
||||
@@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import (
|
||||
dsa_layer_skips_topk,
|
||||
get_dsa_index_head_dim,
|
||||
get_minimax_sparse_attention_config,
|
||||
get_minimax_sparse_disable_value_layer_ids,
|
||||
@@ -209,7 +210,15 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
element_size = torch._utils._element_size(
|
||||
DSATokenToKVPool.index_k_with_scale_buffer_dtype
|
||||
)
|
||||
cell_size += indexer_size_per_token * num_layers * element_size
|
||||
if mr.enable_hisparse or mr.is_draft_worker:
|
||||
num_indexer_layers = num_layers
|
||||
else:
|
||||
num_indexer_layers = sum(
|
||||
1
|
||||
for layer_id in range(mr.start_layer, mr.end_layer)
|
||||
if not dsa_layer_skips_topk(model_config.hf_config, layer_id)
|
||||
)
|
||||
cell_size += indexer_size_per_token * num_indexer_layers * element_size
|
||||
elif is_minimax_sparse(model_config.hf_config):
|
||||
# Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool
|
||||
# (sparse-only, single-head; kv layers store K+V, k-only layers store K).
|
||||
|
||||
Reference in New Issue
Block a user