Introduce ModelRunnerKVCacheMixin to simplify the code. (#15821)
This commit is contained in:
@@ -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
|
||||
|
||||
663
python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py
Normal file
663
python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py
Normal 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"
|
||||
)
|
||||
Reference in New Issue
Block a user