[Bugfix] fix some memory computation bugs for qwen3next with mtp (#16138)
This commit is contained in:
@@ -145,6 +145,23 @@ class ModelRunnerKVCacheMixin:
|
||||
server_args = self.server_args
|
||||
assert config is not None
|
||||
|
||||
# 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
|
||||
|
||||
max_running_requests = server_args.max_running_requests // (
|
||||
self.dp_size if server_args.enable_dp_attention else 1
|
||||
)
|
||||
mamba_state_intermediate_size = (
|
||||
config.mamba2_cache_params.mamba_cache_per_req
|
||||
* max_running_requests
|
||||
* server_args.speculative_num_draft_tokens
|
||||
)
|
||||
total_rest_memory = total_rest_memory - (
|
||||
mamba_state_intermediate_size / (1 << 30)
|
||||
)
|
||||
|
||||
if (
|
||||
server_args.disable_radix_cache
|
||||
or server_args.max_mamba_cache_size is not None
|
||||
@@ -160,19 +177,6 @@ class ModelRunnerKVCacheMixin:
|
||||
)
|
||||
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:
|
||||
@@ -284,13 +288,11 @@ class ModelRunnerKVCacheMixin:
|
||||
|
||||
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.enable_mamba_extra_buffer():
|
||||
if 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:
|
||||
@@ -298,6 +300,14 @@ class ModelRunnerKVCacheMixin:
|
||||
max_num_reqs = min(
|
||||
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
|
||||
)
|
||||
# for dp attention, we need control the max_num_reqs for speculative decoding mamba space
|
||||
if (
|
||||
not self.spec_algorithm.is_none()
|
||||
and self.server_args.enable_dp_attention
|
||||
):
|
||||
max_num_reqs = min(
|
||||
max_num_reqs, self.server_args.max_running_requests // self.dp_size
|
||||
)
|
||||
|
||||
if not self.spec_algorithm.is_none():
|
||||
if self.is_draft_worker:
|
||||
|
||||
Reference in New Issue
Block a user