[Bugfix] fix some memory computation bugs for qwen3next with mtp (#16138)

This commit is contained in:
Yi Zhang
2026-01-05 23:24:21 +08:00
committed by GitHub
parent 130f60ee78
commit a3914e3b3f

View File

@@ -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: