Skip to content

[fix][model_engine][rollout]: blank a capture's cache write targets, and keep the KV pool size across a wake - #2267

Merged
valarLip merged 2 commits into
ROCm:mainfrom
xysheng-AMD:dev/gpu_util
Sep 20, 2026
Merged

valarLip merged 2 commits into
ROCm:mainfrom
xysheng-AMD:dev/gpu_util

Conversation

@xysheng-AMD

@xysheng-AMD xysheng-AMD commented Sep 17, 2026 •

Copy link
Copy Markdown
Contributor

Description

After a sleep/wake cycle the CUDA graph recapture writes cache through the rows
of the last real batch
, and the wake re-sizes the KV pool to something the
scheduler does not know about. Independent defects, fixed separately.

1. A capture writes cache through live write targets

capture_cudagraph runs the model for real — a warmup forward and then the
captured one — so both write cache through the backend's index buffers. The
capture builders synthesize the read side (kv_indptr, kv_indices,
kv_last_page_lens) because a stale value there walks the kernel off the KV list;
the write side is left alone. And forward_vars is not released on sleep.

So at startup those buffers are still zeroed and a capture writes row 0 — safe by
luck, not design. After a wake they hold a served batch's rows. One cause, two
symptoms, depending on whether the pool also changed size:

  • it shrank (30B: 33562 → 8316 blocks): the rows fall outside it, and all
    eight replicas die with Memory access fault during Capturing bs=256.
  • it did not (8B): the writes land on live slots. One decode row comes back
    all-NaN, the sampler draws token 0 out of it, and that single logprob turns the
    step's whole rollout — grad_norm included — into NaN. Generation quality looks
    unchanged, which is why this was hard to spot.

Fix: point the write targets at nothing before capturing. PAD_SLOT_ID is the
established convention — every step already puts it on the tail of a
CUDAGraph-padded batch, and the cache kernels skip it. Replay overwrites the
buffers with real rows, so nothing about the captured graph changes; that is the
same reason the builders can pin the read side to block 0.

Which buffers are write targets is the backend's knowledge, so the enumeration
is AttentionMetadataBuilder.cache_write_targets() rather than a scan of
forward_vars key names. Paged attention aims writes with slot_mapping, and the
default covers it by suffix so TBO's per-ubatch mirrors come along. This matters
beyond tidiness: GDN aims recurrent-state writes with buffers that are not spelled
that way, one of which (non_spec_state_indices_in_tensor) is not in
forward_vars at all, so no naming rule could have reached it.

GDN's own state indices are left uncovered here. Blanking them reads as
correct, but it was tried and Qwen3-Next stopped coming up, dying inside
LLMEngine.__init__ where the startup capture runs; Kimi-K3, which inherits the
same mixin, was unaffected. That is the state GDN was already in on main — this
PR does not change it, and covering it wrongly breaks a model that works today —
so it waits for a fix that can be run against a checkpoint of one. The PR is still
a net improvement for that family: their attention layers are aimed by
slot_mapping, and that half is fixed.

2. A wake re-sizes the pool, and nothing downstream is told

_resume_kv_cache re-derived the block count and min()-ed it with the saved one,
so anything holding device memory at wake time handed the engine back a smaller
pool.

BlockManager sizes its BlockPool in the engine process, from the count
EngineCore.__init__ copies out of the startup reply by hand — with a comment
saying why: sizing runs in the runner subprocess, so nothing it writes to its own
config is visible there. The wake path has no such step, so the scheduler goes on
issuing block ids past the end of the pool. It was not even one count:
resume_memory broadcasts the command rather than a number, so every TP rank sized
itself and they could disagree.

Fix: allocate saved_blocks. The pool's size is settled when the engine is
sized — the scheduler and the captured decode graphs are both built against it — so
the wake path has no business re-deriving it. Net deletion; the headroom the clamp
used to consult is logged instead, so wake-time pressure still leaves a record.

Two behaviour changes. allocate_kv_cache can now raise, so
_kv_cache_num_blocks is cleared only once it has returned — clearing it first
loses the only record of the pool size, and the next wake then reports success for
an engine with no pool. And that raise is not a clean failure:
resume_memory is not a barrier func, so one rank raising kills that worker while
the others block on a collective it will never reach. It beats clamping only
because clamping was silent corruption; a cross-rank preflight is what would make
it clean, and is not in this change.

Measurements

Qwen3-8B DAPO smoke, 12 steps, releasing sleep. Same base, image and config, same
112 sleep/wake cycles; the only variable is which commit is applied:

no fix fix 1 only both
steps with NaN rollout_ppl 11 / 12 0 / 5 0 / 12
grad_norm NaN on 11 steps finite finite on all 12
wake-time pool shrinks 112 48 — still shrinking 0

The middle column is what makes this one root cause: with only the blanking
applied the pool still shrinks on every wake and the NaN is already gone, so the
NaN is the capture's write and not the sizing.

This head was re-run on that config: 0/12 NaN, 0 shrinks, 112 recaptures completed,
0 faults. That run is also what covers the call site, which the non-GPU suite
cannot reach.

Qwen3-30B-A3B, 8× MI355X, resumed from global_step_110 for 31 steps: CUDA graph
recapture from 0/8 all faulted to 272 completed, Memory access fault from 8 to 0,
wake-time shrinks 0.

All GPU arms ran on 6ced6ce51 plus these two commits rather than the rebased
head, because current main needs a newer aiter than our validation image
carries; both apply there with zero conflicts and identical code.

Checks

  • 14 new test cases; each fails with its fix reverted.
  • Full non-GPU suite 6 failed / 5938 passed; all 6 are
    tests/test_heavy_ci_gate.py, which shells out to jq (absent from the
    container, green in real CI) — identical to base's 6 / 5924. Each commit green on
    its own: 5924 → 5935 → 5938.
  • black --check . 870 files, 0 would-reformat; ruff with the diff-context filter
    against the merge-base, 0 findings would block the job.

@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every eligible PR before approval:

  • ✅ Pre Checkin: Black, Ruff, catalog schema validation, non-GPU unit tests

Heavy model tests:

  • ✅ Run after the PR is approved and Pre Checkin passes
  • ✅ Run immediately when an approval review is submitted
  • ✅ Can be requested before approval with labels
Label Tests
ci:full Run all heavy PR model tests: native ATOM, vLLM, and SGLang
ci:atom Run native ATOM model accuracy tests
ci:vllm Run ATOM vLLM OOT model accuracy tests
ci:sglang Run ATOM SGLang model accuracy tests

Heavy jobs are skipped when the PR is not approved and no matching ci:* label is present.
Add labels via the sidebar or gh pr edit 2267 --add-label <label>

@xysheng-AMD xysheng-AMD changed the title fix(model_engine): size the KV budget against our share of the card, not the leftovers [fix][rollout]: a wake mis-sizes the KV pool, and serves what the recapture wrote into it Sep 17, 2026
@zufayu
zufayu requested a review from yitingw1 September 18, 2026 01:16
@yitingw1

Copy link
Copy Markdown
Contributor

Two real defects, and the diagnosis of 2 is a good catch. But both fixes land on the symptom rather than the shared root cause, and one of them has a hole. Details below.

1. The stated difference between startup and wake does not hold

The comment says "on startup the bytes land in a pool no request has reached yet; here the pool is one step from serving." Both paths go through the same allocate_kv_cache → _carve_paged_pool, where the buffer is torch.zeros(total, ...) (model_runner.py:1859) — so the recapture faces exactly the state the startup capture faces. The new test's stub encodes this too (allocate_kv_cache sets torch.zeros(8)).

The actual difference is the address written, not the pool: forward_vars is not released on sleep. At startup slot_mapping is a zeroed CpuGpuBuffer, so a capture pollutes physical slot 0 only; across a sleep it still holds the last served batch's slot ids, so the recapture scatters writes over thousands of live slots. Startup is therefore safe by luck, and is still uncovered by this fix.

2. build_for_cudagraph_capture already owns this problem — for the read side only

aiter_mla.py:2856-2871 synthesizes self-consistent capture metadata precisely because stale buffers break a capture ("walks the kernel off the KV list (illegal access)"). It pins kv_indptr, kv_indices, kv_last_page_lens, g_kv_indptr. The one buffer left at its live value is slot_mapping — the write side.

Filling it with -1 before the capture writes nothing at all, covers startup and wake in one place, and replaces a full-pool zero_() + synchronize() per wake with a scheduled_tokens-sized fill. -1 is the established skip convention here (eagle_proposer.py:831-833 relies on it for pad rows). Worth confirming every cache-write kernel honours it — MLA does; MHA/GDN/V4 I have not checked. If it cannot be guaranteed, moving clear_kv_cache() to the end of capture_cudagraph() still covers both paths and cannot be forgotten by a future caller.

3. A failed capture pollutes the pool and skips the cleanup

clear_kv_cache() sits after capture_cudagraph() inside the try. The capture loops over (bs, max_q_len), so an exception partway through leaves the earlier sizes' K/V already written, skips the cleanup, and pins the runner to eager — serving from a polluted pool with no later recapture to clean it. That is the branch your own "0/8, all faulted" run took. A finally fixes it.

4. The causal chain for defect 1's fault isn't established

"A resized pool invalidates the decode graphs ... and recapturing them faults" — the first half is right, but a recapture is a fresh capture, not a replay, and the capture builder points every sequence at block 0, so a smaller pool alone should not put it out of range. _warn_if_recapture_will_fault in the same file already attributes wake-time recapture faults to PYTORCH_CUDA_ALLOC_CONF=expandable_segments. Was it set in the failing run, and did that warning fire?

This doesn't affect whether the budget fix should land — charging a peer's memory against our own utilization share is wrong on its own terms — but it decides #5.

5. min(available_for_kv_budget, free) can still shrink the pool on a wake

free is a live reading, so a trainer still holding memory at wake time shrinks the pool anyway; your run just had 176GB free. _resume_kv_cache then only warns and proceeds into the recapture. If a resized pool really is fatal, that warning should raise or force enforce_eager; otherwise the same crash returns under different timing.

6. The baseline assumes the first sizing is alone, and a wrong one is now permanent

Nothing enforces "the number was still its alone" — a trainer that initializes before the rollout engine gets charged to us forever, which is the same double-charge frozen in. Could you put the two non_torch readings (startup vs. wake) from the 30B run into the description? That is the evidence for the assumption, and the block counts don't show it.

One thing that does support it: _maybe_warmup() runs inside ModelRunner.__init__ (model_runner.py:811), strictly before the get_num_blocks RPC, so RCCL buffers are up by the time the baseline is taken. Worth stating in the docstring.

@valarLip

Copy link
Copy Markdown
Collaborator

Review — sizing the KV budget against our own share

Reviewed at 6b0f89cb3. New tests pass (5 passed) and ruff/black are clean on all three files.

The thesis is right and the fix is correct as far as it goes. gpu_memory_utilization is this engine's share of the card, not what is left after everyone else, and latching the baseline does remove the 33556 → 8316 shrink. Everything below is about the cut being shallower than the problem.

Provenance: [verified] means I read the deciding code myself on this SHA. [reported] means the shape matches the code but I did not trace the whole chain.


1. The fix is at the wrong altitude

1.1 A wake-time shrink never reaches the engine process, so the corruption path stays open [verified]

The startup path deliberately carries the number across the process boundary by hand, and says why:

# atom/model_engine/engine_core.py:128-144
block_info = self.runner_mgr.call_func("get_num_blocks", wait_out=True)
num_blocks = block_info["num_kvcache_blocks"]
# Sizing happens in the runner subprocess, so nothing it wrote to
# its own `config` is visible here. ...
config.num_kvcache_blocks = num_blocks

BlockManager is then sized from that engine-process value (block_manager.py:107-108).

The wake path does none of it:

# atom/rollout/memory_manager.py:408-415   -- runs in the RUNNER subprocess
available_blocks = self.get_num_blocks()["num_kvcache_blocks"]
num_blocks = min(saved_blocks, available_blocks)
if num_blocks < saved_blocks:
    logger.warning(f"... KV cache blocks reduced from {saved_blocks} to {num_blocks} ...")
self.allocate_kv_cache(num_blocks)

allocate_kv_cache sets only the subprocess's config.num_kvcache_blocks. Nothing in _handle_resume_memory, broadcast_utility_command_sync or AsyncLLMEngine.wake_up sends the new count back, so the engine-process BlockManager keeps the count it was built with at startup and goes on issuing block ids past the end of the reallocated pool. With saved=33556 and available=8316 the scheduler hands out ids up to 33555 into a pool of 8316 — attention writes past the buffer, which is the Memory access fault by GPU node-N in the PR description.

This PR removes the shrink for one configuration. It does not close the path, and the one signal an operator gets is that logger.warning.

1.2 On wake, each TP rank sizes itself independently [reported]

At startup exactly one reply is taken from call_func("get_num_blocks", wait_out=True) and allocate_kv_cache(num_blocks) is then broadcast, so TP ranks agree by construction. resume_memory broadcasts the command instead, and every rank runs _resume_kv_cache → get_num_blocks() against its own mem_get_info(). The all_reduce at model_runner.py:1681 only fires when pipeline_parallel_size > 1. A colocated trainer that allocates unevenly across cards therefore gives rank 3 a smaller pool than rank 0, and a block id valid on rank 0 is out of bounds on rank 3.

The PR reduces how often ranks disagree; it does not make them agree.

1.3 One of four live inputs was latched [reported]

For the captured decode graphs to stay valid the wake pool must equal the sizing pool. get_num_blocks() has four live inputs: non_torch (now latched), peak_torch and _estimate_cudagraph_overhead() (both ≈0 right after the reset_peak_memory_stats() at memory_manager.py:407), and the min(available_for_kv_budget, free) clamp at :1622, which still reads live free. That last one still shrinks — at utilization values where available_for_kv_budget > free, which is exactly what the error path at :1648 exists for, a colocated trainer still drives available_blocks < saved_blocks and 1.1/1.2 recur.

Allocating saved_blocks unconditionally on wake, or latching the whole computed budget, is the shape that matches the invariant the graphs need.


2. The principle is established but not swept

CLAUDE.md: "Fix-then-sweep: after fixing a bug, immediately grep for the same pattern across the codebase and fix all occurrences in one pass." Two occurrences of the exact pattern this PR names as a bug are untouched.

2.1 _piecewise_skip_capture charges a peer against us [verified]

# atom/model_engine/model_runner.py:3623
free = torch.cuda.mem_get_info()[0]

Raw device free, used to decide which CUDA-graph buckets are skipped as not fitting. On the shared card this PR is about, the trainer's footprint shrinks free and makes ranks skip buckets that would have fit in our own share — one function away from the recapture fault in the PR description. Worse under DP: the dist.all_reduce(free_t, op=dist.ReduceOp.MIN, ...) just below propagates the worst-contaminated rank's reading to every rank in the group.

2.2 The post-init validation advises the opposite of the PR's thesis [verified]

# atom/model_engine/model_runner.py:4065-4082
actual_usage = total_after - free_after                              # every process
target_usage = int(total_after * self.config.gpu_memory_utilization) # our share
...
if usage_ratio > self.config.gpu_memory_utilization + 0.02:
    logger.warning(... "Consider reducing gpu_memory_utilization.")

A colocated trainer guarantees this fires. Under the PR's own thesis the warning is wrong in precisely the way the budget was, and it tells the operator to shrink their own share because of someone else's memory.


3. The baseline's premise is asserted, never checked [reported]

kv_budget.py:38 states that the first reading is ours alone. The latch happens on the first get_num_blocks(), inside LLMEngine.__init__. In a colocated veRL/RLHF job the actor model and optimizer state are commonly already resident on the card when the rollout engine is constructed — so max((total-free)-reserved, 0) at latch time already contains the trainer's tens of GB, and every sizing from then on charges them to us. That is the symptom this PR fixes on the wake path, relocated to startup and made permanent.

There is no detection (the divergence log can only fire after a baseline exists), no reset, and no warning. _non_torch_baseline has no invalidation path at all — release_memory / resume_memory never touch it — so a legitimately changed out-of-allocator footprint (MORI arena resized on an EP change, disagg IPC handles released, RCCL comms rebuilt) leaves a wrong number that only a process restart clears.


4. Tests

4.1 The clamp test is green on both sides of the clamp [verified]

# tests/test_kv_budget_own_share.py:62
assert own_non_torch_bytes(TOTAL, 0, TOTAL, None) == 0

The inner expression is (288GB - 0) - 288GB = 0, so max(0, 0) never does any work. Delete the max(..., 0) from kv_budget.py and all five tests still pass. A case that exercises it needs torch_reserved > total - free, e.g. own_non_torch_bytes(TOTAL, 10 * GB, TOTAL, None).

4.2 The latch state machine has no coverage at all [verified]

The new tests import only the pure own_non_torch_bytes and pass baseline explicitly, so ModelRunner._own_non_torch_bytes — which decides when to latch and owns the field — is never exercised. tests/test_rollout_memory_manager_sleep.py:63 stubs get_num_blocks() to a constant dict, so the sleep/wake path does not reach it either.

The most likely future regression is someone resetting _non_torch_baseline on wake for symmetry with the reset_peak_memory_stats() two lines above it in memory_manager.py. That reintroduces the 33556 → 8316 shrink with a fully green suite. The missing test is one call pair: size twice with a changed free, assert the block count did not move.

4.3 A regression test pins the pre-fix behaviour [verified]

test_a_peer_allocating_later_is_not_charged_to_us:45 asserts own_non_torch_bytes(TOTAL, free_shared, RESERVED, None) == 54 * GB — the un-latched reading that charges the trainer's 38GB to us. Only line 46 tests the fix. If the function ever learns to isolate a peer without a baseline, line 45 fails and reads as "the regression is back" when the opposite happened. Worth splitting under its own name, or stating as a comment rather than an assertion.


5. The divergence log [verified]

# atom/model_engine/model_runner.py:895-904
live = max((total - free) - reserved, 0)
if live != charged:
    logger.info("%s: %.2fGB held outside our allocator, %.2fGB of it another "
                "process's; charging the %.2fGB that is ours",
                self.label, live / (1 << 30), (live - charged) / (1 << 30), ...)

Two problems, both in the direction of misleading the exact investigation this log exists for.

  • live - charged is unguarded. When a peer releases — a trainer finishing, or the torch.cuda.empty_cache() at memory_manager.py:406 dropping reserved — live < charged and the line renders as 0.00GB held outside our allocator, -16.00GB of it another process's. That case is exercised by test_the_baseline_holds_when_a_peer_releases_too, which checks only the return value.
  • Our own post-latch growth is attributed to a peer. allocate_kv_cache:2041 runs a torch.distributed.barrier() whose in-repo comment says it forces lazy NCCL communicator creation and its CUDA allocations, and it runs after get_num_blocks; aiter per-shape JIT code objects land later still during capture. Then live - charged > 0 and the bytes are ours, but the log names a peer.

Related: the method docstring at :885-889 says the readings "only [disagree] when something else is on the card". Both cases above contradict that, and since charged is frozen, once they diverge the info line fires on every subsequent sizing, on every rank, forever.


6. Shape

own_non_torch_bytes ignores three of its four parameters whenever baseline is set, so the name promises a measurement it does not make on the taken branch — and model_runner.py:895 re-implements max((total - free) - reserved, 0) verbatim to compute live. Change the clamp in one place and the logged "live" silently uses a different formula from the charged value. live_non_torch_bytes(total, free, reserved) plus a one-line latch at the call site keeps the same CI-runnable test surface the PR wanted, with one owner for the formula.

_own_non_torch_bytes also reads as a pure query but mutates self._non_torch_baseline on first call; the latch is documented only in an unrelated __init__ comment at :646. Any future pre-check that calls it before the real sizing — logging headroom, a subclass override, a test double — locks the baseline at a moment of different memory pressure, silently and irreversibly.


Of the above, 1.1 is the one I would not merge without: as it stands the PR lands labelled as fixing the fault while the cross-process path that produces it is untouched, and the only signal is a warning line. 2.1 and 2.2 are small changes that finish the sweep the PR's own thesis implies, and 4.2 decides whether this fix survives the next refactor.

@xysheng-AMD

Copy link
Copy Markdown
Contributor Author

@yitingw1 @valarLip — thank you both. You arrived from different sides at the
same objection, and following @yitingw1's #1 showed the two "independent defects"
are one root cause with two symptoms. The KV-budget latch fixed neither and
has been removed.

Root cause. capture_cudagraph runs the model for real, so it writes K and V
wherever slot_mapping points. The capture builders pin the read side
(kv_indptr, kv_indices, kv_last_page_lens) and leave the write side at
whatever the last real batch left — and forward_vars is not released on sleep.
At startup the buffer is still zeroed, so only slot 0 is hit; after a wake it
holds a served batch's slot ids. If the pool also came back smaller those ids
land outside it (Memory access fault); if it kept its size they land on live
slots (the NaN logprob). The budget latch only shrank the shrink, which
downgraded the fault into the NaN.

Two fixes.

  1. capture_cudagraph blanks every slot_mapping mirror to -1 before
    capturing, so the capture writes nothing — on startup and on wake alike. -1
    is the repo's existing "written nowhere" sentinel and every CUDA-graph decode
    step already passes it to the same cache kernels.
  2. _resume_kv_cache allocates saved_blocks instead of re-deriving the count.
    The pool's size is settled when the engine is sized, and the scheduler's
    BlockPool lives in the engine process where no wake-time count ever arrives.
    Net deletion.

Evidence. Qwen3-8B smoke, 12 steps, releasing sleep. Same base, image and
config; the only variable is which commit is applied:

no fix fix 1 only both
steps with NaN rollout_ppl 11 / 12 0 / 5 0 / 12
grad_norm NaN on 11 steps finite finite on all 12
wake-time pool shrinks 112 48 — still shrinking 0

The middle column is the point: with only the blanking applied the pool still
shrinks on every wake and the NaN is already gone, so the NaN is the capture's
write and not the sizing.

On the 30B config the fault path is closed by fix 2: the wake-time count never
reached BlockManager, so a shrink left the scheduler issuing block ids past the
end of the pool. Measured 33562 → 8316 with 176GB still free, then eight
Memory access fault. (For the record on @yitingw1 #4: expandable_segments was
not set and _warn_if_recapture_will_fault never fired, so the old
"resize invalidates the graphs" chain was wrong.)

Review closure. @yitingw1 #1–#5 are addressed (#3 by construction — nothing
is written, so there is no cleanup to skip); #6 no longer applies.
@valarLip 1.1, 1.2 and 1.3 are all closed by fix 2, and §4.2's missing test now
exists. §2–§6 concerned the latch: with the wake call gone get_num_blocks()
runs once per process, the latch branch is unreachable and the module reduces to
the line it replaced, so it is deleted rather than repaired. The remaining
non_torch sizing question is a startup concern and needs something that really
isolates our own footprint — happy to take it, and §2.1/§2.2 with it, as a
follow-up PR or issue.

Behaviour change worth naming: if the memory is genuinely gone,
allocate_kv_cache now raises rather than clamping. The clamp was not a safety
net — it converted an OOM into an out-of-bounds write.

Gate. All four Pre Checkin jobs green on 802af4c13. Locally in a
CI-equivalent container: black 870 files / 0 reformats; ruff diff-context against
the merge-base 0 blocking; non-GPU suite 6 failed / 5934 passed, the 6 being
tests/test_heavy_ci_gate.py (needs jq, green in real CI) and identical to
base. Each commit green on its own; each new test fails with its fix reverted.
Two commits, 5 files, +243/−12, disjoint file sets.

One disclosure: the GPU arms ran on 6ced6ce51 + these two commits rather than
the rebased head, because current main needs a newer aiter than our
validation image carries. Both commits apply there with zero conflicts and
identical code; a 30B run on the new fix is underway.

@valarLip this should close 1.1. If you both agree, an approval would let the
heavy jobs run.

yitingw1
yitingw1 previously approved these changes Sep 20, 2026
@valarLip

Copy link
Copy Markdown
Collaborator

Second pass — reviewed at 802af4c13

First, a correction to my own earlier comment. That review was against 6b0f89cb3, which was a different patch: it added atom/model_engine/kv_budget.py and latched the non_torch reading. The head no longer contains that file. Sections 1.3, 3, 4.1–4.3, 5 and 6 of my previous comment are about code that is not here any more — please ignore them rather than hunting for the lines. Sections 2.1 and 2.2 do still apply and are restated at the bottom.

And the direction of the rewrite is right. The new comment block in _resume_kv_cache states the two reasons a wake must not re-derive, and allocating saved_blocks unconditionally is the shape that matches the invariant the captured graphs need. Everything below is about consequences of that rewrite, plus one thing that is not about the code at all.

Provenance: [verified] means I read the deciding code myself at this SHA. [reported] means the shape matches the code but I did not trace the whole chain.


0. The PR description does not describe this diff [verified]

The body advertises two fixes. Neither is in the code:

  • "Latch the reading once, while the engine is being sized" — model_runner.py:1616 is still non_torch = max((total - free) - torch.cuda.memory_reserved(), 0), a live reading, with no baseline anywhere.
  • "clear_kv_cache() after a successful recapture, in _recapture_cudagraphs_if_needed" — grep -n clear_kv_cache atom/rollout/memory_manager.py returns only its own definition at :217 and two mentions in comments. _recapture_cudagraphs_if_needed does not call it.

The body also cites tests/test_kv_budget_own_share.py (5 tests); that file does not exist at this head. The diff adds tests/test_cudagraph_capture_slot_mapping.py (8 tests) plus two cases in the sleep test.

So the measurement tables in the body were produced by a different patch and are not evidence for this one. That needs fixing before merge independently of anything below — a reviewer or a future bisector reading that body will be looking for code that was never pushed.


1. A failed wake destroys the only record of the pool size [verified]

# atom/rollout/memory_manager.py:404-427
saved_blocks = self._kv_cache_num_blocks
self._kv_cache_num_blocks = None          # cleared here
torch.cuda.empty_cache()
...
self.allocate_kv_cache(saved_blocks)      # and this is now deliberately failable

The rewrite makes allocate_kv_cache the thing that raises when the memory is gone — that is the stated design. But the field is cleared before it, and allocate_kv_cache performs the real allocation (:1901, "Primary KV cache allocation: one buffer, one region per builder"), so it can OOM.

After that raise the state is _kv_cache_num_blocks=None, kv_cache=None. A second wake_up(tags=['kv_cache']) then:

  1. takes the guard at :398-403, logs No KV cache num_blocks to resume from, and returns;
  2. _recapture_cudagraphs_if_needed early-returns because kv_cache is None;
  3. resume_memory returns True (:271).

An engine reports itself awake with no KV pool at all, and faults on the first forward instead of at the point of failure. Clearing the field only after the allocation returns fixes it.

2. "this raises here, which is the honest failure" is not what happens [verified]

The comment at :424-426 justifies dropping the clamp on the grounds that a genuine shortage now surfaces as a raise. Two things make that raise something other than an honest failure:

# atom/model_engine/async_proc.py:247-250
_BARRIER_FUNCS: ClassVar[set[str]] = {
    "update_weights_from_ipc",
    "update_weights_from_shm",
}
# atom/model_engine/async_proc.py:261
out = func(*args)            # no try/except

resume_memory is not a barrier func, and the dispatch is unguarded. So a per-rank OOM kills that worker process. Ranks that survived return and enter capture_cudagraph, which issues collectives inside graph_capture() and blocks forever on the dead peer — a hang, which means the except Exception in _recapture_cudagraphs_if_needed never runs. If the rank that OOMs is rank 0, _handle_resume_memory blocks in call_func(wait_out=True) and the operator sees only a TimeoutError naming neither memory nor a block count.

This is the same failure shape the comment above it rejects re-derivation for: per-rank divergence with no agreement mechanism. A preflight comparison in the engine process — ask every rank for its available count, compare, and raise one message carrying both numbers — would be the honest failure the comment describes.

3. The blanking covers one spelling, and claims to cover every backend [verified]

blank_slot_mappings_for_capture matches by suffix:

# atom/model_engine/cudagraph_capture.py:43-46
for name, buf in forward_vars.items():
    if name.endswith("slot_mapping"):

GDN registers its recurrent-state write indices into the same dict under names that do not end that way:

# atom/model_ops/attentions/gdn_attn.py:306-316
gdn_metadata = {
    "spec_state_indices": self.spec_state_indices_tensor,
    "non_spec_state_indices": self.non_spec_state_indices_tensor,
    ...
}
self.model_runner.forward_vars.update(gdn_metadata)

They are in forward_vars; the suffix rule simply does not select them. And the capture builder hands the graph raw live views of exactly those buffers, with no fill:

# atom/model_ops/attentions/gdn_attn.py:1455-1475
non_spec_state_indices_tensor=self.non_spec_state_indices_tensor.gpu[:bs],
non_spec_state_indices_in_tensor=(self.non_spec_state_indices_in_tensor.gpu[:bs]),
...
gdn_metadata.slot_idx = self.non_spec_state_indices_tensor.gpu[:bs]

So on a wake recapture the warmup forward writes conv/SSM state through live sequences' state rows — Qwen3-Next / GLM-hybrid requests resume from corrupted recurrent state, silently, with no fault. kimi_mla_gdn_attn.py inherits the same builder.

Worth noting in the same breath: the replayssm block a few lines below does reason about this hazard —

"Capture-time only wires up the (address-stable) buffers; the cursor is deliberately NOT advanced here. Warmup and capture replay dummy batches, and letting them commit would leave real sequences resuming from records that were never written."

— for the write cursor, but not for the state indices themselves. The axis is half-recognised in the backend that needs it most.

Two consequences for the module:

  • cudagraph_capture.py:7-8 says the invariant "is enforced once, for every backend". That is not true as written, and it is the sentence that tells the next reader not to audit their backend. [reported]: the review this pass is built on also flags V4 (deepseek_v4_attn.py builds no slot_mapping at all, so the call is a silent no-op there while its capture aims real compressor-state and SWA writes at physical planes), Qwen4-Exp (qwen4_exp_attn.py:686 deliberately points state writes at slots 0..bs-1), and the DSpark paged draft warmup (derives slots from the live block_tables into a drafter-private tensor). I verified the GDN one; I did not trace those three.
  • The sentinel is not even the right instrument for GDN: ssm_state[idx] = ... is an index_put, where -1 wraps to the last state row rather than being skipped.

The shape that would settle it is to declare the axis rather than the spelling — a blank_cache_write_targets() on AttentionMetadataBuilder, defaulting to the slot mappings and overridden where the write target is a state index, so the enumeration lives where the buffers are created instead of being inferred from key names in model_engine/.

4. A dangling cross-reference, in the sentence that justifies the sentinel [verified]

# atom/model_engine/cudagraph_capture.py:13
# "Written nowhere" -- see `AttentionMetaDataBuilder._prepare_slot_mapping`.

grep -rn _prepare_slot_mapping atom/ returns that comment and nothing else, and the class is spelled AttentionMetadataBuilder (backends.py:120) — both halves of the pointer are wrong. A reader who follows it to validate the -1 contract finds nothing. The real source is _write_prefill_slots (backends.py:631), whose docstring states it directly.


5. Still open from the previous pass [verified at this head]

Both are the same double-charging the PR's original thesis named as a bug, in functions the rewrite did not touch. Line numbers have shifted:

  • model_runner.py:3639 — free = torch.cuda.mem_get_info()[0] in the piecewise capture skip guard: raw device free deciding which buckets get captured, so a colocated trainer makes ranks skip buckets that fit in our own share. Under DP the all_reduce(..., ReduceOp.MIN) just below propagates the worst-contaminated rank's reading to the whole group.
  • model_runner.py:4085-4086 — actual_usage = total_after - free_after (every process) compared against total_after * gpu_memory_utilization (our share), warning the operator to "Consider reducing gpu_memory_utilization". A colocated trainer guarantees it fires, and the advice is backwards on exactly the deployment this PR is about.

6. Also reported, not independently traced [reported]

  • BlockManager.clear_cache() has no callers, so nothing drops the prefix-cache hash index when the pool is zeroed or freed. Now that the pool is guaranteed to come back the same size, the stale ids stay in range and nothing faults — a request sharing a prefix is served cached blocks and attends to zeroed K/V. The loud version of this failure is what the PR removes.
  • The capture builders pin kv_indptr / kv_indices / kv_last_page_lens but pass block_tables and context_lens through as raw views of the last real batch (aiter_attention.py:1387-1388, aiter_mla.py:2946-2947, gdn_attn.py:1637-1638). paged_attention_{asm,triton} reads those directly. The docstring's claim that the builders synthesize the read side is half true.
  • _recapture_cudagraphs_if_needed's failure handler clears only self.graphs, not the other stores release_cudagraphs exists to clear, then sets enforce_eager = True — which makes every future release_cudagraphs return at its first line while _release_kv_cache still frees the pool.
  • Dropping the clamp removes the only wake-time consultation of the physical bound, so a wake can succeed by eating the non-torch reserve and then OOM inside the recapture, where the handler turns it into permanent eager mode plus a warning. The deleted logger.warning was also the only record that wake-time pressure existed; logging free/total alongside the count would keep it.
  • The helper duck-types on .np with no isinstance filter, and forward_vars is heterogeneous (ints, bare tensors, CpuGpuBuffer(with_numpy=False) whose .np raises). A future *slot_mapping registered as a plain tensor is an AttributeError inside capture_cudagraph; zero matches is indistinguishable from full coverage. An assertion that at least one key matched turns both into decisions.
  • All eight new tests call the helper directly against a fake buffer; none exercises capture_cudagraph, so moving or deleting the call at model_runner.py:3861 keeps CI green.
  • SLOT_WRITTEN_NOWHERE = -1 is a fourth name for a value already spelled PAD_SLOT_ID in two modules and as a bare literal in several more — including the sites the docstring cites as its authority.

Findings 1 and 3 are the two I would hold on: the first turns a recoverable failure into an engine that reports success with no pool, and the second leaves the corruption this PR exists to fix live on GDN-family models while the module asserts it is fixed everywhere. Item 0 is separate from the code but should not ship as-is either way.

@xysheng-AMD

xysheng-AMD commented Sep 20, 2026 •

Copy link
Copy Markdown
Contributor Author

@valarLip thank you — #3 is the one that mattered.

#0 — the body did not describe the diff. Rewritten. Process failure on my
side: I changed the code and never updated it.

#3 — correct, and the interface is rebuilt in the shape you named. The
enumeration is now AttentionMetadataBuilder.cache_write_targets(); the default
returns the paged write targets and model_engine/ only fills them. The "enforced
once, for every backend"
sentence is gone.

There is a layer past name-matching, too: non_spec_state_indices_in_tensor is not
in forward_vars at all — it lives on the builder — so no rule that scanned that
dict could have reached it. That is why the enumeration had to move rather than get
a better pattern.

But GDN is not covered, and that is after trying. I added the override filling
PAD_SLOT_ID, which reads as correct — a capture's metadata is num_prefills=0,
so the decode branch runs and its kernels skip a negative state_idx. Heavy CI
disagreed: Qwen3-Next stopped coming up, dying inside LLMEngine.__init__,
where the startup capture runs. Kimi-K3 inherits the same mixin and was unaffected,
so the cause is more than the sentinel. Removing the override brought Qwen3-Next
back to green.

That leaves GDN's recurrent state exactly where it was on main — this PR does not
change it, and covering it wrongly breaks a model that works today, so it waits for
a fix that can be run against a checkpoint of one. The PR is still a net gain for
that family: their attention layers are aimed by slot_mapping, and that half is
fixed.

One correction to your mechanism: ssm_state[idx] = ... would indeed wrap on a
negative index, but that line is in the num_prefills > 0 branch, which a capture
cannot reach. It reinforces your point rather than weakening it — the safe value
varies by backend and by branch, so it has to be the builder's call.

On the three you marked [reported]: I checked two, and the new interface does
not rescue them. V4 and Qwen4-Exp do not leave stale values — their capture
builders deliberately compute real state rows 0..bs-1, and they do it after the
blanking (qwen4_exp_attn.py:688, deepseek_v4_attn.py:4126, whose own comment
already names the hazard). They need capture-dedicated scratch rows rather than a
sentinel. I have not traced the DSpark draft warmup.

#1 — correct, fixed. _kv_cache_num_blocks is cleared only after
allocate_kv_cache returns, with a regression test that fails when the order is put
back.

#2 — correct, the comment now says so instead of claiming the failure is clean.
It beats clamping only because clamping was silent corruption. The cross-rank
preflight needs a new engine-process RPC, so I kept it out of these two fixes — say
the word and I will add it. Your §6 point about the lost record is taken: free/total
is logged before the allocation.

#4 — fixed, to _write_prefill_slots at backends.py:643.

§6, taken: SLOT_WRITTEN_NOWHERE is now PAD_SLOT_ID (not imported from
fla_ops.replayssm only because that would pull Triton into a module that must load
without a GPU); the .np duck-typing over a heterogeneous dict goes away with the
scan, and so does zero-matches-looks-like-full-coverage — zero is now a builder's
declaration.

Not in this PR, and I will take any of them: the GDN and V4 / Qwen4-Exp state write targets, which are one
axis with cache_write_targets as the hook — and #2273 on main suggests the
same direction, fixing the plugin path by not assigning real slots for warmup
batches rather than by a sentinel; the cross-rank preflight;
BlockManager.clear_cache() having no callers, where you are right that a
same-size pool makes it quieter rather than louder, but I did not want to touch
cache invalidation here; and the block_tables / context_lens live views, where
you are right that "builders synthesize the read side" was half true.

Verification

Non-GPU: 6 failed / 5938 passed, the 6 being test_heavy_ci_gate.py (needs jq,
green in real CI) and identical to base's 6 / 5924; per-commit 5924 → 5935 → 5938;
black 870 files 0 reformats; ruff diff-context 0 blocking. Each new test fails
with its fix reverted.

GPU, on this head: Qwen3-8B 12 steps — 0/12 NaN steps, 0 wake-time shrinks, 112
recaptures completed, 0 faults. That run is also what covers the call site, which
the non-GPU suite cannot reach (model_runner imports AITER). Against the same
base without the fix it is 11/12 NaN and 112 shrinks; with only the blanking, 0 NaN
while the pool still shrinks — which is what places the NaN on the write rather
than on the sizing. Qwen3-30B-A3B resumed from global_step_110 for 31 steps: 272
recaptures completed, 0 faults, 0 shrinks.

CI: Pre Checkin all green, Qwen3-Next back to success. The remaining reds sit in
container-lifecycle steps (Start CI container, Kill all Docker containers) and
move between models run to run — last round Qwen3-Next failed and DeepSeek-V4-Pro
MTP passed, this round the reverse — so they look like runner flake. GLM-5.2 last
round was 0.9189 < 0.92, a 0.0011 miss, which reads as threshold margin. A re-run
would settle both.

@xysheng-AMD xysheng-AMD changed the title [fix][rollout]: a wake mis-sizes the KV pool, and serves what the recapture wrote into it [fix][model_engine][rollout]: blank a capture's cache write targets, and keep the KV pool size across a wake Sep 20, 2026
xysheng-AMD and others added 2 commits September 20, 2026 13:43
…batch's

`capture_cudagraph` runs the model for real -- a warmup forward and then the
captured one -- so both write cache wherever the backend's index buffers point.
The capture builders synthesize the *read* side (`kv_indptr`, `kv_indices`,
`kv_last_page_lens`) precisely because a stale value there walks the kernel off
the KV list, and leave the write side at whatever the last real batch left.

At startup those buffers are freshly zeroed, so a capture writes row 0 and
nothing else. `forward_vars` is not released on sleep, so on a wake they still
hold a served batch's rows and the capture scatters writes across live entries
-- or past the end of the pool, if it came back smaller.

Both were measured. On Qwen3-30B-A3B, where a wake also resized the pool from
33562 blocks to 8316, the stale ids fell outside it and all eight replicas died
with `Memory access fault by GPU node-N` during `Capturing bs=256` --
`expandable_segments` was not set and `_warn_if_recapture_will_fault` never
fired, so the alloc-config path the warning exists for is not this. On Qwen3-8B
at 4k, where the pool kept its size, the same writes landed on live slots
instead: one decode row per batch came back NaN across the whole 151669-entry
vocabulary while its `context_lens` incremented normally, the sampler drew
token 0 out of that row, and that single logprob turned the step's entire
rollout -- every aggregate, and `grad_norm` with it -- into NaN, so all eight
actors skipped the update. Ten of twelve steps were lost that way; with this,
none of twelve. Applying only this commit and leaving the pool free to shrink
on every wake is also 0 of 12, which is what places the NaN on the write rather
than on the sizing.

So the capture is pointed at nothing first. `PAD_SLOT_ID` is the established
sentinel: every step already puts it on the tail of a CUDAGraph-padded batch,
the paged cache kernels skip it (`slot_idx < 0`), and so do GDN's
`fused_recurrent_gated_delta_rule` and replayssm kernels. Replay overwrites the
buffers with real rows, so nothing about the captured graph changes -- the same
reason the builders can pin the read side to block 0.

Which buffers those are is the backend's knowledge, so the enumeration is
`AttentionMetadataBuilder.cache_write_targets` rather than a scan of
`forward_vars` key names. Paged attention aims writes with `slot_mapping`, and
the default covers it by suffix so TBO's per-ubatch mirrors come along. GDN aims
recurrent-state writes with `spec_state_indices` / `non_spec_state_indices`,
which are not spelled that way, and with `non_spec_state_indices_in_tensor`,
which is not in `forward_vars` at all -- no naming rule could have reached the
second kind, which is why the enumeration moved rather than got a better regex.

GDN's own state indices stay uncovered: blanking them reads as correct, but it
was tried and Qwen3-Next stopped coming up inside `LLMEngine.__init__`, where the
startup capture runs. Uncovered is the state that backend was already in, so it
waits for a fix that can be run against a checkpoint of one.

Both live on the builder, next to the capture metadata they belong with.

Co-authored-by: Cursor <cursoragent@cursor.com>
`_resume_kv_cache` re-derived the block count and `min()`-ed it with the saved
one, so anything holding device memory at wake time silently handed the engine
back a smaller pool. Nothing downstream is told.

`BlockManager` sizes its `BlockPool` in the ENGINE process, from the
`config.num_kvcache_blocks` that `EngineCore.__init__` copies out of the startup
reply by hand -- with a comment saying why: sizing runs in the runner
subprocess, so nothing it writes to its own `config` is visible there. The wake
path has no such step. `allocate_kv_cache` sets the subprocess's copy, and
neither `_handle_resume_memory`, `broadcast_utility_command_sync` nor
`AsyncLLMEngine.wake_up` carries the count back, so the scheduler goes on
issuing block ids past the end of the reallocated pool.

It was not even one count. `resume_memory` broadcasts the command rather than a
number, so every TP rank sized itself against its own `mem_get_info()`; the
`all_reduce` in `get_num_blocks` only fires under pipeline parallelism. A peer
allocating unevenly across cards could therefore leave a block id valid on rank
0 and out of bounds on rank 3.

Measured on Qwen3-30B-A3B at utilization 0.45: 33562 blocks at startup, 8316 on
the first wake while 176GB was still free, and eight `Memory access fault by GPU
node-N`.

The pool's size is settled when the engine is sized -- the scheduler and the
captured decode graphs are both built against it -- so the wake path has no
business re-deriving it. `reset_peak_memory_stats()` goes with the re-derivation,
`get_num_blocks` having been its only reader here; the headroom it used to
consult is logged instead, so wake-time pressure still leaves a record.

That makes `allocate_kv_cache` the call that can fail, which changes two things.
`_kv_cache_num_blocks` is now cleared only once it has returned: clearing it
first loses the only record of the pool size on the way out, and the next wake
then takes the "nothing to resume from" guard and lets `resume_memory` report
success for an engine with no pool at all. And the raise is not a clean failure
-- `resume_memory` is not a barrier func and `AsyncIOProc` does not guard the
dispatch, so one rank raising kills that worker while the others block on a
collective it will never reach. The comment says so rather than claiming
otherwise; what makes this preferable to clamping is that clamping was silent
corruption, and a cross-rank preflight is the thing that would make it clean.

Co-authored-by: Cursor <cursoragent@cursor.com>
@valarLip
valarLip merged commit 0795f0e into ROCm:main Sep 20, 2026
20 of 36 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants