Skip to content

fix(PP): size the mamba pool per pipeline stage, not per whole model - #33666

Merged
YAMY1234 merged 1 commit into
sgl-project:mainfrom
YAMY1234:fix/pp-mamba-pool-sizing
Aug 6, 2026
Merged

YAMY1234 merged 1 commit into
sgl-project:mainfrom
YAMY1234:fix/pp-mamba-pool-sizing

Conversation

@YAMY1234

@YAMY1234 YAMY1234 commented Aug 5, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

KVCacheConfigurator._handle_max_mamba_cache sizes the mamba state pool from config.mamba2_cache_params.mamba_cache_per_req, which is computed over the whole model's mamba layer list:

) * len(self.layers)      # configs/mamba_utils.py

Under pipeline parallelism each rank only allocates state for the layers in its own [start_layer, end_layer) slice — the allocation paths in this same file already filter on that range. Charging the whole model against a single stage over-estimates the per-slot cost by roughly pp_size, so the pool is sized for a fraction of the requests that actually fit.

On Kimi-K3 (93 layers, 69 linear-attention layers) at pp_size=8 this resolves to a 26-slot pool, which clamps max_running_requests to 6. That propagates: resolve_max_num_reqs clamps max_running_requests to max_mamba_cache_size // ratio, which in turn sets pp_max_micro_batch_size.

Modifications

Scale the per-request cost (and the ReplaySSM ring, which is sized the same way) by the stage's share of the mamba layers.

The scale is taken from the stage holding the most mamba layers, not from this rank's own count. Both give the same system-wide capacity — the system is limited by the heaviest stage either way — but the maximum is identical on every rank, so max_mamba_cache_size, max_running_requests and pp_max_micro_batch_size stay uniform across stages without a collective. Scaling by each rank's own share instead makes those values diverge whenever the mamba layers do not split evenly, which is the common case: at pp_size=8 a 93-layer model distributes its 69 mamba layers as [9, 8, 8, 9, 9, 9, 9, 8].

Sweeping 145 mamba budgets from 4 to 40 GiB on that layout:

scaling budgets where pp_max_micro_batch_size diverges across stages
this rank's own share 95/145 (65.5%)
largest stage's share (this PR) 0/145

pp_max_micro_batch_size has to agree on every stage: it decides how a batch is split into microbatches, so a stage that picks a different split disagrees with its neighbours about which sequences are in flight.

Every rank computes the layer distribution locally from get_pp_indices, so nothing is communicated. At pp_size=1 the rank holds every mamba layer and the scale is exactly 1.

Accuracy Tests

End to end on Kimi-K3 with real weights, pp_size=8 tp_size=1, 2 nodes, --attention-backend flashinfer, GSM8K 200 questions:

max_mamba_cache_size max_running_requests GSM8K
before, --chunked-prefill-size 16384 26 6 0.985
after, --chunked-prefill-size 16384 210 52 0.990
before, --chunked-prefill-size 65536 26 6 0.985
after, --chunked-prefill-size 65536 210 52 0.985

The pool grows 8.1x and the concurrency ceiling 8.7x, while the score stays in the same band (the 0.985/0.990 spread is one question out of 200). The change only lifts a capacity limit; it does not touch any compute path.

What that "before" column looks like in the startup log: the unpatched tree allocates a per-stage pool but charges the whole model against it, so every stage lands on the same starved cap.

PP0  max_mamba_cache_size: 26   conv_state 0.05GB   ssm_state 1.42GB
PP1  max_mamba_cache_size: 26   conv_state 0.04GB   ssm_state 1.27GB
PP2  max_mamba_cache_size: 26   conv_state 0.04GB   ssm_state 1.27GB
PP3  max_mamba_cache_size: 26   conv_state 0.05GB   ssm_state 1.42GB
PP4  max_mamba_cache_size: 26   conv_state 0.05GB   ssm_state 1.42GB
PP5  max_mamba_cache_size: 26   conv_state 0.05GB   ssm_state 1.42GB
PP6  max_mamba_cache_size: 26   conv_state 0.05GB   ssm_state 1.42GB
PP7  max_mamba_cache_size: 26   conv_state 0.04GB   ssm_state 1.27GB

all stages: max_running_requests is capped to 6 by the mamba state cache
            (max_mamba_cache_size=26, 4 state slots per request)

The allocated ssm_state already varies with what each stage holds — 1.42 GB on the stages with 9 mamba layers, 1.27 GB on those with 8 (the layers distribute as [9, 8, 8, 9, 9, 9, 9, 8]). The pool size does not: all eight stages are sized as if they held all 69.

The same over-charge also applies when the pool size is given explicitly. --max-mamba-cache-size skips the auto solve, so the pool itself comes out right, but the memory the solver then books for it still uses the whole-model per-request cost — leaving too little for the KV pool. Same layout, 70 GiB handed to the solver:

--max-mamba-cache-size KV budget left, before after
112 39.36 GiB 66.00 GiB
144 30.69 GiB 64.87 GiB
210 12.80 GiB 62.54 GiB
328 -19.20 GiB 58.37 GiB

At 328 the budget goes negative and there is no KV pool left to build, which is what an operator sees as "raising --max-mamba-cache-size makes startup fail".

Driving _handle_max_mamba_cache once per rank on the same layer layout with an 8 GiB budget isolates the same effect:

before after
pp_size=1 pool 8 8 (unchanged)
pp_size=8 pool, per rank 8, 8, 8, 8, 8, 8, 8, 8 74, 74, 74, 74, 74, 74, 74, 74

TestPPMambaPoolSizing in test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py adds two CPU-only tests. test_stage_is_not_charged_for_the_whole_model fails on the unpatched tree and passes with this change; running that file goes from 1 failed, 7 passed to 8 passed.

Speed Tests and Profiling

The pool size this PR unlocks is what an agentic serving configuration actually wants. On a 1P1D deployment (8 GPUs running PP8 prefill, 8 running TP8 decode) with a trace-replay workload of 98,827 requests — mean input 219k tokens, p90 549k, ~98% theoretical prefix reuse — a 144-slot mamba pool with a 14.5M-token KV pool sustains:

client concurrency effective total per GPU TTFT p50 prompt cache read
16 70,756 tok/s 4,422 — 94.7%
24 78,437 tok/s 4,902 7.05 s 93.6%
32 26,184 tok/s 1,637 85.10 s 92.8%
32, decode --decode-context-parallel-size 8 82,977 tok/s 5,186 6.8 s 71.2%

"Effective total" counts prompt tokens including cache hits, which is what the client sees; at 93.6% reuse the recomputed share is far smaller. The per-GPU column divides by all 16 GPUs in the deployment. The last row is the mean of two runs; the drop at concurrency 32 is a decode-side KV limit, not a mamba one, and decode context parallelism removes it.

That operating point is not reachable without this change: the auto solve gives roughly 1/pp_size of it, and setting --max-mamba-cache-size by hand instead runs into the accounting problem above. The numbers are a capability statement, not a before/after delta — there is no matched "before" run at this pool size because the configuration does not come up.

Per-token cost is unchanged; this PR only lifts a capacity limit.

Note for existing deployments

Correcting the over-charge frees real memory, so at an unchanged --mem-fraction-static the solver hands the KV pool what the mamba pool was previously over-booking. On a Kimi-K3 PP8 prefill instance that moved max_total_tokens from 2,123,072 to 21,287,872 at the same flag value, and the instance then failed to start. The old over-charge was acting as unintended headroom. Deployments tuned against the previous accounting may need to lower --mem-fraction-static or pin --max-total-tokens after this change.

Checklist

Merge order

This PR raises the concurrency ceiling, which makes packed prefill batches larger. On the same hardware that turns out to be what makes the int32 overflow in #33665 reachable in practice: before this change max_running_requests is clamped to 6 and a batch never gets near the 2^31 threshold, after it the ceiling is 52. Worth landing that one first.


CI States

Latest PR Test (Base): ✅ Run #30984169609
Latest PR Test (Extra): ✅ Run #31123531254

_handle_max_mamba_cache charges mamba_cache_per_req, which is sized from
the whole model's mamba layer list, but each PP rank only allocates state
for the layers in its own [start_layer, end_layer) slice. Every stage is
therefore charged roughly pp_size times what it holds and the pool is
starved -- on Kimi-K3 (93 layers, 69 linear-attention layers) at pp 8 the
pool resolves to 26 slots and clamps max_running_requests to 6, where the
same run with this change gets 210 slots and 52.

Scale the per-request cost by the share held by the stage with the most
mamba layers. That is the capacity the system is limited by anyway, and
taking the maximum rather than each rank's own count keeps the derived
max_mamba_cache_size, max_running_requests and pp_max_micro_batch_size
identical on every stage without a collective. Scaling by each rank's own
share instead makes those values diverge whenever the mamba layers do not
split evenly, which is the common case.

No-op at pp_size 1, where the rank holds every mamba layer.
@YAMY1234

YAMY1234 commented Aug 6, 2026

Copy link
Copy Markdown
Collaborator Author

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Aug 6, 2026
@YAMY1234 YAMY1234 changed the title fix(mamba): size the mamba pool per pipeline stage, not per whole model fix(PP): size the mamba pool per pipeline stage, not per whole model Aug 6, 2026
@YAMY1234
YAMY1234 merged commit 2fc5572 into sgl-project:main Aug 6, 2026
276 of 306 checks passed
Fridge003 pushed a commit that referenced this pull request Aug 7, 2026
…eline stage, not per whole model (#33666) (#34035)

Co-authored-by: YAMY <74099316+YAMY1234@users.noreply.github.com>
sagearc pushed a commit to sagearc/sglang that referenced this pull request Aug 13, 2026
…eline stage, not per whole model (sgl-project#33666) (sgl-project#34035)

Co-authored-by: YAMY <74099316+YAMY1234@users.noreply.github.com>
Signed-off-by: Sage Ahrac <sagiahrak@gmail.com>
@YAMY1234
YAMY1234 deleted the fix/pp-mamba-pool-sizing branch August 25, 2026 16:16
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

bypass-fastfail run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants