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.device_config import DeviceConfig
|
||||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl
|
||||||
AttentionArch,
|
|
||||||
ModelConfig,
|
|
||||||
ModelImpl,
|
|
||||||
get_nsa_index_head_dim,
|
|
||||||
is_deepseek_nsa,
|
|
||||||
)
|
|
||||||
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
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.constants import GPU_MEMORY_TYPE_WEIGHTS
|
||||||
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
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 (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
get_attention_tp_group,
|
get_attention_tp_group,
|
||||||
get_attention_tp_size,
|
|
||||||
initialize_dp_attention,
|
initialize_dp_attention,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
set_is_extend_in_batch,
|
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.layers.torchao_utils import apply_torchao_config_to_model
|
||||||
from sglang.srt.lora.lora_manager import LoRAManager
|
from sglang.srt.lora.lora_manager import LoRAManager
|
||||||
from sglang.srt.lora.lora_registry import LoRARef
|
from sglang.srt.lora.lora_registry import LoRARef
|
||||||
from sglang.srt.mem_cache.allocator import (
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
BaseTokenToKVPoolAllocator,
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
PagedTokenToKVPoolAllocator,
|
|
||||||
SWATokenToKVPoolAllocator,
|
|
||||||
TokenToKVPoolAllocator,
|
|
||||||
)
|
|
||||||
from sglang.srt.mem_cache.memory_pool import (
|
|
||||||
DoubleSparseTokenToKVPool,
|
|
||||||
HybridLinearKVPool,
|
|
||||||
HybridReqToTokenPool,
|
|
||||||
MHATokenToKVPool,
|
|
||||||
MHATokenToKVPoolFP4,
|
|
||||||
MLATokenToKVPool,
|
|
||||||
MLATokenToKVPoolFP4,
|
|
||||||
NSATokenToKVPool,
|
|
||||||
ReqToTokenPool,
|
|
||||||
SWAKVPool,
|
|
||||||
)
|
|
||||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||||
from sglang.srt.model_executor.cuda_graph_runner import (
|
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||||
CudaGraphRunner,
|
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.hook_manager import register_forward_hooks
|
||||||
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
|
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 (
|
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||||
PiecewiseCudaGraphRunner,
|
PiecewiseCudaGraphRunner,
|
||||||
)
|
)
|
||||||
@@ -167,7 +147,6 @@ from sglang.srt.utils import (
|
|||||||
get_cpu_ids_by_node,
|
get_cpu_ids_by_node,
|
||||||
get_local_ip_auto,
|
get_local_ip_auto,
|
||||||
init_custom_process_group,
|
init_custom_process_group,
|
||||||
is_float4_e2m1fn_x2,
|
|
||||||
is_hip,
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
@@ -245,10 +224,6 @@ def add_chunked_prefix_cache_attention_backend(backend_name):
|
|||||||
# Detect stragger ranks in model loading
|
# Detect stragger ranks in model loading
|
||||||
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -280,7 +255,7 @@ class ModelRunnerOutput:
|
|||||||
expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None
|
expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None
|
||||||
|
|
||||||
|
|
||||||
class ModelRunner:
|
class ModelRunner(ModelRunnerKVCacheMixin):
|
||||||
"""ModelRunner runs the forward passes of the models."""
|
"""ModelRunner runs the forward passes of the models."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -1488,159 +1463,6 @@ class ModelRunner:
|
|||||||
|
|
||||||
return result
|
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
|
@property
|
||||||
def qwen3_next_config(self):
|
def qwen3_next_config(self):
|
||||||
config = self.model_config.hf_config
|
config = self.model_config.hf_config
|
||||||
@@ -1683,76 +1505,6 @@ class ModelRunner:
|
|||||||
def mambaish_config(self):
|
def mambaish_config(self):
|
||||||
return self.mamba2_config or self.hybrid_gdn_config or self.kimi_linear_config
|
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):
|
def can_run_piecewise_cuda_graph(self):
|
||||||
if self.server_args.enable_torch_compile:
|
if self.server_args.enable_torch_compile:
|
||||||
log_info_on_rank0(
|
log_info_on_rank0(
|
||||||
@@ -1818,397 +1570,6 @@ class ModelRunner:
|
|||||||
|
|
||||||
log_info_on_rank0(logger, f"Using KV cache dtype: {self.kv_cache_dtype}")
|
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):
|
def init_cublas(self):
|
||||||
"""We need to run a small matmul to init cublas. Otherwise, it will raise some errors later."""
|
"""We need to run a small matmul to init cublas. Otherwise, it will raise some errors later."""
|
||||||
dtype = torch.float16
|
dtype = torch.float16
|
||||||
|
|||||||
@@ -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