[Qwen3-next] support mamba radix cache for overlap scheduler (#14792)
This commit is contained in:
@@ -238,8 +238,10 @@ 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, it will be safe when it is larger than 2 (yizhang2077)
|
||||
# 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__)
|
||||
|
||||
@@ -1446,14 +1448,9 @@ class ModelRunner:
|
||||
server_args = self.server_args
|
||||
assert config is not None
|
||||
|
||||
speculativa_ratio = (
|
||||
0
|
||||
if server_args.speculative_num_draft_tokens is None
|
||||
else server_args.speculative_num_draft_tokens
|
||||
)
|
||||
if (
|
||||
server_args.disable_radix_cache
|
||||
or config.mamba2_cache_params.mamba_cache_per_req == 0
|
||||
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:
|
||||
@@ -1461,7 +1458,25 @@ class ModelRunner:
|
||||
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
|
||||
@@ -1475,21 +1490,22 @@ class ModelRunner:
|
||||
server_args.max_mamba_cache_size = int(
|
||||
(mamba_state_memory_raw * (1 << 30))
|
||||
// config.mamba2_cache_params.mamba_cache_per_req
|
||||
// (1 + speculativa_ratio)
|
||||
)
|
||||
|
||||
if self.hybrid_gdn_config is not None:
|
||||
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
|
||||
server_args.dp_size if server_args.enable_dp_attention else 1
|
||||
)
|
||||
mamba_state_memory = (
|
||||
server_args.max_mamba_cache_size
|
||||
* config.mamba2_cache_params.mamba_cache_per_req
|
||||
* (1 + speculativa_ratio)
|
||||
/ (1 << 30)
|
||||
)
|
||||
return total_rest_memory - mamba_state_memory
|
||||
|
||||
@property
|
||||
def qwen3_next_config(self):
|
||||
config = self.model_config.hf_config
|
||||
if isinstance(config, Qwen3NextConfig):
|
||||
return config
|
||||
return None
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
config = self.model_config.hf_config
|
||||
@@ -1683,11 +1699,18 @@ class ModelRunner:
|
||||
)
|
||||
|
||||
if self.mambaish_config is not None:
|
||||
ratio = (
|
||||
MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO
|
||||
if not self.server_args.disable_radix_cache
|
||||
else 1
|
||||
)
|
||||
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
|
||||
)
|
||||
@@ -1789,11 +1812,13 @@ class ModelRunner:
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user