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,
|
start_layer: Optional[int] = None,
|
||||||
end_layer: Optional[int] = None,
|
end_layer: Optional[int] = None,
|
||||||
index_buf_size: Optional[int] = None,
|
index_buf_size: Optional[int] = None,
|
||||||
|
skip_topk_layers: Optional[List[bool]] = None,
|
||||||
):
|
):
|
||||||
override_dim = (
|
override_dim = (
|
||||||
kv_cache_dim if kv_cache_dim != kv_lora_rank + qk_rope_head_dim else None
|
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
|
# num head == 1 and head dim == 128 for index_k in DSA
|
||||||
assert index_head_dim == 128
|
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 _is_hip:
|
||||||
if aiter_can_use_preshuffle_paged_mqa():
|
if aiter_can_use_preshuffle_paged_mqa():
|
||||||
assert (
|
assert (
|
||||||
@@ -3180,6 +3188,10 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
|||||||
if self.custom_mem_pool
|
if self.custom_mem_pool
|
||||||
else nullcontext()
|
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 = [
|
self.index_k_with_scale_buffer = [
|
||||||
torch.zeros(
|
torch.zeros(
|
||||||
# Layout:
|
# Layout:
|
||||||
@@ -3188,17 +3200,11 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
|||||||
# data: for page i,
|
# data: for page i,
|
||||||
# * buf[i, :page_size * head_dim] for fp8 data
|
# * buf[i, :page_size * head_dim] for fp8 data
|
||||||
# * buf[i, page_size * head_dim:].view(float32) for scale
|
# * buf[i, page_size * head_dim:].view(float32) for scale
|
||||||
(
|
(0 if self.skip_topk_layers[i] else num_pages, cols),
|
||||||
(index_buf_size + page_size + 1) // self.page_size,
|
|
||||||
self.page_size
|
|
||||||
* (
|
|
||||||
index_head_dim + index_head_dim // self.quant_block_size * 4
|
|
||||||
),
|
|
||||||
),
|
|
||||||
dtype=self.index_k_with_scale_buffer_dtype,
|
dtype=self.index_k_with_scale_buffer_dtype,
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
for _ in range(layer_num)
|
for i in range(layer_num)
|
||||||
]
|
]
|
||||||
self._finalize_allocation_log(size)
|
self._finalize_allocation_log(size)
|
||||||
|
|
||||||
@@ -3215,7 +3221,9 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
|||||||
|
|
||||||
tgt_loc_flat = tgt_loc.view(-1).long()
|
tgt_loc_flat = tgt_loc.view(-1).long()
|
||||||
src_loc_flat = src_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]
|
index_k[tgt_loc_flat] = index_k[src_loc_flat]
|
||||||
|
|
||||||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
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
|
chunk_size = self.cpu_offloading_chunk_size
|
||||||
page_chunk_size = max(1, chunk_size // self.page_size)
|
page_chunk_size = max(1, chunk_size // self.page_size)
|
||||||
for layer_id in range(self.layer_num):
|
for layer_id in range(self.layer_num):
|
||||||
|
if self.skip_topk_layers[layer_id]:
|
||||||
|
index_k_cpu.append([])
|
||||||
|
continue
|
||||||
index_k_cpu.append([])
|
index_k_cpu.append([])
|
||||||
for i in range(0, len(page_indices), page_chunk_size):
|
for i in range(0, len(page_indices), page_chunk_size):
|
||||||
chunk_page_indices = page_indices[i : i + 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
|
chunk_size = self.cpu_offloading_chunk_size
|
||||||
page_chunk_size = max(1, chunk_size // self.page_size)
|
page_chunk_size = max(1, chunk_size // self.page_size)
|
||||||
for layer_id in range(self.layer_num):
|
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):
|
for i in range(0, len(page_indices), page_chunk_size):
|
||||||
chunk_page_indices = page_indices[i : i + page_chunk_size]
|
chunk_page_indices = page_indices[i : i + page_chunk_size]
|
||||||
idx_cpu = index_k_cpu[layer_id][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)
|
self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num)
|
||||||
]
|
]
|
||||||
item_lens = [
|
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
|
return data_ptrs, data_lens, item_lens
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
|
dsa_layer_skips_topk,
|
||||||
get_dsa_index_head_dim,
|
get_dsa_index_head_dim,
|
||||||
get_minimax_sparse_attention_config,
|
get_minimax_sparse_attention_config,
|
||||||
get_minimax_sparse_disable_value_layer_ids,
|
get_minimax_sparse_disable_value_layer_ids,
|
||||||
@@ -947,6 +948,11 @@ class ModelRunnerKVCacheMixin:
|
|||||||
pool_kwargs["host_to_device_ratio"] = parse_hisparse_config(
|
pool_kwargs["host_to_device_ratio"] = parse_hisparse_config(
|
||||||
self.server_args
|
self.server_args
|
||||||
).host_to_device_ratio
|
).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.token_to_kv_pool = PoolCls(
|
||||||
self.max_total_num_tokens,
|
self.max_total_num_tokens,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
|
dsa_layer_skips_topk,
|
||||||
get_dsa_index_head_dim,
|
get_dsa_index_head_dim,
|
||||||
get_minimax_sparse_attention_config,
|
get_minimax_sparse_attention_config,
|
||||||
get_minimax_sparse_disable_value_layer_ids,
|
get_minimax_sparse_disable_value_layer_ids,
|
||||||
@@ -209,7 +210,15 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
element_size = torch._utils._element_size(
|
element_size = torch._utils._element_size(
|
||||||
DSATokenToKVPool.index_k_with_scale_buffer_dtype
|
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):
|
elif is_minimax_sparse(model_config.hf_config):
|
||||||
# Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool
|
# Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool
|
||||||
# (sparse-only, single-head; kv layers store K+V, k-only layers store K).
|
# (sparse-only, single-head; kv layers store K+V, k-only layers store K).
|
||||||
|
|||||||
Reference in New Issue
Block a user