Introduce ModelRunnerKVCacheMixin to simplify the code. (#15821)

This commit is contained in:
Liangsheng Yin
2025-12-26 11:17:33 +08:00
committed by GitHub
parent a91e072f33
commit cb1812954a
2 changed files with 670 additions and 646 deletions

View File

@@ -41,13 +41,7 @@ from sglang.srt.configs import (
)
from sglang.srt.configs.device_config import DeviceConfig
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
from sglang.srt.configs.model_config import (
AttentionArch,
ModelConfig,
ModelImpl,
get_nsa_index_head_dim,
is_deepseek_nsa,
)
from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
@@ -91,7 +85,6 @@ from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
get_attention_tp_group,
get_attention_tp_size,
initialize_dp_attention,
set_dp_buffer_len,
set_is_extend_in_batch,
@@ -109,24 +102,8 @@ from sglang.srt.layers.sampler import create_sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.mem_cache.allocator import (
BaseTokenToKVPoolAllocator,
PagedTokenToKVPoolAllocator,
SWATokenToKVPoolAllocator,
TokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.memory_pool import (
DoubleSparseTokenToKVPool,
HybridLinearKVPool,
HybridReqToTokenPool,
MHATokenToKVPool,
MHATokenToKVPoolFP4,
MLATokenToKVPool,
MLATokenToKVPoolFP4,
NSATokenToKVPool,
ReqToTokenPool,
SWAKVPool,
)
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
from sglang.srt.model_executor.cuda_graph_runner import (
CudaGraphRunner,
@@ -140,6 +117,9 @@ from sglang.srt.model_executor.forward_batch_info import (
)
from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
ModelRunnerKVCacheMixin,
)
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
PiecewiseCudaGraphRunner,
)
@@ -167,7 +147,6 @@ from sglang.srt.utils import (
get_cpu_ids_by_node,
get_local_ip_auto,
init_custom_process_group,
is_float4_e2m1fn_x2,
is_hip,
is_npu,
log_info_on_rank0,
@@ -245,10 +224,6 @@ def add_chunked_prefix_cache_attention_backend(backend_name):
# Detect stragger ranks in model loading
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
# the ratio of mamba cache pool size to max_running_requests
MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO = 3
MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP = 2
MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP = 1
logger = logging.getLogger(__name__)
@@ -280,7 +255,7 @@ class ModelRunnerOutput:
expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None
class ModelRunner:
class ModelRunner(ModelRunnerKVCacheMixin):
"""ModelRunner runs the forward passes of the models."""
def __init__(
@@ -1488,159 +1463,6 @@ class ModelRunner:
return result
def get_cell_size_per_token(self, num_layers: int) -> int:
kv_size = torch._utils._element_size(self.kv_cache_dtype)
if self.use_mla_backend:
cell_size = (
(self.model_config.kv_lora_rank + self.model_config.qk_rope_head_dim)
* num_layers
* kv_size
)
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
# kv_scale_buffer
scale_block_size = 16
cell_size = (cell_size // 2) + (
(
(
self.model_config.kv_lora_rank
+ self.model_config.qk_rope_head_dim
)
// scale_block_size
)
* num_layers
* kv_size
)
# Add indexer KV cache overhead for NSA models (DeepSeek V3.2)
if is_deepseek_nsa(self.model_config.hf_config):
index_head_dim = get_nsa_index_head_dim(self.model_config.hf_config)
indexer_size_per_token = (
index_head_dim
+ index_head_dim // NSATokenToKVPool.quant_block_size * 4
)
element_size = torch._utils._element_size(
NSATokenToKVPool.index_k_with_scale_buffer_dtype
)
cell_size += indexer_size_per_token * num_layers * element_size
else:
cell_size = (
self.model_config.get_num_kv_heads(get_attention_tp_size())
* (self.model_config.head_dim + self.model_config.v_head_dim)
* num_layers
* kv_size
)
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
# kv_scale_buffer
scale_block_size = 16
n = self.model_config.get_num_kv_heads(get_attention_tp_size())
k = self.model_config.head_dim
cell_size = (cell_size // 2) + (
(n * k * num_layers * 2 * kv_size) // scale_block_size
)
if self.model_config.hf_config.architectures[0] == "MiMoV2FlashForCausalLM":
cell_size += (
self.model_config.get_swa_num_kv_heads(get_attention_tp_size())
* (
self.model_config.hf_text_config.swa_head_dim
+ self.model_config.hf_text_config.swa_v_head_dim
)
* len(self.model_config.swa_attention_layer_ids)
* kv_size
)
return cell_size
def profile_max_num_token(self, total_gpu_memory: int):
available_gpu_memory = get_available_gpu_memory(
self.device,
self.gpu_id,
distributed=get_world_group().world_size > 1,
cpu_group=get_world_group().cpu_group,
)
# Get the number of layers used for KV cache calculation
if self.is_draft_worker:
num_layers = getattr(
self.model_config.hf_config,
"num_nextn_predict_layers",
self.num_effective_layers,
)
elif mambaish := self.mambaish_config:
num_layers = len(mambaish.full_attention_layer_ids)
elif self.model_config.full_attention_layer_ids:
num_layers = len(self.model_config.full_attention_layer_ids)
else:
num_layers = self.num_effective_layers
cell_size = self.get_cell_size_per_token(num_layers)
rest_memory = available_gpu_memory - total_gpu_memory * (
1 - self.mem_fraction_static
)
if self.mambaish_config is not None:
rest_memory = self.handle_max_mamba_cache(rest_memory)
logger.info(f"The available memory for KV cache is {rest_memory:.2f} GB.")
return int(rest_memory * (1 << 30)) // cell_size
def handle_max_mamba_cache(self, total_rest_memory):
config = self.mambaish_config
server_args = self.server_args
assert config is not None
if (
server_args.disable_radix_cache
or server_args.max_mamba_cache_size is not None
):
# with disable radix cache, sets the max_mamba_cache_size based on the max_running_requests
if server_args.max_mamba_cache_size is None:
if server_args.max_running_requests is not None:
server_args.max_mamba_cache_size = server_args.max_running_requests
else:
server_args.max_mamba_cache_size = 512
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
server_args.dp_size if server_args.enable_dp_attention else 1
)
else:
assert config.mamba2_cache_params.mamba_cache_per_req > 0
# reserve the memory for the intermediate mamba states used for spec dec
if not self.spec_algorithm.is_none():
assert server_args.speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None
mamba_state_intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req
* server_args.max_running_requests
* server_args.speculative_num_draft_tokens
)
total_rest_memory = total_rest_memory - (
mamba_state_intermediate_size / (1 << 30)
)
# allocate the memory based on the ratio between mamba state memory vs. full kv cache memory
# solve the equations:
# 1. mamba_state_memory + full_kv_cache_memory == total_rest_memory
# 2. mamba_state_memory / full_kv_cache_memory == server_args.mamba_full_memory_ratio
mamba_state_memory_raw = (
total_rest_memory
* server_args.mamba_full_memory_ratio
/ (1 + server_args.mamba_full_memory_ratio)
)
# calculate the max_mamba_cache_size based on the given total mamba memory
server_args.max_mamba_cache_size = int(
(mamba_state_memory_raw * (1 << 30))
// config.mamba2_cache_params.mamba_cache_per_req
)
mamba_state_memory = (
server_args.max_mamba_cache_size
* config.mamba2_cache_params.mamba_cache_per_req
/ (1 << 30)
)
return total_rest_memory - mamba_state_memory
@property
def qwen3_next_config(self):
config = self.model_config.hf_config
@@ -1683,76 +1505,6 @@ class ModelRunner:
def mambaish_config(self):
return self.mamba2_config or self.hybrid_gdn_config or self.kimi_linear_config
def set_num_token_hybrid(self):
page_size = self.server_args.page_size
if (
"Llama4ForConditionalGeneration"
in self.model_config.hf_config.architectures
):
temp_ratio = (
(1 - self.is_hybrid_swa)
+ self.is_hybrid_swa
* self.attention_chunk_size
/ self.model_config.context_len
)
self.swa_max_total_num_tokens = (
4 * self.max_total_num_tokens * temp_ratio // (3 * temp_ratio + 1)
)
self.full_max_total_num_tokens = (
4 * self.max_total_num_tokens
- 12 * self.max_total_num_tokens * temp_ratio // (3 * temp_ratio + 1)
)
self.swa_max_total_num_tokens = (
self.swa_max_total_num_tokens // page_size * page_size
)
self.full_max_total_num_tokens = (
self.full_max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.full_max_total_num_tokens
elif "MiMoV2MTP" in self.model_config.hf_config.architectures:
assert self.is_draft_worker
# MiMoV2MTP uses SWA, so set full KV cache to 0
self.full_max_total_num_tokens = 0
self.swa_max_total_num_tokens = (
self.max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.swa_max_total_num_tokens
else:
assert self.sliding_window_size is not None and self.sliding_window_size > 0
full_layers_num = len(self.model_config.full_attention_layer_ids)
swa_layers_num = len(self.model_config.swa_attention_layer_ids)
# Algorithm:
# Existing max_total_num_tokens is per layer and assume all layers have the same number of tokens.
# - Find total # of tokens available across layers.
# - Calculate full_max_total_num_tokens and swa_max_total_num_tokens based on the given swa_full_tokens_ratio.
total_tokens = (
self.max_total_num_tokens * self.model_config.num_hidden_layers
)
swa_full_tokens_ratio = self.server_args.swa_full_tokens_ratio
# Solve the equations:
# 1. swa_max_total_num_tokens * swa_layers_num + full_max_total_num_tokens * full_layers_num == total_tokens
# 2. full_max_total_num_tokens * swa_full_tokens_ratio == swa_max_total_num_tokens
denominator = swa_full_tokens_ratio * swa_layers_num + full_layers_num
self.full_max_total_num_tokens = int(total_tokens / denominator)
self.swa_max_total_num_tokens = int(
self.full_max_total_num_tokens * swa_full_tokens_ratio
)
self.full_max_total_num_tokens = (
self.full_max_total_num_tokens // page_size * page_size
)
self.swa_max_total_num_tokens = (
self.swa_max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.full_max_total_num_tokens
logger.info(
f"Use sliding window memory pool. full_layer_tokens={self.full_max_total_num_tokens}, swa_layer_tokens={self.swa_max_total_num_tokens}"
)
def can_run_piecewise_cuda_graph(self):
if self.server_args.enable_torch_compile:
log_info_on_rank0(
@@ -1818,397 +1570,6 @@ class ModelRunner:
log_info_on_rank0(logger, f"Using KV cache dtype: {self.kv_cache_dtype}")
def init_memory_pool(self, total_gpu_memory: int, server_args: ServerArgs):
max_num_reqs = server_args.max_running_requests
max_total_tokens = server_args.max_total_tokens
self.max_total_num_tokens = self.profile_max_num_token(total_gpu_memory)
if max_num_reqs is None:
max_num_reqs = min(
max(
int(
self.max_total_num_tokens / self.model_config.context_len * 512
),
2048,
),
4096,
)
if self.mambaish_config is not None:
additional_ratio = 0
if (
self.server_args.enable_mamba_extra_buffer()
and not self.spec_algorithm.is_none()
):
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
else:
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
if self.server_args.disable_radix_cache:
ratio = 1
else:
ratio = MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO + additional_ratio
max_num_reqs = min(
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
)
if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone():
if self.is_draft_worker:
self.max_total_num_tokens = self.server_args.draft_runner_cache_size
max_num_reqs = self.server_args.max_num_reqs
else:
self.server_args.draft_runner_cache_size = self.max_total_num_tokens
self.server_args.max_num_reqs = max_num_reqs
if max_total_tokens is not None:
if max_total_tokens > self.max_total_num_tokens:
logging.warning(
f"max_total_tokens={max_total_tokens} is larger than the profiled value "
f"{self.max_total_num_tokens}. "
f"Use the profiled value instead."
)
self.max_total_num_tokens = min(self.max_total_num_tokens, max_total_tokens)
self.max_total_num_tokens = (
self.max_total_num_tokens
// self.server_args.page_size
* self.server_args.page_size
)
# different pp rank may have different num of layers, so we need to reduce the max_total_num_tokens
if self.pp_size > 1:
tensor = torch.tensor(self.max_total_num_tokens, dtype=torch.int64)
torch.distributed.all_reduce(
tensor,
op=torch.distributed.ReduceOp.MIN,
group=get_world_group().cpu_group,
)
self.max_total_num_tokens = tensor.item()
# create token size for hybrid cache
if self.is_hybrid_swa:
self.set_num_token_hybrid()
if self.max_total_num_tokens <= 0:
raise RuntimeError(
f"Not enough memory. Please try to increase --mem-fraction-static. "
f"Current value: {self.server_args.mem_fraction_static=}"
)
# Initialize req_to_token_pool
if self.req_to_token_pool is None:
# FIXME(lsyin): this is the temporary fix for the context length issue when using speculative decoding
extra_max_context_len = 4
if self.server_args.speculative_num_draft_tokens is not None:
extra_max_context_len += self.server_args.speculative_num_draft_tokens
if self.server_args.disaggregation_mode == "decode":
from sglang.srt.disaggregation.decode import (
DecodeReqToTokenPool,
HybridMambaDecodeReqToTokenPool,
)
# subscribe memory for pre-allocated requests
# if max_num_reqs <= 32, we pre-allocate 2x requests
pre_alloc_size = max_num_reqs * 2 if max_num_reqs <= 32 else 0
if config := self.mambaish_config:
self.req_to_token_pool = HybridMambaDecodeReqToTokenPool(
size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=config.mamba2_cache_params,
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
pre_alloc_size=pre_alloc_size,
)
else:
self.req_to_token_pool = DecodeReqToTokenPool(
size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
pre_alloc_size=pre_alloc_size,
)
elif config := self.mambaish_config:
self.req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=self.server_args.max_mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=config.mamba2_cache_params,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
)
else:
self.req_to_token_pool = ReqToTokenPool(
size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
)
else:
# Draft worker shares req_to_token_pool with the target worker.
assert self.is_draft_worker
# Initialize token_to_kv_pool
is_nsa_model = is_deepseek_nsa(self.model_config.hf_config)
if self.server_args.attention_backend == "ascend":
if self.use_mla_backend:
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMLATokenToKVPool,
)
self.token_to_kv_pool = NPUMLATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
index_head_dim=self.model_config.index_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
else:
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
self.token_to_kv_pool = NPUMHATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
elif self.use_mla_backend and is_nsa_model:
self.token_to_kv_pool = NSATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
index_head_dim=get_nsa_index_head_dim(self.model_config.hf_config),
)
elif self.use_mla_backend and not self.mambaish_config:
assert not is_nsa_model
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
self.token_to_kv_pool = MLATokenToKVPoolFP4(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
else:
self.token_to_kv_pool = MLATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
elif self.server_args.enable_double_sparsity:
self.token_to_kv_pool = DoubleSparseTokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_attention_tp_size()),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
heavy_channel_num=self.server_args.ds_heavy_channel_num,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
else:
if self.is_hybrid_swa:
kwargs = {}
if self.is_hybrid_swa_compress:
kwargs = {
"swa_head_num": max(
1,
self.model_config.hf_text_config.swa_num_key_value_heads
// get_attention_tp_size(),
),
"swa_head_dim": self.model_config.hf_text_config.swa_head_dim,
"swa_v_head_dim": self.model_config.hf_text_config.swa_v_head_dim,
"v_head_dim": self.model_config.hf_text_config.v_head_dim,
}
self.token_to_kv_pool = SWAKVPool(
size=self.full_max_total_num_tokens,
size_swa=self.swa_max_total_num_tokens,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
swa_attention_layer_ids=self.model_config.swa_attention_layer_ids,
full_attention_layer_ids=self.model_config.full_attention_layer_ids,
enable_kvcache_transpose=False,
device=self.device,
**kwargs,
)
elif config := self.mambaish_config:
extra_args = {}
if self.use_mla_backend:
extra_args = {
"kv_lora_rank": self.model_config.kv_lora_rank,
"qk_rope_head_dim": self.model_config.qk_rope_head_dim,
}
self.token_to_kv_pool = HybridLinearKVPool(
page_size=self.page_size,
size=self.max_total_num_tokens,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
# if draft worker, we only need 1 attention layer's kv pool
full_attention_layer_ids=(
[0] if self.is_draft_worker else config.full_attention_layer_ids
),
enable_kvcache_transpose=False,
device=self.device,
mamba_pool=self.req_to_token_pool.mamba_pool,
enable_memory_saver=self.server_args.enable_memory_saver,
use_mla=self.use_mla_backend,
**extra_args,
)
else:
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
self.token_to_kv_pool = MHATokenToKVPoolFP4(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(
self.server_args.speculative_algorithm is not None
),
)
else:
self.token_to_kv_pool = MHATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(
self.server_args.speculative_algorithm is not None
),
)
# Initialize token_to_kv_pool_allocator
need_sort = self.server_args.disaggregation_mode in ("decode", "prefill")
if self.token_to_kv_pool_allocator is None:
if _is_npu and (
self.server_args.attention_backend == "ascend"
or self.hybrid_gdn_config is not None
):
from sglang.srt.hardware_backend.npu.allocator_npu import (
NPUPagedTokenToKVPoolAllocator,
)
self.token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
if self.page_size == 1:
if self.is_hybrid_swa:
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
self.full_max_total_num_tokens,
self.swa_max_total_num_tokens,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
self.token_to_kv_pool_allocator = TokenToKVPoolAllocator(
self.max_total_num_tokens,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
assert not self.is_hybrid_swa
self.token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
assert self.is_draft_worker
if self.is_hybrid_swa:
assert (
self.token_to_kv_pool_allocator.__class__
== SWATokenToKVPoolAllocator
)
self.token_to_kv_pool.full_to_swa_index_mapping = (
self.token_to_kv_pool_allocator.full_to_swa_index_mapping
)
logger.info(
f"Memory pool end. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
)
def init_cublas(self):
"""We need to run a small matmul to init cublas. Otherwise, it will raise some errors later."""
dtype = torch.float16

View File

@@ -0,0 +1,663 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
import torch
from sglang.srt.configs.model_config import get_nsa_index_head_dim, is_deepseek_nsa
from sglang.srt.distributed.parallel_state import get_world_group
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.mem_cache.allocator import (
PagedTokenToKVPoolAllocator,
SWATokenToKVPoolAllocator,
TokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.memory_pool import (
DoubleSparseTokenToKVPool,
HybridLinearKVPool,
HybridReqToTokenPool,
MHATokenToKVPool,
MHATokenToKVPoolFP4,
MLATokenToKVPool,
MLATokenToKVPoolFP4,
NSATokenToKVPool,
ReqToTokenPool,
SWAKVPool,
)
from sglang.srt.utils.common import (
get_available_gpu_memory,
is_float4_e2m1fn_x2,
is_npu,
)
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.server_args import ServerArgs
# the ratio of mamba cache pool size to max_running_requests
MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO = 3
MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP = 2
MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP = 1
logger = logging.getLogger(__name__)
_is_npu = is_npu()
class ModelRunnerKVCacheMixin:
def get_cell_size_per_token(self: ModelRunner, num_layers: int) -> int:
kv_size = torch._utils._element_size(self.kv_cache_dtype)
if self.use_mla_backend:
cell_size = (
(self.model_config.kv_lora_rank + self.model_config.qk_rope_head_dim)
* num_layers
* kv_size
)
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
# kv_scale_buffer
scale_block_size = 16
cell_size = (cell_size // 2) + (
(
(
self.model_config.kv_lora_rank
+ self.model_config.qk_rope_head_dim
)
// scale_block_size
)
* num_layers
* kv_size
)
# Add indexer KV cache overhead for NSA models (DeepSeek V3.2)
if is_deepseek_nsa(self.model_config.hf_config):
index_head_dim = get_nsa_index_head_dim(self.model_config.hf_config)
indexer_size_per_token = (
index_head_dim
+ index_head_dim // NSATokenToKVPool.quant_block_size * 4
)
element_size = torch._utils._element_size(
NSATokenToKVPool.index_k_with_scale_buffer_dtype
)
cell_size += indexer_size_per_token * num_layers * element_size
else:
cell_size = (
self.model_config.get_num_kv_heads(get_attention_tp_size())
* (self.model_config.head_dim + self.model_config.v_head_dim)
* num_layers
* kv_size
)
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
# kv_scale_buffer
scale_block_size = 16
n = self.model_config.get_num_kv_heads(get_attention_tp_size())
k = self.model_config.head_dim
cell_size = (cell_size // 2) + (
(n * k * num_layers * 2 * kv_size) // scale_block_size
)
if "MiMoV2FlashForCausalLM" in self.model_config.hf_config.architectures:
cell_size += (
self.model_config.get_swa_num_kv_heads(get_attention_tp_size())
* (
self.model_config.hf_text_config.swa_head_dim
+ self.model_config.hf_text_config.swa_v_head_dim
)
* len(self.model_config.swa_attention_layer_ids)
* kv_size
)
return cell_size
def profile_max_num_token(self: ModelRunner, total_gpu_memory: int):
available_gpu_memory = get_available_gpu_memory(
self.device,
self.gpu_id,
distributed=get_world_group().world_size > 1,
cpu_group=get_world_group().cpu_group,
)
# Get the number of layers used for KV cache calculation
if self.is_draft_worker:
num_layers = getattr(
self.model_config.hf_config,
"num_nextn_predict_layers",
self.num_effective_layers,
)
elif mambaish := self.mambaish_config:
num_layers = len(mambaish.full_attention_layer_ids)
elif self.model_config.full_attention_layer_ids:
num_layers = len(self.model_config.full_attention_layer_ids)
else:
num_layers = self.num_effective_layers
cell_size = self.get_cell_size_per_token(num_layers)
rest_memory = available_gpu_memory - total_gpu_memory * (
1 - self.mem_fraction_static
)
if self.mambaish_config is not None:
rest_memory = self.handle_max_mamba_cache(rest_memory)
logger.info(f"The available memory for KV cache is {rest_memory:.2f} GB.")
return int(rest_memory * (1 << 30)) // cell_size
def handle_max_mamba_cache(self: ModelRunner, total_rest_memory):
config = self.mambaish_config
server_args = self.server_args
assert config is not None
if (
server_args.disable_radix_cache
or server_args.max_mamba_cache_size is not None
):
# with disable radix cache, sets the max_mamba_cache_size based on the max_running_requests
if server_args.max_mamba_cache_size is None:
if server_args.max_running_requests is not None:
server_args.max_mamba_cache_size = server_args.max_running_requests
else:
server_args.max_mamba_cache_size = 512
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
server_args.dp_size if server_args.enable_dp_attention else 1
)
else:
assert config.mamba2_cache_params.mamba_cache_per_req > 0
# reserve the memory for the intermediate mamba states used for spec dec
if not self.spec_algorithm.is_none():
assert server_args.speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None
mamba_state_intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req
* server_args.max_running_requests
* server_args.speculative_num_draft_tokens
)
total_rest_memory = total_rest_memory - (
mamba_state_intermediate_size / (1 << 30)
)
# allocate the memory based on the ratio between mamba state memory vs. full kv cache memory
# solve the equations:
# 1. mamba_state_memory + full_kv_cache_memory == total_rest_memory
# 2. mamba_state_memory / full_kv_cache_memory == server_args.mamba_full_memory_ratio
mamba_state_memory_raw = (
total_rest_memory
* server_args.mamba_full_memory_ratio
/ (1 + server_args.mamba_full_memory_ratio)
)
# calculate the max_mamba_cache_size based on the given total mamba memory
server_args.max_mamba_cache_size = int(
(mamba_state_memory_raw * (1 << 30))
// config.mamba2_cache_params.mamba_cache_per_req
)
mamba_state_memory = (
server_args.max_mamba_cache_size
* config.mamba2_cache_params.mamba_cache_per_req
/ (1 << 30)
)
return total_rest_memory - mamba_state_memory
def set_num_tokens_hybrid_swa(self: ModelRunner):
page_size = self.server_args.page_size
if (
"Llama4ForConditionalGeneration"
in self.model_config.hf_config.architectures
):
temp_ratio = (
(1 - self.is_hybrid_swa)
+ self.is_hybrid_swa
* self.attention_chunk_size
/ self.model_config.context_len
)
self.swa_max_total_num_tokens = (
4 * self.max_total_num_tokens * temp_ratio // (3 * temp_ratio + 1)
)
self.full_max_total_num_tokens = (
4 * self.max_total_num_tokens
- 12 * self.max_total_num_tokens * temp_ratio // (3 * temp_ratio + 1)
)
self.swa_max_total_num_tokens = (
self.swa_max_total_num_tokens // page_size * page_size
)
self.full_max_total_num_tokens = (
self.full_max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.full_max_total_num_tokens
elif "MiMoV2MTP" in self.model_config.hf_config.architectures:
assert self.is_draft_worker
# MiMoV2MTP uses SWA, so set full KV cache to 0
self.full_max_total_num_tokens = 0
self.swa_max_total_num_tokens = (
self.max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.swa_max_total_num_tokens
else:
assert self.sliding_window_size is not None and self.sliding_window_size > 0
full_layers_num = len(self.model_config.full_attention_layer_ids)
swa_layers_num = len(self.model_config.swa_attention_layer_ids)
# Algorithm:
# Existing max_total_num_tokens is per layer and assume all layers have the same number of tokens.
# - Find total # of tokens available across layers.
# - Calculate full_max_total_num_tokens and swa_max_total_num_tokens based on the given swa_full_tokens_ratio.
total_tokens = (
self.max_total_num_tokens * self.model_config.num_hidden_layers
)
swa_full_tokens_ratio = self.server_args.swa_full_tokens_ratio
# Solve the equations:
# 1. swa_max_total_num_tokens * swa_layers_num + full_max_total_num_tokens * full_layers_num == total_tokens
# 2. full_max_total_num_tokens * swa_full_tokens_ratio == swa_max_total_num_tokens
denominator = swa_full_tokens_ratio * swa_layers_num + full_layers_num
self.full_max_total_num_tokens = int(total_tokens / denominator)
self.swa_max_total_num_tokens = int(
self.full_max_total_num_tokens * swa_full_tokens_ratio
)
self.full_max_total_num_tokens = (
self.full_max_total_num_tokens // page_size * page_size
)
self.swa_max_total_num_tokens = (
self.swa_max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.full_max_total_num_tokens
logger.info(
f"Use sliding window memory pool. full_layer_tokens={self.full_max_total_num_tokens}, swa_layer_tokens={self.swa_max_total_num_tokens}"
)
def init_memory_pool(
self: ModelRunner, total_gpu_memory: int, server_args: ServerArgs
):
max_num_reqs = server_args.max_running_requests
max_total_tokens = server_args.max_total_tokens
self.max_total_num_tokens = self.profile_max_num_token(total_gpu_memory)
if max_num_reqs is None:
max_num_reqs = min(
max(
int(
self.max_total_num_tokens / self.model_config.context_len * 512
),
2048,
),
4096,
)
if self.mambaish_config is not None:
additional_ratio = 0
if (
self.server_args.enable_mamba_extra_buffer()
and not self.spec_algorithm.is_none()
):
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
else:
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
if self.server_args.disable_radix_cache:
ratio = 1
else:
ratio = MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO + additional_ratio
max_num_reqs = min(
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
)
if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone():
if self.is_draft_worker:
self.max_total_num_tokens = self.server_args.draft_runner_cache_size
max_num_reqs = self.server_args.max_num_reqs
else:
self.server_args.draft_runner_cache_size = self.max_total_num_tokens
self.server_args.max_num_reqs = max_num_reqs
if max_total_tokens is not None:
if max_total_tokens > self.max_total_num_tokens:
logging.warning(
f"max_total_tokens={max_total_tokens} is larger than the profiled value "
f"{self.max_total_num_tokens}. "
f"Use the profiled value instead."
)
self.max_total_num_tokens = min(self.max_total_num_tokens, max_total_tokens)
self.max_total_num_tokens = (
self.max_total_num_tokens
// self.server_args.page_size
* self.server_args.page_size
)
# different pp rank may have different num of layers, so we need to reduce the max_total_num_tokens
if self.pp_size > 1:
tensor = torch.tensor(self.max_total_num_tokens, dtype=torch.int64)
torch.distributed.all_reduce(
tensor,
op=torch.distributed.ReduceOp.MIN,
group=get_world_group().cpu_group,
)
self.max_total_num_tokens = tensor.item()
# create token size for hybrid cache
if self.is_hybrid_swa:
self.set_num_tokens_hybrid_swa()
if self.max_total_num_tokens <= 0:
raise RuntimeError(
f"Not enough memory. Please try to increase --mem-fraction-static. "
f"Current value: {self.server_args.mem_fraction_static=}"
)
# Initialize req_to_token_pool
if self.req_to_token_pool is None:
# FIXME(lsyin): this is the temporary fix for the context length issue when using speculative decoding
extra_max_context_len = 4
if self.server_args.speculative_num_draft_tokens is not None:
extra_max_context_len += self.server_args.speculative_num_draft_tokens
if self.server_args.disaggregation_mode == "decode":
from sglang.srt.disaggregation.decode import (
DecodeReqToTokenPool,
HybridMambaDecodeReqToTokenPool,
)
# subscribe memory for pre-allocated requests
# if max_num_reqs <= 32, we pre-allocate 2x requests
pre_alloc_size = max_num_reqs * 2 if max_num_reqs <= 32 else 0
if config := self.mambaish_config:
self.req_to_token_pool = HybridMambaDecodeReqToTokenPool(
size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=config.mamba2_cache_params,
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
pre_alloc_size=pre_alloc_size,
)
else:
self.req_to_token_pool = DecodeReqToTokenPool(
size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
pre_alloc_size=pre_alloc_size,
)
elif config := self.mambaish_config:
self.req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=self.server_args.max_mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=config.mamba2_cache_params,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
)
else:
self.req_to_token_pool = ReqToTokenPool(
size=max_num_reqs,
max_context_len=self.model_config.context_len
+ extra_max_context_len,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
)
else:
# Draft worker shares req_to_token_pool with the target worker.
assert self.is_draft_worker
# Initialize token_to_kv_pool
is_nsa_model = is_deepseek_nsa(self.model_config.hf_config)
if self.server_args.attention_backend == "ascend":
if self.use_mla_backend:
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMLATokenToKVPool,
)
self.token_to_kv_pool = NPUMLATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
index_head_dim=self.model_config.index_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
else:
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
self.token_to_kv_pool = NPUMHATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
elif self.use_mla_backend and is_nsa_model:
self.token_to_kv_pool = NSATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
index_head_dim=get_nsa_index_head_dim(self.model_config.hf_config),
)
elif self.use_mla_backend and not self.mambaish_config:
assert not is_nsa_model
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
self.token_to_kv_pool = MLATokenToKVPoolFP4(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
else:
self.token_to_kv_pool = MLATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
elif self.server_args.enable_double_sparsity:
self.token_to_kv_pool = DoubleSparseTokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_attention_tp_size()),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
heavy_channel_num=self.server_args.ds_heavy_channel_num,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
else:
if self.is_hybrid_swa:
kwargs = {}
if self.is_hybrid_swa_compress:
kwargs = {
"swa_head_num": max(
1,
self.model_config.hf_text_config.swa_num_key_value_heads
// get_attention_tp_size(),
),
"swa_head_dim": self.model_config.hf_text_config.swa_head_dim,
"swa_v_head_dim": self.model_config.hf_text_config.swa_v_head_dim,
"v_head_dim": self.model_config.hf_text_config.v_head_dim,
}
self.token_to_kv_pool = SWAKVPool(
size=self.full_max_total_num_tokens,
size_swa=self.swa_max_total_num_tokens,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
swa_attention_layer_ids=self.model_config.swa_attention_layer_ids,
full_attention_layer_ids=self.model_config.full_attention_layer_ids,
enable_kvcache_transpose=False,
device=self.device,
**kwargs,
)
elif config := self.mambaish_config:
extra_args = {}
if self.use_mla_backend:
extra_args = {
"kv_lora_rank": self.model_config.kv_lora_rank,
"qk_rope_head_dim": self.model_config.qk_rope_head_dim,
}
self.token_to_kv_pool = HybridLinearKVPool(
page_size=self.page_size,
size=self.max_total_num_tokens,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
# if draft worker, we only need 1 attention layer's kv pool
full_attention_layer_ids=(
[0] if self.is_draft_worker else config.full_attention_layer_ids
),
enable_kvcache_transpose=False,
device=self.device,
mamba_pool=self.req_to_token_pool.mamba_pool,
enable_memory_saver=self.server_args.enable_memory_saver,
use_mla=self.use_mla_backend,
**extra_args,
)
else:
if is_float4_e2m1fn_x2(self.kv_cache_dtype):
self.token_to_kv_pool = MHATokenToKVPoolFP4(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(
self.server_args.speculative_algorithm is not None
),
)
else:
self.token_to_kv_pool = MHATokenToKVPool(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(
self.server_args.speculative_algorithm is not None
),
)
# Initialize token_to_kv_pool_allocator
need_sort = self.server_args.disaggregation_mode in ("decode", "prefill")
if self.token_to_kv_pool_allocator is None:
if _is_npu and (
self.server_args.attention_backend == "ascend"
or self.hybrid_gdn_config is not None
):
from sglang.srt.hardware_backend.npu.allocator_npu import (
NPUPagedTokenToKVPoolAllocator,
)
self.token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
if self.page_size == 1:
if self.is_hybrid_swa:
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
self.full_max_total_num_tokens,
self.swa_max_total_num_tokens,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
self.token_to_kv_pool_allocator = TokenToKVPoolAllocator(
self.max_total_num_tokens,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
assert not self.is_hybrid_swa
self.token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
device=self.device,
kvcache=self.token_to_kv_pool,
need_sort=need_sort,
)
else:
assert self.is_draft_worker
if self.is_hybrid_swa:
assert (
self.token_to_kv_pool_allocator.__class__
== SWATokenToKVPoolAllocator
)
self.token_to_kv_pool.full_to_swa_index_mapping = (
self.token_to_kv_pool_allocator.full_to_swa_index_mapping
)
logger.info(
f"Memory pool end. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
)