[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
|
server_args = self.server_args
|
||||||
assert config is not None
|
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 (
|
if (
|
||||||
server_args.disable_radix_cache
|
server_args.disable_radix_cache
|
||||||
or server_args.max_mamba_cache_size is not None
|
or server_args.max_mamba_cache_size is not None
|
||||||
@@ -160,19 +177,6 @@ class ModelRunnerKVCacheMixin:
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
assert config.mamba2_cache_params.mamba_cache_per_req > 0
|
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
|
# allocate the memory based on the ratio between mamba state memory vs. full kv cache memory
|
||||||
# solve the equations:
|
# solve the equations:
|
||||||
@@ -284,13 +288,11 @@ class ModelRunnerKVCacheMixin:
|
|||||||
|
|
||||||
if self.mambaish_config is not None:
|
if self.mambaish_config is not None:
|
||||||
additional_ratio = 0
|
additional_ratio = 0
|
||||||
if (
|
if self.server_args.enable_mamba_extra_buffer():
|
||||||
self.server_args.enable_mamba_extra_buffer()
|
if not self.spec_algorithm.is_none():
|
||||||
and not self.spec_algorithm.is_none()
|
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
|
||||||
):
|
else:
|
||||||
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
|
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
|
||||||
else:
|
|
||||||
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
|
|
||||||
if self.server_args.disable_radix_cache:
|
if self.server_args.disable_radix_cache:
|
||||||
ratio = 1
|
ratio = 1
|
||||||
else:
|
else:
|
||||||
@@ -298,6 +300,14 @@ class ModelRunnerKVCacheMixin:
|
|||||||
max_num_reqs = min(
|
max_num_reqs = min(
|
||||||
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
|
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 not self.spec_algorithm.is_none():
|
||||||
if self.is_draft_worker:
|
if self.is_draft_worker:
|
||||||
|
|||||||
Reference in New Issue
Block a user