[Fix] Reserve the mamba pool's +1 padding slot in the memory budget solve (#32184)
This commit is contained in:
@@ -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)
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user