[Fix] Reserve the mamba pool's +1 padding slot in the memory budget solve (#32184)

This commit is contained in:
Liangsheng Yin
2026-07-23 14:55:35 -07:00
committed by GitHub
parent 3d0c6bf57f
commit 59ef3b15cc
@@ -1717,7 +1717,7 @@ class KVCacheConfigurator:
max_mamba_cache_size=server_args.max_mamba_cache_size max_mamba_cache_size=server_args.max_mamba_cache_size
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
) )
# Reserve intermediate memory based on capped max_num_reqs # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
if has_spec_dec: if has_spec_dec:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
capped_reqs = min( capped_reqs = min(
@@ -1726,7 +1726,7 @@ class KVCacheConfigurator:
) )
intermediate_size = ( intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req config.mamba2_cache_params.mamba_cache_per_req
* capped_reqs * (capped_reqs + 1)
* server_args.speculative_num_draft_tokens * server_args.speculative_num_draft_tokens
) )
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
@@ -1740,11 +1740,11 @@ class KVCacheConfigurator:
max_mamba_cache_size=server_args.max_running_requests max_mamba_cache_size=server_args.max_running_requests
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
) )
# Reserve intermediate memory based on capped max_num_reqs # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
if has_spec_dec: if has_spec_dec:
intermediate_size = ( intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req config.mamba2_cache_params.mamba_cache_per_req
* server_args.max_mamba_cache_size * (server_args.max_mamba_cache_size + 1)
* server_args.speculative_num_draft_tokens * server_args.speculative_num_draft_tokens
) )
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
@@ -1753,11 +1753,9 @@ class KVCacheConfigurator:
assert config.mamba2_cache_params.mamba_cache_per_req > 0 assert config.mamba2_cache_params.mamba_cache_per_req > 0
per_req = config.mamba2_cache_params.mamba_cache_per_req per_req = config.mamba2_cache_params.mamba_cache_per_req
# Solve jointly for max_mamba_cache_size accounting for intermediate memory. # Solve jointly for max_mamba_cache_size (K), including the pool's
# The mamba budget (from the ratio split) must cover both: # +1 padding slot on both buffers (see memory_pool.py):
# 1. main mamba state: max_mamba_cache_size * per_req # (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes
# 2. intermediate states: (max_mamba_cache_size / ratio) * D * per_req
# So: max_mamba_cache_size * per_req * (1 + D/ratio) = mamba_budget_bytes
mamba_budget = ( mamba_budget = (
total_rest_memory total_rest_memory
* server_args.mamba_full_memory_ratio * server_args.mamba_full_memory_ratio
@@ -1772,7 +1770,8 @@ class KVCacheConfigurator:
server_args.override( server_args.override(
"mamba_pool.memory_budget_spec", "mamba_pool.memory_budget_spec",
max_mamba_cache_size=int( max_mamba_cache_size=int(
mamba_budget_bytes // (per_req * (1 + D / ratio)) (mamba_budget_bytes - per_req * (1 + D))
// (per_req * (1 + D / ratio))
), ),
) )
# Intermediate memory is included in mamba_budget, subtract it # Intermediate memory is included in mamba_budget, subtract it
@@ -1781,12 +1780,12 @@ class KVCacheConfigurator:
server_args.max_running_requests // self.ps.attn_dp_size, server_args.max_running_requests // self.ps.attn_dp_size,
server_args.max_mamba_cache_size // ratio, server_args.max_mamba_cache_size // ratio,
) )
intermediate_size = per_req * capped_reqs * D intermediate_size = per_req * (capped_reqs + 1) * D
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else: else:
server_args.override( server_args.override(
"mamba_pool.memory_budget", "mamba_pool.memory_budget",
max_mamba_cache_size=int(mamba_budget_bytes // per_req), max_mamba_cache_size=int((mamba_budget_bytes - per_req) // per_req),
) )
# Validate: max_mamba_cache_size must be positive after memory allocation. # Validate: max_mamba_cache_size must be positive after memory allocation.
@@ -1805,8 +1804,9 @@ class KVCacheConfigurator:
f"(4) use GPUs with more memory." f"(4) use GPUs with more memory."
) )
# +1: the pool's padding slot
mamba_state_memory = ( mamba_state_memory = (
server_args.max_mamba_cache_size (server_args.max_mamba_cache_size + 1)
* config.mamba2_cache_params.mamba_cache_per_req * config.mamba2_cache_params.mamba_cache_per_req
/ (1 << 30) / (1 << 30)
) )