diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index cd6b55ec4..3835f5fc7 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -1717,7 +1717,7 @@ class KVCacheConfigurator: max_mamba_cache_size=server_args.max_mamba_cache_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: ratio = self._calculate_mamba_ratio() capped_reqs = min( @@ -1726,7 +1726,7 @@ class KVCacheConfigurator: ) intermediate_size = ( config.mamba2_cache_params.mamba_cache_per_req - * capped_reqs + * (capped_reqs + 1) * server_args.speculative_num_draft_tokens ) 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 // 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: intermediate_size = ( 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 ) 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 per_req = config.mamba2_cache_params.mamba_cache_per_req - # Solve jointly for max_mamba_cache_size accounting for intermediate memory. - # The mamba budget (from the ratio split) must cover both: - # 1. main mamba state: max_mamba_cache_size * per_req - # 2. intermediate states: (max_mamba_cache_size / ratio) * D * per_req - # So: max_mamba_cache_size * per_req * (1 + D/ratio) = mamba_budget_bytes + # Solve jointly for max_mamba_cache_size (K), including the pool's + # +1 padding slot on both buffers (see memory_pool.py): + # (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes mamba_budget = ( total_rest_memory * server_args.mamba_full_memory_ratio @@ -1772,7 +1770,8 @@ class KVCacheConfigurator: server_args.override( "mamba_pool.memory_budget_spec", 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 @@ -1781,12 +1780,12 @@ class KVCacheConfigurator: server_args.max_running_requests // self.ps.attn_dp_size, 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)) else: server_args.override( "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. @@ -1805,8 +1804,9 @@ class KVCacheConfigurator: f"(4) use GPUs with more memory." ) + # +1: the pool's padding slot mamba_state_memory = ( - server_args.max_mamba_cache_size + (server_args.max_mamba_cache_size + 1) * config.mamba2_cache_params.mamba_cache_per_req / (1 << 30) )