Skip to content

[Lumen-RL] Improve FP8 rollout weight synchronization and CUDA Graph stability - #2028

Merged
valarLip merged 28 commits into
ROCm:mainfrom
ZhangDanyang-AMD:lumen-rl
Sep 16, 2026
Merged

valarLip merged 28 commits into
ROCm:mainfrom
ZhangDanyang-AMD:lumen-rl

Conversation

@ZhangDanyang-AMD

@ZhangDanyang-AMD ZhangDanyang-AMD commented Aug 26, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

Make ATOM's rollout weight synchronization correct for a colocated RL trainer:
every tensor the trainer sends reaches the engine, in the layout the kernel
expects, without invalidating captured CUDA graphs.

Validating that on Lumen-RL's ATOM examples turned up three correctness bugs
that are silently wrong today. Two are in the scheduler and fire only under
preemption; the third is in this PR's own weight-sync path and only shows up
with FP8 online quantization. All three are fixed here, each in its own commit.

Rebased onto main at 82f0b24. Verified on Qwen3-8B-Base and Qwen3-30B-A3B
on MI355X (gfx950).

Part 1 — rollout weight synchronization

Routed-expert weights were never synced, and nothing said so

A model's routed experts arrive one tensor per expert and belong in the fused
w13_weight / w2_weight of the layer's FusedMoE, which is neither the
incoming name nor anything packed_modules_mapping describes. They matched
nothing, were counted as skipped at debug level, and the rollout went on
serving whatever the experts held at load time. On Qwen3-30B-A3B that is 96
tensors per replica per sync, 48 layers x 2.

Routing them is only half of it. FusedMoE.weight_loader writes plain
row-major bytes over buffers the aiter kernel reads through a permutation, so a
sync has to re-establish that layout exactly as the initial load does — but not
by re-running process_weights_after_loading. Those hooks are initialisation:
they hand the module new Parameter objects while a captured graph and
_param_to_module still point at the old ones, and several are not idempotent
(Fp8MoEMethod's per-tensor path collapses w13_weight_scale from [E, 2] to
[E] on its first call, so a second raises IndexError).

shuffle_expert_slices() is the layout step alone, applied in place to only
the slices this sync rewrote, once after the last shard. The pending set
accumulates across buckets, so an expert whose w1 and w3 arrive in different
buckets is relaid out exactly once, and a half-rewritten expert raises instead
of mixing two layouts.

Combinations this path does not implement raise before the write rather than
reporting updated=1: a quantized expert buffer (the loader would byte-copy or
numerically cast, leaving the scale describing the old weight), expert
parallelism, and redundant expert replicas.

A trainer that keeps experts fused had to reshape tensors on ATOM's behalf

A trainer whose transformers keeps MoE experts fused sends one 3D tensor per
layer instead of three per expert: (E, 2I, H) experts.gate_up_proj and
(E, H, I) experts.down_proj. Same buffers, same dim order, w13's first half
along the intermediate dim being the gate projection — only the leaf name
differs.

Accepting both conventions means a caller does not have to know which one ATOM
happens to use, rename its tensors, split them, or pre-apply the kernel layout.
One fused tensor covers every expert and both halves of w13, so it is driven
through weight_loader once per half: a 3D loaded_weight puts the loader on
its full-load path, where the expert dimension is written whole and the
intermediate dimension is still narrowed by TP rank. The halves go in as chunk
views rather than copies — the loader's copy handles a strided source, and
materialising them would double the largest tensor in the sync.

The shuffle layout was decided in two places that disagreed

Whether a quantized 2D GEMM weight is held preshuffled was decided in
LinearBase.process_weights_after_loading for the initial load and in
WeightUpdaterMixin._post_process_fp8_weight for an online update. The two
disagreed wherever the rule was not simply the env var: the module-level
needs_preshuffled_weight exception that DeepSeek's fused qkv_a_proj sets,
the triton a8w8 per_Token GEMM that wants the unshuffled weight, the
non-shuffle FP4 blockscale GEMM. Each disagreement leaves a synced weight in a
layout the loaded one would never have had, which the kernel then reads through
the wrong permutation.

Both sides now call one function, weight_is_stored_preshuffled(). The sync
side also gets the rank check the load already had: only 2D weights are
shuffled, because Qwen3-Next's GDN conv1d expands its weight to 3D and has to
stay row-major.

Two supporting fixes in the same area:

  • shuffle_weights rebound tensor.data to aiter's return value for a 2D
    weight, handing the parameter a new address while a captured decode graph
    still held the old one. It now writes through the existing storage, the way
    its own 3D branch already did.
  • Online quantization replaces self.weight and self.weight_scale with fresh
    Parameters, which carry none of the attributes __init__ hung on the
    originals. weight_loader() reads weight_loader_process off the parameter
    it is handed, so a later weight update raised on it.

A level-1 sleep left decode graphs pointing at a freed KV pool

AsyncLLMEngine.sleep(level=1) — the default level — releases the KV cache and
nothing else. The decode graphs captured the base address of that pool, but
only _release_weights cleared them and recorded _graphs_backup_keys, so a
level-1 sleep left them in place and _recapture_cudagraphs_if_needed returned
early on wake with nothing to recapture. _resume_kv_cache then allocated a
pool of a different size at a different address and the next decode replayed
the old graphs against it.

On Qwen3-8B-Base on MI355X that is not a wrong answer, it is a crash:

Model Runner0/1: KV cache blocks reduced from 47740 to 44626
Memory access fault by GPU node-2 on address 0x7c9f5ca20000
[ModelRunner0/1] proc died unexpectedly (exitcode=-6)

The graph release moves into release_cudagraphs() and is called from both
release paths, so whichever of the two allocations a sleep frees, the graphs
that captured it are dropped and wake recaptures them. This predates the
branch; main has the same shape.

Graph recapture on wake faults under expandable_segments

Sleep frees the weights and the KV pool, and wake recaptures the decode graphs
against their new addresses. That recapture faults under
PYTORCH_CUDA_ALLOC_CONF=expandable_segments, so
Config.sleep_keeps_memory_resident offers a way out: keep both allocated,
nothing the graphs captured moves, and there is nothing to recapture.

The cost is exactly the memory a colocated trainer sleeps the rollout engine to
reclaim, so it is opt-in rather than the default for every non-eager
deployment, and it has no effect under enforce_eager. A host that defines
neither enforce_eager nor the config field gets the behaviour every caller
had before this option existed. This is the only behavioural change in the
series, and it is last in the history so it can be dropped on its own if
upstream would rather fix the recapture fault than work around it.

A padded embedding matrix let the sampler return an undecodable id

A checkpoint whose embedding matrix is padded up to a friendlier width — Qwen3
rounds 151665 real tokens up to 151936 — leaves the tail rows holding whatever
the padding was initialised to. On that checkpoint they are copies of an
existing embedding rather than zero or -inf, so the sampler reaches them and
can return an id the tokenizer cannot decode, while the trainer masks exactly
those positions on its side.

The number used to arrive in LUMENRL_ATOM_TRUE_VOCAB_SIZE, an environment
variable named after a downstream project. It is a property of the checkpoint,
not of the deployment, so it belongs in Config, and a downstream-prefixed
name has no place in ATOM. Config.true_vocab_size defaults to 0, meaning
"this vocabulary is not padded" — the right answer for every model whose
embedding matrix matches its tokenizer, at the cost of one comparison per step.

A value that would mask nothing is refused rather than silently accepted: it
counts the tokenizer's real tokens, so it can be neither negative nor larger
than the rows the embedding matrix has, and either way round it disables the
mask — the exact failure this path exists to prevent. The override also moved
out of model construction into RLHF post-processing, so compiled Qwen3 stays
integration-agnostic.

Two AITER fallbacks that named no version and were never exercised

moe.py guarded from aiter.ops.shuffle import moe_shuffle_scale with an
ImportError fallback to shuffle_scale. The two are not aliases:
moe_shuffle_scale dispatches on the chip and calls shuffle_scale_n32k4 on
gfx1250, so the fallback would have quietly laid out MoE scales for the wrong
kernel there. AITER has exported moe_shuffle_scale since #3756 (2026-06-20)
and main imports it directly, so that is restored.

layernorm.py called AITER's RMSNorm with use_model_sensitive_rmsnorm=1
inside a try/except TypeError that parsed the exception message and cached
the verdict in two module-level globals. That parameter has been part of
rmsnorm2d_fwd and rmsnorm2d_fwd_with_add since #647 (2025-07-17), and both
are @torch_compile_guard custom ops where probing from inside the traced
region is fragile. It is now passed straight through: the argument defaults to
0 in both signatures, so passing the still-default-off
ATOM_USE_MODEL_SENSITIVE_RMSNORM unconditionally leaves the default path
exactly as main has it, and there is nothing for a branch to choose between.
Verified bit-identical to main's call on MI355X.

add_request fan-out had no test

LLMEngine.add_request routes through preprocess_fanout so that
SamplingParams.n > 1 produces n sequences instead of one. Nothing asserted
that, and the failure mode — silently returning a single sequence — is
invisible to the request counters.

Part 2 — three correctness fixes

The deferred-output placeholder survived preemption

preempt only stripped the trailing placeholder token when speculative
decoding was on. Without speculation it left one behind, and the placeholder is
eos_token_id.

The placeholder is not a speculation-only artifact: postprocess appends one
whenever need_placeholder holds, which includes is_deferred_out — defined
as pipeline_parallel_size == 1 (model_runner.py:189), i.e. true for every
TP-only engine. The scheme is safe because the next step's postprocess
overwrites the placeholder in place. A sequence preempted on this step never
reaches that next step, so the overwrite never happens:

  • the re-prefill feeds the model a context ending in <|endoftext|>, so it
    starts a fresh document instead of continuing the answer;
  • the same token is returned to the caller as generated output, and under
    ignore_eos=False the request terminates right there.

Symptom: a reply that reads coherently and then stops before it answers
anything.

Fix: on the non-speculative path strip seq.num_placeholder_tokens, the
counter postprocess sets where it appends and preempt clears just below.
The matching 0.0 entries the placeholders pushed onto seq.logprobs go with
them — postprocess patches those in place too, so removing tokens without
them would leave the two lists describing different positions.

A chunked recompute was declared final at the prompt boundary

is_final_chunk compared progress against seq.num_prompt_tokens. Correct for
a first admission, wrong for a sequence re-admitted after preempt, which must
recompute prompt plus every token it had already generated. Every other
length in the scheduler already accounts for this — Phase 1's remaining,
Phase 2's num_new_tokens, and the next_token_ids loop directly below, which
carries a comment saying exactly that.

So a 337-token recompute split into 256 + 81 was declared complete after the
first chunk, because 256 >= 128 held against the prompt length. The remaining
81 positions never got their KV computed. The sequence then decoded against
blocks still holding whatever the previous owner wrote into them, and carried
on producing that other request's answer.

postprocess re-derived the same predicate as
seq.num_cached_tokens < seq.num_prompt_tokens. It now reads the
is_final_chunk the scheduler already froze at pre-advance offsets, rather
than re-deriving it: by the time postprocess runs, num_tokens may have
grown by this step's sampled token, so neither length is a valid bound. A
length-based fallback stays for callers that pass a batch without the field.

Both fire only when the KV pool is too small to hold max_num_seqs sequences
at once, which is why they have gone unnoticed: a comfortably sized pool never
reaches either. This path is common under memory pressure — 256 sequences
recomputing ~900 tokens each against max_num_batched_tokens=8192 fits 9 per
batch.

A weight update overwrote weights their readers were still reading

This one is in this PR's own code. A weight update rewrites the parameter
buffer in place, and the FP8 path does it twice: once for the quantized bytes,
again for the kernel's shuffled layout.

In place is deliberate, and the alternative is worse. shuffle_weights used to
rebind tensor.data to aiter's return value, which moves the address out from
under a captured decode graph; the graph then replays against the old buffer.
On the ATOM FP8 rollout example that is total loss — 192 of 192 sequences come
back as !!!!!!!!, every one running to the length cap, and DAPO dies with
"filter_groups collected no valid groups".

What keeping the address costs is that the write lands in a buffer which is
still live, and nothing waited for the work reading it. An update follows
generation immediately, so the last decode replays of the step that just ended
can still be in flight. Overlap one with the shuffle and the graph reads a
half-permuted weight. Generation continues, every sequence past that point is
token soup, and nothing raises.

Waiting once per update is not enough — measured over seven weight syncs it
still lost five of them. The wait has to sit with the write, so it goes before
the first in-place write of each parameter and again in the layout
post-process, which the direct-copy call sites reach without passing through
the requantize.

The BF16 rollout was never affected, which is why this went unnoticed: it
writes the weight once and never shuffles.

Test Plan

python -m compileall -q atom

# Part 1 and the FP8 ordering fix
python -m pytest -q \
  tests/model_ops/ \
  tests/test_llm_engine_add_request_fanout.py \
  tests/test_rollout_expert_weight_sync.py \
  tests/test_rollout_memory_manager_sleep.py \
  tests/test_rollout_vocab_mask.py \
  tests/test_weight_sync_shuffle_layout.py \
  tests/test_weight_sync_inplace_ordering.py \
  tests/test_envs.py

# Part 2, the scheduler
python -m pytest -q \
  tests/test_scheduler.py tests/test_block_manager.py \
  tests/test_prefill_scheduler.py tests/test_scheduler_partial_prefill_tail.py \
  tests/test_per_req_cache_decoupling.py tests/test_scheduler_mtp_max_tokens.py \
  tests/test_dspark_scheduler.py

Test Result

Python compilation passed. Part 1 + FP8 ordering: 239 passed. Scheduler:
301 passed. Black clean across the repo; Ruff reports nothing on the lines
this PR adds.

The expert-sync tests build the fused target from scratch through the initial
load path and assert the synced buffers are bit-identical to it, which is the
only check that catches a layout that is merely self-consistent. The
sleep/wake tests assert on data_ptr() and object identity across two
sleep/wake cycles, because the point of sleep_keeps_memory_resident is that
nothing moves — something a test on updated/released counters cannot see.

The FP8 weight-sync fix

An 8-step ATOM FP8 DAPO smoke (Qwen3-8B-Base, 8 colocated replicas, 24 prompts
x 8 samples, gpu_memory_utilization=0.30) performs seven weight syncs:

wait syncs corrupted mean step
none 7 of 7 62.0s
once per update, at the entry points 5 of 7 56.3s
before the in-place layout rewrite 0 of 7 41.8s
before the first in-place write 0 of 7 38.2s

Corruption is mismatch/k3_kl — the trainer's KL between the logprobs the
rollout reported and the actor's own for the same tokens — crossing from
0.003-0.05 to 47-92, with entropy following 0.5 -> 7.0, grad_norm 0.8 -> 10.7,
and the share of unparseable answers roughly doubling. The wait is not what
costs time: weight sync goes from 1.08s to 1.13s per update, and the step gets
faster, because a corrupted rollout generates garbage until it hits the
length cap.

Three steps is not enough to judge this: it samples the event twice, and the
first fix that looked clean over three steps lost five of seven over eight.

The scheduler fixes: a generation whose correctness does not depend on floating point luck

Greedy token-by-token comparison cannot be used here: continuous batching
reshapes every forward and bf16 reductions are not associative, so two
correct runs of this engine already disagree on 74 of 96 prompts.

Instead each request counts upward from its own disjoint range. A break in the
run of consecutive integers means the context stopped being its own; a
consecutive run of numbers from another request's range means it is attending
to that request's KV. Neither is reachable by rounding.

Qwen3-8B-Base, 1x MI355X, 64 requests, max_num_seqs=64, max_tokens=512,
ignore_eos=True, prefix caching off, chunked prefill on, KV pool size forced
directly so the only variable is pool size.

KV blocks max_num_batched_tokens scheduler preemptions clean drifted into another request
6000 8192 both fixes 0 64/64 0
800 8192 neither fix 74 26/64 0
800 8192 both fixes 84 64/64 0
420 8192 both fixes 200 64/64 0
800 256 neither fix 74 19/64 9
800 256 first fix only 84 32/64 9
800 256 both fixes 84 64/64 0

Attribution, from the scheduler's own preemption and re-prefill logs:

  • In the 800 / 8192 / neither run, 45 of 64 requests were preempted and 38
    broke. Every broken request was one of the preempted ones; not a single
    request that was never preempted broke.

  • The first fix only row isolates the second bug. All 32 broken requests
    there were chunked recomputes, and 9 continued another request's count with
    no <|endoftext|> anywhere near the break:

    • req 21 counts correctly to 3100061, then emits 1200000, 1200001, 1200002, ... — req 2's range.
    • req 22 to 3200057, then 2200003, 2200004, 2200005, ... — req 12's.
    • req 24 to 3400053, then 24000000, 24000001, 2400002, ... — req 14's.

    These are consecutive runs inside another request's range, not a stray round
    number, which is what a lost-context model would produce.

Excluded as confounders, all 64/64 clean with a roomy pool: prompt length not
16-aligned, prompts extended to ~600 tokens, and plain chunked prefill on new
prompts with no preemption. Preemption is the only trigger.

CPU-only scheduler simulator

Driving the real Scheduler and BlockManager with a fake forward, checking
that no block is held by two live sequences, that every forwarded position has
a block, and that returned completions match the fake model's token stream:

KV blocks requests preemptions requests preempted prefill/decode result
4675 768 2571 492 (64.1%) 1.08 no invariant broken
34164 768 0 0 0.16 no invariant broken

This is a no-regression check only — it produces byte-identical output on the
unpatched tree, because a fake forward cannot reproduce reading another
request's KV. It does quantify what preemption costs: for the same 655k decode
tokens, prefill work goes from 103k to 710k tokens, a 6.9x increase.

End-to-end

BF16 ATOM rollout, 8-step DAPO smoke: clean throughout, mismatch/k3_kl
0.0007-0.0012.

Qwen3-30B-A3B MoE, 8 GPUs, 6144 requests, max_num_seqs=256, 4096 max length.
Measured before the rebase, with the scheduler fixes only:

configuration answered correctly parse failure generation s
gpu_memory_utilization=0.30, no scheduler fixes 0.0449 47.4% 244.1
0.30, both fixes (4 runs) 0.1055 / 0.1087 / 0.1117 / 0.1159 22.5% 177-192
0.45 (pool large enough, no preemption) 0.1073-0.1174 21.8-23.4% 108-111
vLLM 0.45, same inputs 0.1069-0.1113 20.2-21.5% 151-156

Preemptions per replica per step, from get_metrics_statistics(): 2265-2596 at
0.30, 0 at 0.45. After the fixes an undersized KV pool costs only throughput
and no longer costs correctness. A 60-step DAPO training run on that diff
completed cleanly, reward accuracy 0.46-0.51 at step 60.

Notes for reviewers

Thanks @valarLip for the RMSNorm simplification — the branch that was there had
nothing to choose between, since the argument defaults to 0.

Two items from the earlier revision are no longer in the diff, because main
has since grown them independently: the 128x128 block-scale FP8 online
quantization in atom/quantization/quark/utils.py and
atom/model_ops/linear.py, and ATOM_LOG_LEVEL. There is nothing to review
for either. The Copilot reviews earlier in this thread were written against
that 12-file revision; the current head touches 19 files.

The 22 commits this branch accumulated before it merged main at 0d4e96605
are squashed into the first commit. They were written against a much older
main and several only fix up the commit before them — including several that
later commits in this series revert outright — so replaying them individually
onto today's main produces intermediate trees that never existed and
conflicts against code main has since evolved on its own. The resulting tree
is byte-identical to the pre-rebase branch tree plus the three fixes. The eight
topical commits that follow are unchanged and are the ones worth reading.

The pre-rebase head is preserved at lumen-rl-pre-rebase-20260914
(d6b9e147c) if anything needs comparing against it.

Submission Checklist

@ZhangDanyang-AMD
ZhangDanyang-AMD requested a lite review from Copilot August 26, 2026 08:15
@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 2028 --add-label <label>

@ZhangDanyang-AMD ZhangDanyang-AMD changed the title Lumen rl [Lumen-RL] Improve FP8 rollout weight synchronization and CUDA Graph stability Aug 26, 2026

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

This PR improves ATOM’s Lumen-RL integration by tightening FP8 online quantization + rollout weight sync behavior and adjusting CUDA Graph / memory policies to reduce instability, while also adding request-preprocess fan-out and configurable console log verbosity.

Changes:

  • Add/align FP8 “true 128×128 block-scale” online quantization paths and post-sync layout handling to match initial load behavior.
  • Adjust rollout weight update + sleep/wake memory behavior to keep weights/KV-cache/graphs resident in non-eager mode and change CUDA-graph recapture failure handling.
  • Add request preprocessing fan-out support and make console logging threshold configurable via ATOM_LOG_LEVEL (plus a Qwen3 logits masking hook for padded vocab).

Reviewed changes

Copilot reviewed 12 out of 12 changed files in this pull request and generated 6 comments.

Show a summary per file
File Description
tests/test_envs.py Adds coverage for ATOM_LOG_LEVEL default/override behavior.
atom/utils/envs.py Introduces ATOM_LOG_LEVEL env var with default WARNING.
atom/utils/init.py Makes console handler level configurable (but needs logger-level alignment fix).
atom/rollout/weight_updater.py Adds CUDA-graph invalidation hook and improves packed/FP8 post-processing logic.
atom/rollout/memory_manager.py Changes no-eager sleep policy and alters CUDA-graph recapture failure behavior.
atom/quantization/quark/utils.py Adds 128×128 FP8 blockscale quant routine (currently duplicated).
atom/models/qwen3.py Masks logits beyond a “true vocab size” set via env var.
atom/model_ops/utils.py Preserves parameter storage during shuffle to keep CUDA-graph pointers stable.
atom/model_ops/moe.py Adds compatibility fallback for older AITER shuffle_scale API naming.
atom/model_ops/linear.py Uses the 128×128 FP8 blockscale quant path for per_1x128 (missing import currently).
atom/model_ops/layernorm.py Enables model-sensitive RMSNorm flag in AITER calls.
atom/model_engine/llm_engine.py Switches request submission to preprocess_fanout() for n>1 support.
Suppressed comments (1)

atom/rollout/memory_manager.py:104

  • In the enforce_eager branch of _release_weights(), the condition if not self.enforce_eager and ... can never be true, so CUDA graphs will not be released before weights are discarded. This can keep GPU memory pinned unexpectedly.
        # Release CUDA graphs first — they hold references to weight memory
        # and prevent freeing GPU memory.
        if not self.enforce_eager and hasattr(self, "graphs") and self.graphs:
            self._graphs_backup_keys = list(self.graphs.keys())
            self.graphs.clear()

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread atom/rollout/memory_manager.py Outdated
Comment thread atom/model_ops/linear.py Outdated
Comment thread atom/utils/__init__.py Outdated
Comment thread atom/models/qwen3.py
Comment thread atom/quantization/quark/utils.py
Comment thread atom/rollout/weight_updater.py Outdated

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

Copilot reviewed 12 out of 12 changed files in this pull request and generated no new comments.

Suppressed comments (5)

Previously missed (2) — in code that hasn't changed since the last review.

atom/rollout/weight_updater.py:31

  • The docstring says this method “Drop[s] stale CUDA graphs”, but the implementation is a no-op for both eager and non-eager paths. Updating the docstring to match the actual behavior would prevent confusion for future maintainers.

This issue also appears on line 38 of the same file.

        """Drop stale CUDA graphs after online weight updates.

        Recapture is intentionally deferred to ``resume_memory``/wake-up, where
        MemoryManagerMixin verifies that both weights and KV cache are resident
        on GPU.  This avoids recapturing against an incomplete post-update

atom/rollout/memory_manager.py:102

  • After adding the early return for not self.enforce_eager, the CUDA-graph release block below can only run in eager mode, but its condition still checks not self.enforce_eager, so graphs would never be cleared. This can keep stale graphs alive and prevent GPU memory from being released during sleep.
            )
            return
        # Release CUDA graphs first — they hold references to weight memory
        # and prevent freeing GPU memory.
        if not self.enforce_eager and hasattr(self, "graphs") and self.graphs:

atom/rollout/weight_updater.py:42

  • This helper currently has an unconditional return in non-eager mode, which makes the CUDA-graph invalidation code below it unreachable dead code. Please either remove the unreachable block or gate it behind a condition so the function’s behavior is clear.
        # valid. Keep the graphs resident instead of dropping+recapturing them,
        # which under expandable_segments faults during post-wake graph capture.
        return

        torch.cuda.synchronize()

atom/quantization/quark/utils.py:390

  • quantize_weight_to_fp8_128x128_blockscale is defined twice in this module. The second definition overrides the first and will likely trigger Ruff/flake8 redefinition checks (e.g., F811), while also making it unclear which implementation is intended.
def quantize_weight_to_fp8_128x128_blockscale(weight, quant_dtype):
    """Quantize a 2D weight to FP8 with 128x128 block scales.

    Returns:
        q_weight: quantized weight with the same shape as input ``weight``.

atom/utils/init.py:1045

  • getLogger() now defaults the console handler to WARNING, but the logger itself is still set to INFO. That means logger.info(...) will still build LogRecords (and then be dropped by the handler), which contradicts the performance note above. Set the logger level to match the chosen handler level.
        console_handler.setFormatter(formatter)
        console_handler.setLevel(
            getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
        )

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

Copilot reviewed 11 out of 11 changed files in this pull request and generated no new comments.

Suppressed comments (7)

Previously missed (4) — in code that hasn't changed since the last review.

atom/rollout/weight_updater.py:28

  • The docstring says this method “drops stale CUDA graphs”, but the implementation returns without invalidating graphs in both eager and non-eager modes. Please update the docstring to reflect the current behavior (no-op in non-eager due to in-place weight updates) so callers aren’t misled.

This issue also appears on line 42 of the same file.

    def _invalidate_cudagraphs_after_weight_update(self) -> None:
        """Drop stale CUDA graphs after online weight updates.

        Recapture is intentionally deferred to ``resume_memory``/wake-up, where
        MemoryManagerMixin verifies that both weights and KV cache are resident

atom/rollout/weight_updater.py:372

  • This shuffle guard now skips FP8 weights with dim==3. shuffle_weights() explicitly supports both 2D and 3D parameters (e.g., MoE expert weights), so this can leave some updated FP8 weights in the wrong layout.
        if needs_shuffle and param.dim() == 2:
            shuffle_weights(param)

atom/rollout/memory_manager.py:99

  • In no-eager mode this method returns early, so the subsequent if not self.enforce_eager and hasattr(self, "graphs") ... block is now dead code and CUDA graphs won’t be released even when _release_weights() is trying to free GPU memory. Consider clearing graphs unconditionally when present (or remove the unreachable branch).

This issue also appears on line 231 of the same file.

        if not self.enforce_eager:
            logger.info(
                f"{self.label}: no-eager sleep keeps weights + CUDA graphs resident"
            )
            return

atom/models/qwen3.py:363

  • This file contains a @support_torch_compile-decorated model class (Qwen3Model). Adding env-dependent control flow and an in-place logits mask here risks breaking Dynamo/CUDA-graph capture assumptions for compiled models. Please move this “true vocab size” masking to a non-compiled call site (e.g., the sampling/logits post-processing path) instead of modifying the model implementation.
    ) -> torch.Tensor:
        logits = self.lm_head(hidden_states)
        true_vocab_size = int(os.environ.get("LUMENRL_ATOM_TRUE_VOCAB_SIZE", "0") or 0)
        if true_vocab_size > 0 and logits.shape[-1] > true_vocab_size:

atom/rollout/weight_updater.py:46

  • Code below is unreachable because the function returns unconditionally just above. This dead CUDA-graph invalidation block should be removed (or the early return removed) to avoid confusion and potential lint failures.
        torch.cuda.synchronize()
        graphs = getattr(self, "graphs", None)
        if graphs:
            self._graphs_backup_keys = list(graphs.keys())
            graphs.clear()

atom/rollout/memory_manager.py:235

  • On CUDA graph recapture failure this now raises, removing the previous fallback to enforce_eager=True. Since resume_memory() is invoked via EngineUtility without an obvious recovery path, this can turn a transient recapture issue into a hard failure of the wake-up flow. Restoring the eager fallback (or otherwise handling the exception) would improve operational robustness.
            logger.error(
                f"{self.label}: CUDA graph recapture failed: {e}",
                exc_info=True,
            )
            if hasattr(self, "_graphs_backup_keys"):

atom/utils/init.py:1045

  • console_handler now defaults to WARNING via ATOM_LOG_LEVEL, but the logger itself is still set to INFO. That means INFO log calls will still construct LogRecords and be filtered at the handler, which defeats the performance intent described in the comment above. Set the logger level to the same ATOM_LOG_LEVEL you apply to the console handler.
        console_handler.setLevel(
            getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
        )

Copilot AI review requested due to automatic review settings August 26, 2026 09:09

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

Copilot reviewed 11 out of 11 changed files in this pull request and generated 1 comment.

Suppressed comments (5)

Previously missed (1) — in code that hasn't changed since the last review.

atom/rollout/memory_manager.py:103

  • After the if not self.enforce_eager: ... return guard, control flow implies enforce_eager is True, so if not self.enforce_eager and ... is always false and CUDA graphs will never be released in the eager sleep path. Drop the redundant not self.enforce_eager check so the release logic can run when appropriate.
            return
        # Release CUDA graphs first — they hold references to weight memory
        # and prevent freeing GPU memory.
        if not self.enforce_eager and hasattr(self, "graphs") and self.graphs:
            self._graphs_backup_keys = list(self.graphs.keys())

atom/utils/init.py:1045

  • getLogger() still sets the logger level to INFO, but the console handler now defaults to WARNING via ATOM_LOG_LEVEL. That means INFO logs still build LogRecords and get filtered at the handler, contradicting the nearby comment about avoiding this overhead. Set logger.setLevel() to the same computed console_level as the handler.
        console_handler.setLevel(
            getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
        )

atom/rollout/weight_updater.py:29

  • The docstring says this method “drops stale CUDA graphs” and defers recapture, but the current implementation returns early for non-eager mode and never invalidates graphs. Update the docstring to match the intended no-eager behavior (keep graphs resident because weights are updated in-place).
        """Drop stale CUDA graphs after online weight updates.

        Recapture is intentionally deferred to ``resume_memory``/wake-up, where
        MemoryManagerMixin verifies that both weights and KV cache are resident
        on GPU.  This avoids recapturing against an incomplete post-update

atom/rollout/weight_updater.py:42

  • There is unreachable code after the early return (torch.cuda.synchronize(), graphs.clear(), empty_cache(), etc.). Keeping dead code here is error-prone and makes it unclear whether graphs should be invalidated or preserved. Remove the unreachable block (or remove the early return if invalidation is actually intended).
        # valid. Keep the graphs resident instead of dropping+recapturing them,
        # which under expandable_segments faults during post-wake graph capture.
        return

        torch.cuda.synchronize()

atom/rollout/memory_manager.py:237

  • Recapture failures now re-raise, which will propagate through resume_memory() and can take down the worker process. Unless this is intentionally fail-fast, consider restoring the previous fallback to enforce_eager=True on recapture failure so the system can keep serving (albeit without CUDA graphs) instead of crashing.
                exc_info=True,
            )
            if hasattr(self, "_graphs_backup_keys"):
                del self._graphs_backup_keys
            raise

Comment thread atom/models/qwen3.py
@zufayu
zufayu requested a review from ZhangLirong-amd August 27, 2026 01:23
Comment thread atom/rollout/weight_updater.py Outdated
Copilot AI review requested due to automatic review settings August 27, 2026 03:02

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

Copilot reviewed 11 out of 11 changed files in this pull request and generated 2 comments.

Suppressed comments (3)

atom/utils/init.py:1045

  • ATOM_LOG_LEVEL currently only changes the console handler level; the logger itself is still hard-coded to INFO. That means logger.info(...) calls will still build LogRecords and then be dropped by the WARNING handler, which undermines the goal of avoiding per-request INFO overhead by default. Consider deriving a single level from ATOM_LOG_LEVEL and applying it to both the logger and the handler.
        console_handler.setLevel(
            getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
        )

atom/rollout/memory_manager.py:104

  • In _release_weights(), the new early-return when not self.enforce_eager makes the subsequent CUDA-graph release block unreachable, and in the eager path (self.enforce_eager is True) the current condition if not self.enforce_eager and ... will never run. If self.graphs is populated while enforce_eager=True, this will prevent graphs from being cleared and can block GPU memory from being released during sleep.
        # Release CUDA graphs first — they hold references to weight memory
        # and prevent freeing GPU memory.
        if not self.enforce_eager and hasattr(self, "graphs") and self.graphs:
            self._graphs_backup_keys = list(self.graphs.keys())
            self.graphs.clear()

atom/model_ops/layernorm.py:76

  • Same compatibility concern as above for rmsnorm2d_fwd_with_add: adding use_model_sensitive_rmsnorm will crash on older AITER versions that don't support this kwarg. A try/except fallback keeps the callsite forward-compatible without changing behavior on newer AITER.
    rmsnorm2d_fwd_with_add(
        out, x, residual, residual_out, weight, eps, use_model_sensitive_rmsnorm=1
    )

Comment thread atom/rollout/memory_manager.py Outdated
Comment thread atom/model_ops/layernorm.py Outdated

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

Copilot reviewed 11 out of 11 changed files in this pull request and generated no new comments.

Suppressed comments (4)

atom/utils/init.py:1045

  • The console handler now defaults to WARNING, but the logger itself is still set to INFO earlier in this block. That means every logger.info() call still builds a LogRecord and then gets filtered by the handler, which defeats the performance intent described in the comment above. Set the logger level to the same env-derived level used by the handler (and compute it once).
        console_handler.setLevel(
            getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
        )

atom/rollout/memory_manager.py:104

  • After the new early-return for no-eager mode, this block only runs when enforce_eager=True, but the CUDA graph release conditional still checks not self.enforce_eager, making it unreachable. If self.graphs is populated, it will never be cleared here and can keep weight storage referenced while you try to discard weights.

This issue also appears on line 238 of the same file.

        # Release CUDA graphs first — they hold references to weight memory
        # and prevent freeing GPU memory.
        if not self.enforce_eager and hasattr(self, "graphs") and self.graphs:
            self._graphs_backup_keys = list(self.graphs.keys())
            self.graphs.clear()

atom/rollout/memory_manager.py:242

  • This exception path now re-raises, which will likely take down the runner process (AsyncProcRunner.call_func(wait_out=True) does not wrap exceptions into a structured result). If recapture fails in the field, falling back to eager mode is safer than crashing the whole engine, especially since this path is already best-effort recovery after sleep/wake.
                exc_info=True,
            )
            if hasattr(self, "_graphs_backup_keys"):
                del self._graphs_backup_keys
            raise

atom/models/qwen3.py:364

  • Reading and parsing LUMENRL_ATOM_TRUE_VOCAB_SIZE on every compute_logits() call adds overhead on the decode hot path. Cache the parsed value on first use (or in init) and reuse it for subsequent steps.
        logits = self.lm_head(hidden_states)
        true_vocab_size = int(os.environ.get("LUMENRL_ATOM_TRUE_VOCAB_SIZE", "0") or 0)
        if true_vocab_size > 0 and logits.shape[-1] > true_vocab_size:
            logits[..., true_vocab_size:] = float("-inf")

Copilot AI review requested due to automatic review settings August 27, 2026 05:41

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Pull request overview

Copilot reviewed 11 out of 11 changed files in this pull request and generated no new comments.

Suppressed comments (3)

atom/models/qwen3.py:365

  • This change modifies a model file that contains an @support_torch_compile-decorated model (Qwen3Model). Introducing os.environ reads and an in-place logits masking op in this file risks Dynamo/torch.compile graph breaks and violates the repo rule to avoid editing these files. Consider moving the vocab masking to a call site outside the compiled model code path (e.g., postprocess or a runner-level logits hook) and passing the true vocab size in via config/context instead of reading env vars in compute_logits().
    def compute_logits(
        self,
        hidden_states: torch.Tensor,
    ) -> torch.Tensor:
        logits = self.lm_head(hidden_states)
        true_vocab_size = int(os.environ.get("LUMENRL_ATOM_TRUE_VOCAB_SIZE", "0") or 0)
        if true_vocab_size > 0 and logits.shape[-1] > true_vocab_size:
            logits[..., true_vocab_size:] = float("-inf")
        return logits

atom/rollout/memory_manager.py:239

  • _recapture_cudagraphs_if_needed() now raises on recapture failure, which can make resume_memory() fail hard and potentially bring down the service during sleep/wake scenarios. Previously this path fell back to enforce_eager=True to preserve availability. If recapture is still a possible runtime path, consider restoring the eager fallback (or gating the raise behind a strict/debug option).
        except Exception as e:
            logger.error(
                f"{self.label}: CUDA graph recapture failed: {e}",
                exc_info=True,
            )

atom/utils/init.py:1045

  • The console handler level is now driven by ATOM_LOG_LEVEL, but the logger itself is still hard-coded to INFO. This means INFO logs still create LogRecords (and get discarded by the WARNING handler), which contradicts the comment about avoiding per-request logging overhead. Set the logger level to the same resolved level as the handler so filtering happens before LogRecord creation.
        console_handler.setLevel(
            getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
        )

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review is ineligible. To be eligible to request a review, you need a paid Copilot license, or your organization must enable Copilot code review.

@valarLip

Copy link
Copy Markdown
Collaborator

Reviewed at 44054b3d3. This is a selection, not an exhaustive list — the items below are the ones I judged worth acting on.

Provenance: [verified] means I read the code at the cited line myself (and for #6, ran the failure on this box). [reported] means the mechanism is traced through the diff but I did not open the far end of the call chain — worth confirming before acting.

Cleared, so nobody re-litigates them: is_final_chunk index alignment is sound (req_ids = list(seqs.keys()) is the same dict); the weight_is_stored_preshuffled truth table is line-for-line equivalent to the inlined original; aiter's shuffle_weight ends in .contiguous() so copy_ is not self-aliasing; the unquantized MoE load does shuffle, so _finalize_expert_weight_sync is re-establishing a real layout; RLHFModelRunner.postprocess matches the base signature parameter-for-parameter; no @support_torch_compile file is touched; and there is no collection-time ImportError anywhere, so the CI run is not at risk.


The three headline fixes are each undone one or two functions away, in the same file

That is the pattern worth looking at before the individual items:

# (a) A fence was added -- but on the FP8 path it fires AFTER the write it fences,
#     and four other writes have none at all.
param.data.copy_(tensor)                        # weight_updater.py:677, the overwrite
self._post_process_fp8_weight(module, param)    # :678 -- _await_readers_of lives in here

# (b) The `param.data` rebind was removed from shuffle_weights, because a live
#     decode graph holds the old address. One function away, on the same path:
param.data, weight_scale.data, _ = normalize_e4m3fn_to_e4m3fnuz(param.data, ...)   # :605

# (c) Nine lines of comment explain the EOS-placeholder bug in preempt() -- and only
#     the non-speculative branch is fixed.
num_placeholder = self.mtp_k
if is_deferred_out: num_placeholder += 1        # scheduler.py:2612-2614
...
if self.spec_decode_local and self.mtp_k > 0:
    strip = self.mtp_k + seq.num_rejected       # :2399 -- the old formula, unchanged
else:
    strip = seq.num_placeholder_tokens

1. [verified] preempt() still leaves exactly one EOS on an MTP engine — the case the commit exists to fix

The comment block at scheduler.py:2390-2398 states the failure precisely:

The placeholder is eos_token_id ... Left in place it stops being a placeholder: the recompute prefills a context ending in <|endoftext|> and the model duly starts a new document, and the same token is handed back to the caller as generated output.

Then:

if self.spec_decode_local and self.mtp_k > 0:
    strip = self.mtp_k + seq.num_rejected
else:
    strip = seq.num_placeholder_tokens

postprocess sets the authoritative width at :2612-2614:

num_placeholder = self.mtp_k
if is_deferred_out:
    num_placeholder += 1

and is_deferred_out is pipeline_parallel_size == 1 (model_runner.py:189) — true on every TP-only engine, which is the normal MTP deployment.

So with num_rejected == 0, mtp_k + 1 placeholders exist and the branch strips mtp_k: one eos_token_id survives, exactly as the comment describes. With num_rejected = r > 0 it over-strips by 2r - 1, deleting real tokens and — new in this PR — real logprobs.

del seq.output_tokens[-strip:] is also unclamped: with strip > len(output_tokens) Python clears the list and seq.num_tokens -= strip drops below num_prompt_tokens, so the sequence re-prefills a truncated prompt and num_completion_tokens goes negative.

2. [verified] _await_readers_of covers 2 of ~6 in-place writes, and on the FP8 path it fires after the write

$ grep -n "_await_readers_of" atom/rollout/weight_updater.py
468:    def _await_readers_of(self, param: torch.nn.Parameter) -> None:
541:        self._await_readers_of(param)
597:        self._await_readers_of(param)

:597 is inside _post_process_fp8_weight, which update_weights calls at :678 — one line after param.data.copy_(tensor) at :677. The first overwrite still races in-flight decode replays. The SHM (:800) and IPC (:956) paths have the same shape.

Completely unfenced: the bf16 param.data.copy_ at :683/:804/:962, _try_shard_weight at :464, the weight_loader(param, tensor) fallback at :686, and the entire new MoE expert path (weight_loader at :262/:321, in-place shuffle_expert_slices at :422).

That last one matters most: _check_expert_sync_supported requires an unquantized MoE, so the headline feature of this PR can never reach the helper that carries the wait. By this PR's own measured mechanism (5 of 7 syncs lost), the graph reads a half-permuted weight and every sequence past that point is token soup with nothing raised.

3. [verified] _post_process_fp8_weight still rebinds param.data, on the online-update path

# weight_updater.py:605
param.data, weight_scale.data, _ = normalize_e4m3fn_to_e4m3fnuz(param.data, weight_scale.data)

This is the exact address move the PR removed from shuffle_weights — and here there is nothing to recapture, because the decode graph is live and holds the old address. Per this PR's own evidence for the shuffle_weights fix, the graph then replays against the stale buffer: 192 of 192 sequences come back as !!!!!!!!.

Gated on need_normalize_e4m3fn_to_e4m3fnuz, i.e. the e4m3fnuz per-token path, which is the normal ROCm FP8 configuration.

4. [verified] Fix-then-sweep: the offload-resume admission branch still sizes with num_prompt_tokens

# scheduler.py:1306
if offload_resume:
    num_new_tokens = seq.num_prompt_tokens - seq.num_cached_tokens

versus its sibling fifty lines down, which already carries the reason:

# scheduler.py:1354-1358
# ... so num_tokens > num_prompt_tokens and those tokens still need KV recomputed.
num_new_tokens = seq.num_tokens - num_cached_blocks * self.block_manager.hash_block_size

Before this PR a sequence could not be is_partial_prefill with num_cached_tokens >= num_prompt_tokens; after it, that is the normal state of a preempted sequence. Let the offload connector park such a sequence and re-admit it: num_new_tokens <= 0 → _prefill_chunk_for_budget returns None → waiting.appendleft(seq); break every tick forever, starving every request behind it. (Phase 1 at :1221 uses num_tokens, so it never rescues it.) With chunked prefill off, _assert_positive_prefill_chunk (:2127) raises in the engine loop instead.

5. [verified] release_cudagraphs leaves the TBO graph store holding the freed KV pool

$ grep -rn "tbo_graphs" atom/rollout/memory_manager.py atom/model_engine/model_runner.py
atom/model_engine/model_runner.py:979:            self.model.tbo_graphs.clear()

ModelRunner.exit() clears both stores; release_cudagraphs clears only runner.graphs. This PR newly calls release_cudagraphs from _release_kv_cache, so with TBO on, sleep(level=1) now frees the pool while the TBO graphs survive holding its base and the old weight addresses — and _recapture_cudagraphs_if_needed keys entirely off _graphs_backup_keys, so it never rebuilds them.

The first two-batch-overlap forward after wake replays against reallocated memory: use-after-free / memory access fault, or silently stale KV. Not reachable before this PR, because level-1 sleep did not release graphs at all.

6. [verified] The new 2D weight.copy_(shuffled) is not dtype-portable, where the rebind it replaced was

Ran on this box (torch 2.9.1+rocm7.1.1):

CPU fp4 copy_: NotImplementedError: "copy_kernel" not implemented for 'Float4_e2m1fn_x2'

CUDA is fine. linear.py:899/911 call shuffle_weights on exactly that dtype for MXFP4 checkpoints. Any path where the weight is not yet on GPU — CPU weight prep, a CPU unit test, param_device="meta" from _stream_online_quant (linear.py:532) — now either raises or, on meta, silently no-ops while tensor.is_shuffled = True is still set at :177, so the kernel reads an unshuffled weight through the preshuffle permutation.

The new test only exercises bfloat16 and is CUDA-gated. Falling back to the rebind when copy_ is unsupported would keep both properties.

7. [verified] weight_is_stored_preshuffled compares QuantType with ==, where the code it replaces deliberately used .value

# linear.py:455/461/464
if quant_type == QuantType.per_Token:

versus the documented reason, in this same repo:

# attention_residual.py:49-51
# QuantType is compared by .value throughout ATOM: the enum can be re-imported
# under a different module identity, which breaks `is`/`==` on the members.
if getattr(quant_type, "value", None) != QuantType.per_Token.value:

QuantType is a pybind enum from the compiled aiter.jit.module_aiter_core, and atom/quant_spec.py resolves it lazily for the same reason. If a module's quant_type came from a differently-identified aiter import (plugin process, re-import, lazy proxy), every branch falls through to return False; on the update path (weight_updater.py:617) the FP8 weight is then written and never re-shuffled, and the preshuffle GEMM reads a row-major weight. Silent token soup after the first sync. The old .value comparison was immune.


8. [reported] An expert tensor sent under ATOM's own fused names bypasses the entire new relayout machinery

_get_param_to_module_mapping keys on named_modules × named_parameters, so model.layers.N.mlp.experts.w13_weight is a key — such a name never reaches _apply_unmatched_weight. It takes elif tensor.shape == param.shape: param.data.copy_(tensor) (:681) for bf16, or _post_process_fp8_weight, which skips the shuffle because param.dim() == 3.

UnquantizedFusedMoEMethod.process_weights_after_loading (moe.py:884) had left aiter's per-expert 16×16 permutation there; it is now plain row-major, _pending_expert_relayout stays empty so finalize returns at if not pending, and the sync reports updated=1 per layer. _check_expert_sync_supported never runs on this route either, so an EP or quantized MoE is not rejected.

9. [reported] _pending_expert_relayout is cleared only on the success path

_finalize_expert_weight_sync shuffles buffer A at :422, then raises the missing-shard RuntimeError at :415 for buffer B — pending.clear() at :428 is never reached and A's entry survives with A already shuffled. Same for the NotImplementedErrors in _check_expert_sync_supported mid-sync, and for SHM/IPC syncs that abort before is_last (:831/:989).

The next successful sync re-shuffles those slices. Per the function's own docstring: "shuffling an already-shuffled slice does not undo the first shuffle, it produces a third layout." Silent, with updated=N logged as success.

Needs try/finally, and a non-empty pending at sync start should be treated as evidence of a prior failure rather than ignored.

10. [reported] Adding release_cudagraphs to _release_kv_cache makes the default sleep recapture — the operation this PR's opt-in exists to avoid — and a failed recapture then disables the opt-in permanently

sleep_keeps_memory_resident defaults to False, so a default-configured rollout under PYTORCH_CUDA_ALLOC_CONF=expandable_segments now takes a recapture on every level-1 sleep/wake where before it took none, and hits the fault the opt-in exists to avoid.

If capture_cudagraph() throws instead, the failure handler at :321 sets self.enforce_eager = True permanently — and sleep_keeps_memory_resident() reads enforce_eager first, so from that point the runner is pinned to eager and the operator's sleep_keeps_memory_resident=True stops taking effect, behind one logger.warning.

The graph-release fix itself is right; it needs to be paired with the opt-in being reachable, or with detecting expandable_segments at startup.


Test coverage

Measured by running the suite with HIP_VISIBLE_DEVICES="": 89 of the ~108 new tests never execute in CI. Four of the seven new files skip at module level on not torch.cuda.is_available(), including the two that cover the flagship expert-sync and shuffle-layout work. The "239 passed" in the PR body can only come from a GPU box.

The two scheduler commits — the highest-blast-radius changes in the branch, and where #1 and #4 live — have zero tests: TestPreempt never sets num_placeholder_tokens, and every TestPostprocess case calls postprocess(seqs, output) without the batch kwarg, so the new final[i] branch is never entered.

The skip guards are correctly placed above the aiter imports, so there is no collection-abort risk to the wider CI run.

xysheng-AMD and others added 2 commits September 15, 2026 03:04
`preempt()` re-derived the number of trailing placeholder slots as `mtp_k +
seq.num_rejected`. `postprocess` appends `mtp_k + is_deferred_out -
num_rejected` of them and records the count in `seq.num_placeholder_tokens`,
so the two agree only when deferred output is off AND nothing was rejected.

`is_deferred_out` is `pipeline_parallel_size == 1`, true on every TP-only
engine, which is the normal MTP deployment. There the speculative branch
strips `mtp_k` of `mtp_k + 1` placeholders and one `eos_token_id` survives
into the recompute -- the exact failure the nine lines of comment above it
describe, on the branch they do not fix. The re-prefill then feeds the model a
context ending in `<|endoftext|>`, so it starts a new document, and the same
token is handed back as generated output; under `ignore_eos=False` the request
stops there, which reads as a coherent answer that ends before it answers
anything.

With `num_rejected = r > 0` the same formula errs the other way and strips
`2r - 1` too many, deleting real tokens and -- new in this series -- real
logprobs. And `del seq.output_tokens[-strip:]` was unclamped: an oversized
`strip` clears the list outright, `seq.num_tokens` drops below
`num_prompt_tokens`, and `num_completion_tokens` goes negative.

Both branches now read `seq.num_placeholder_tokens`, which is the only width
that describes what is actually there, bounded by what is present and by the
prompt. The P/D first-decode path appends the remote's drafts to those same
trailing slots and recorded nothing, so it records the count too -- the remote
may send fewer than `mtp_k`, which the old formula over-stripped.

`TestPreempt` never set `num_placeholder_tokens` and the commit that
introduced the strip shipped no test at all; 8 of the 9 added here fail
without this change.

Co-authored-by: Cursor <cursoragent@cursor.com>
`is_final_chunk` is measured against the admitted length and `postprocess`
reads it rather than re-deriving it, both from the commit that ended a chunked
prefill at the admitted length. Neither had a test: every `TestPostprocess`
case called `postprocess(seqs, output)` with no `batch=`, so the `final[i]`
branch was never entered, and nothing drove a recompute whose prefill runs
past the prompt boundary.

Four cases, no GPU: the two-chunk split of a re-admitted sequence and the
frozen `is_final_chunk` it produces, `postprocess` honouring `final[i]` where
`num_cached_tokens < num_prompt_tokens` would call a middle chunk final, the
other side of that branch, and the length fallback for a caller that hands a
batch without the field. The first two fail against the pre-fix scheduler.

Co-authored-by: Cursor <cursoragent@cursor.com>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Unresolved critical scheduler, CUDA Graph lifecycle, and weight-synchronization issues remain.

Get a fresh assessment by requesting another Copilot review.

Review details

Suppressed comments (2)

atom/model_engine/scheduler.py:1323

  • When an offload load covers the entire sequence (num_cached_tokens == seq.num_tokens, which can happen when LMCache has the full context), num_new_tokens is 0. _prefill_chunk_for_budget returns None for this value, so this branch puts the sequence back at the head and breaks; every scheduling tick repeats and no decode is ever scheduled. Handle a fully loaded request by transitioning it to decode or bypassing the offload-prefill branch before calculating a prefill chunk.
                num_new_tokens = seq.num_tokens - seq.num_cached_tokens
                budget_remaining = self.max_num_batched_tokens - num_batched_tokens
                chunk = self._prefill_chunk_for_budget(
                    num_new_tokens, budget_remaining, num_batched_tokens
                )

atom/rollout/memory_manager.py:363

  • On a partial TBO recapture, capture_cudagraph() may already have stored earlier entries in model.tbo_graphs when a later shape fails. This fallback removes the normal graph backup but never clears that parallel TBO store; after switching to eager those graph/context/output objects are no longer replayed or recaptured and can keep the private CUDA graph pool allocated. Clear model.tbo_graphs along with the regular graph state here.
            if hasattr(self, "_graphs_backup_keys"):
                del self._graphs_backup_keys
            logger.warning(f"{self.label}: Falling back to enforce_eager=True")
  • Files reviewed: 21/21 changed files
  • Comments generated: 6
  • Review effort level: Lite

Comment thread atom/rollout/memory_manager.py Outdated
Comment thread atom/rollout/weight_updater.py
Comment thread atom/rollout/weight_updater.py
Comment thread atom/rollout/weight_updater.py
Comment thread atom/rollout/weight_updater.py
Comment thread atom/rollout/memory_manager.py Outdated
@xysheng-AMD

Copy link
Copy Markdown
Contributor

@valarLip The central observation is correct and holds for all three headline fixes: a
fence placed after the write it guards, a param.data rebind removed from one
function and left in the next, and a nine-line comment describing a bug on the
branch that does not fix it. The pattern, not the individual lines, is treated
as the finding.

Seven items fixed, three answered. Two of the answered ones carry a mechanism
correction: the reported cause is not the one that bites, and the actual
defect is worse in both cases. Every item was also run against the unfixed
tree — a regression test that passes without its fix is not one.


Fixed

1 — preempt() leaves one EOS on an MTP engine. Both branches now strip
seq.num_placeholder_tokens, the count postprocess records where it
appends. mtp_k + num_rejected matches that only with deferred output off and
nothing rejected; at num_rejected = r it over-strips by 2r - 1 real tokens
and their logprobs. The del is clamped by what is present and by the prompt.
The P/D first-decode path appends the remote's drafts to the same slots and
recorded nothing, so it records the count too — the remote may send fewer than
mtp_k, which the old formula over-stripped.

2 — _await_readers_of covers 2 of ~6 writes. Correct on all three
points, including that _check_expert_sync_supported requires an unquantized
MoE, so the expert path can never reach the helper carrying the wait.

Four more call sites would not fix it, since a forgotten call is the failure
mode. Every in-place write to a live parameter now goes through
_copy_into_param or _load_into_param, which wait and then write. No bare
param.data.copy_ or weight_loader(param, ...) remains on the named-tensor,
SHM or IPC path; the one bare loader call left writes a local float32
accumulation buffer, which has no readers.

4 — offload-resume sizes with num_prompt_tokens. Both failure modes are
reachable with chunked prefill on, one token apart. Above the boundary the
width is negative and _assert_positive_prefill_chunk raises; on the boundary
it is zero, reported as None, so the sequence returns to the head of
waiting and the loop breaks — the silent starvation. Now seq.num_tokens,
matching Phase 1 above and the sibling branch below.

6 — weight.copy_(shuffled) is not dtype-portable. Accepted. The symptom
is version-bound — on torch 2.10 in the validation image, fp4 copy_ works on
CPU, CUDA and meta — but the property holds: the rebind needed no copy kernel
and the in-place write does. It falls back to the rebind on
NotImplementedError. A weight shuffled on a device with no copy kernel for
it is not one a captured graph replays against, so no address needs
protecting.

7 — QuantType compared with ==. Changed to .value. One correction to
the premise, since it is load-bearing for later readers: the replaced code did
not uniformly use .value. In main, linear.py reaches the load-time
decision with == (524/544/550/564) and uses .value only at runtime (942,
956, 973, 983, 995, 1001); weight_updater.py uses .value throughout. So
weight_is_stored_preshuffled unified a == call site with a .value one
and kept == — a regression for the sync side specifically, not a reversal of
a deliberate choice. The concern is unaffected.

Behaviour-neutral where the enum has one module identity, as in the validation
image: the 30-row truth table is identical before and after. Insurance against
the re-import case, not a live change there.

8 — an expert tensor under ATOM's own fused name bypasses the relayout.
w13_weight is a real parameter of the FusedMoE, so a trainer mirroring
ATOM's state dict resolves in _get_param_to_module_mapping and takes the
plain copy_: row-major bytes into a buffer the kernel reads through the
permutation, _pending_expert_relayout empty so finalize returns at if not pending, updated counting it as success, and _check_expert_sync_supported
never running. Now routed through the expert path, with the support check and
every slice registered.

Ships inside the fence commit: in all three dispatch blocks the new branch and
the conversion of the copies below it are the same hunks.

9 — _pending_expert_relayout cleared only on the success path. Correct,
including the consequence — the next sync shuffles an already-shuffled slice
into a third layout, with updated=N logged as success.

The blanket try/finally clear trades one silent wrong answer for another: an
entry the loop never reached describes a buffer that is row-major now, and a
later sync re-establishing its layout is the only fix. So each entry is
dropped as its own shuffle completes, inside a finally, and anything still
pending is logged at error.

The second suggestion — treating a non-empty pending at sync start as
evidence of prior failure — was left out. For the bucketed SHM/IPC paths that
is the normal case: it is how an expert whose w1 and w3 arrive in different
buckets is relaid out exactly once. The error log at the end of a failed
finalize is the same signal without the false positives.


Answered, mechanism corrected

3 — _post_process_fp8_weight rebinds param.data. A real bug, fixed,
but not for the stated reason. normalize_e4m3fn_to_e4m3fnuz does not move
the weight's address: it fixes the NaN bytes through an int8 view of the same
storage and returns that storage with a reinterpreted dtype, so data_ptr()
is unchanged and a captured graph cannot observe the rebind. Measured:

call 1: weight ptr unchanged = True  | dtype torch.float8_e4m3fnuz
call 1: scale  ptr unchanged = False | value 6.0
call 2: scale value = 12.0 (started at 3.0)

Two defects there are worse than an address move. weight_scale is freshly
allocated per call, so that address does move out from under the graph — the
described hazard, on the other buffer. And
need_normalize_e4m3fn_to_e4m3fnuz is a static property of the layer
(params_dtype == torch.float8_e4m3fnuz, set once in create_weights) that
nothing clears after the load-time conversion, so the conversion re-ran every
sync on an already-converted parameter and doubled the scale again each time.
After N syncs the dequantized weight is 2**N too large.

The weight needs no conversion on this path: _requantize_fp8_weight
quantizes into param.dtype against finfo(e4m3fnuz).max, and the
direct-copy path receives bytes already in param.dtype. Gated on
param.dtype == torch.float8_e4m3fn now, with the scale written in place.

Inert where the FP8 smoke runs — the flag is False on gfx950 — so the
end-to-end numbers neither confirm nor refute it. Unit tests cover it,
including 3.0 → 3.0 → 3.0 across three syncs.

5 — release_cudagraphs leaves the TBO graph store holding the freed pool.
The gap is real and is fixed in the same place, so both release paths now match
ModelRunner.exit().

The failure mode is not use-after-free. Under TBO the replayable handle still
lands in runner.graphs — capture_cudagraph stores what capture_tbo_graph
returns — so clearing that stops the replay; tbo_graphs is a parallel
keep-alive store for the per-ubatch contexts and the captured output. What
survives is a pin on the graph's private memory pool, precisely the footprint
the caller went to sleep to reclaim, held until a recapture overwrites the same
key.

10 — the default sleep now takes a recapture. Agreed on both halves,
including that the graph release is right. The first half is the intended
trade: before it, a level-1 sleep freed the KV pool and left the graphs to
replay against it, which is wrong unconditionally; after it the default takes
a recapture, which faults only under expandable_segments.

The second half is now addressed where it can be acted on rather than where it
is discovered. Releasing graphs with expandable_segments set and the opt-in
off warns before the wake and names the option; the existing handler speaks
only after the recapture has failed. Since that handler sets
enforce_eager = True permanently while sleep_keeps_memory_resident() reads
enforce_eager first, an operator who set the option silently stops getting
it — self-consistent, with no graphs left to keep valid, but not the
configured state, so it is logged. Startup detection was not added: the
release site knows both facts, startup knows only one.


Test coverage

The strongest item in the review. The two scheduler commits had no tests:
TestPreempt never set num_placeholder_tokens and no TestPostprocess case
passed batch=, so the final[i] branch was never entered. Item 1 lived in
one of those commits.

Added, all runnable without a GPU, with the count that fails unfixed:

test covers fails unfixed
TestPreemptStripsExactlyThePlaceholders item 1: strip width vs. what postprocess appended, over mtp_k and num_rejected, plus the clamp and the P/D drafts 8 of 9
TestChunkedPrefillFinality finality against the admitted length; postprocess reading final[i] 2 of 4
TestOffloadResumeAdmission item 4, including a request starved behind the resume 4 of 4
test_weight_sync_inplace_ordering.py (rewritten) item 2, every write path separately 9 of 13
test_weight_sync_expert_routing.py (new) items 8 and 9 8 of 11
test_rollout_memory_manager_sleep.py (extended) item 5's TBO store, item 10's two warnings 5 of 21

The ordering file was rewritten because its old tests asserted that a wait
happened — which the shipped bug satisfied, one line after the write. Each
test now records the buffer's contents when its fence fires and asserts it
still held the old bytes.

On the module-level skips: torch.cuda.is_available() is the wrong predicate,
but the correct one does not help. atom.model_ops.linear and
atom.model_ops.utils import aiter, whose utility.dtypes calls
get_gfx_runtime() at module scope, and that function ignores GPU_ARCHS and
always shells out to rocminfo. Those modules cannot be imported on a CPU
runner at all; a GPU CI job would fix it, or making aiter.utility.dtypes
lazy upstream. What did move is everything that never needed aiter: measured
in the ROCm validation image with HIP_VISIBLE_DEVICES=, the weight-sync
files go from 60 to 84 executed and the scheduler group from 299 to 316.
Run non-GPU unit tests is green on the pushed head.

The "no collection-abort risk" note is correct. Worth recording how it can
look otherwise: with HIP_VISIBLE_DEVICES= on a box where aiter is
installed, pytest.importorskip("aiter") succeeds and the import then fails
inside aiter on the masked rocminfo, so correctly guarded files present as
collection errors. CI's runner has no aiter and the same guard skips cleanly.


Verification

  • 8-step ATOM FP8 DAPO smoke: 0 of 7 syncs corrupted. Eight steps rather than
    three, because three samples the event twice — the first version of the
    fence fix was clean over three and lost five of seven over eight. Run twice
    on this tree and once on the pre-round tree, since the clean band is wider
    run to run than one run suggests: maxima of 0.258, 0.015 and 0.024 against a
    corruption threshold of 1. The spread is variance, and cannot be otherwise
    in this configuration — the e4m3fnuz gate is inert, the QuantType table is
    unchanged, nothing is preempted, and the model is dense.
  • Count probe, four pool sizes: 64/64 clean and 0 drifted in all four,
    preemptions 0 / 84 / 200 / 84.
  • CPU scheduler simulator at both pool sizes: no invariant broken.
  • Each of the nine commits is green on its own, with the executed count
    growing monotonically (453 before the round, then 462, 466, 470, 482, 486,
    490, 500, 500, 500).
  • Black clean repo-wide; ruff reports nothing within the diff context.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 Needs a closer look

Unresolved moderate issues remain in CUDA Graph invalidation and TP-aware FP8/MoE weight synchronization.

Review details

Suppressed comments (11)

atom/rollout/memory_manager.py:77

  • PIECEWISE captures are not stored in runner.graphs: each CUDAGraphWrapper keeps its concrete_cudagraph_entries, and the piecewise capture path can therefore leave this dictionary empty. This guard returns without setting _graphs_backup_keys or invalidating those graphs (and their graph pools), so a level-1 sleep can free the KV cache and wake will skip recapture, replaying graphs against freed storage. Invalidate the piecewise stores and mark them for recapture here, or apply the resident-memory path to them as well.
    if not getattr(runner, "graphs", None):
        return

atom/rollout/memory_manager.py:80

  • Standard single-rank captures keep their output tensors in runner.graph_logits separately from runner.graphs (see ModelRunner.capture_cudagraph). Clearing only runner.graphs leaves these captured output buffers referenced during sleep, so empty_cache() cannot reclaim their graph-private memory and wake can recapture on top of it. Clear graph_logits whenever the graph stores are invalidated, including when runner.graphs is empty.
    if not getattr(runner, "graphs", None):
        return
    runner._graphs_backup_keys = list(runner.graphs.keys())
    runner.graphs.clear()
    runner.graph_pool = None

atom/rollout/weight_updater.py:325

  • With TP > 1, this sends each fused 3D chunk directly to FusedMoE.weight_loader(). Rank-3 inputs set full_load=True, and _load_w13()/_load_w2() only narrow the local destination; they do not slice the full source by self.tp_rank. For example, TP=2 with a full (E, 2I, H) gate_up_proj produces (E, I, H) chunks while each rank's half-buffer is only (E, I/2, H), so the load overruns/fails instead of updating the rank-local weights. Slice the fused source along dim 1 for w1/w3 and dim 2 for w2 before invoking the full-load loader, and add a TP>1 regression test.
        for shard_id, chunk in zip(shard_ids, gpu.chunk(len(shard_ids), dim=1)):
            self._load_into_param(
                param,
                weight_loader,
                chunk,

atom/rollout/weight_updater.py:818

  • update_weights documents named_tensors as full, unsharded weights, and _requantize_fp8_weight explicitly narrows such tensors for world_size > 1. This same-dtype FP8 branch bypasses both that narrowing and _try_shard_weight, then copies the full tensor directly into the rank-local parameter, so a quantized full tensor on TP > 1 raises a shape mismatch instead of updating the rank. Apply the same TP slicing before _copy_into_param, then run the layout post-processing on the local shard.
            elif self._is_fp8_param(module, param) and tensor.dtype == param.dtype:
                tensor = tensor.to(device=self.device)
                self._copy_into_param(param, tensor)
                self._post_process_fp8_weight(module, param)

atom/rollout/weight_updater.py:944

  • The SHM update path repeats the same TP hole as update_weights: its full-tensor contract is documented above, but the same-dtype FP8 branch copies directly into the rank-local parameter and bypasses both _requantize_fp8_weight's TP narrowing and _try_shard_weight. With TP > 1 and an FP8 tensor from the trainer, this path raises on the full/local shape mismatch. Narrow the tensor for this rank before copying and post-processing.
                elif self._is_fp8_param(module, param) and tensor.dtype == param.dtype:
                    tensor = tensor.to(device=self.device)
                    self._copy_into_param(param, tensor)
                    self._post_process_fp8_weight(module, param)

atom/rollout/weight_updater.py:1105

  • The IPC update path also bypasses TP slicing for same-dtype FP8 inputs: the documented full tensor is copied directly into the rank-local parameter here, while only the dtype-mismatch requantization path narrows by world_size. A TP rollout receiving already-quantized full weights therefore fails with a shape mismatch instead of updating its shard. Apply the same per-rank narrowing before _copy_into_param and then run the layout post-process.
            elif self._is_fp8_param(module, param) and tensor.dtype == param.dtype:
                self._copy_into_param(param, tensor)
                self._post_process_fp8_weight(module, param)
                updated += 1

atom/rollout/weight_updater.py:371

  • The documented full/unsharded input contract is also broken for the ATOM-named w13_weight/w2_weight route: this exact-shape check compares a global trainer tensor with the local TP-sharded parameter and rejects it before any TP slicing. On TP > 1 this raises NotImplementedError for a valid full buffer, unlike the normal per-shard loader path. Author must make this route TP-aware before copying and relayouting the local slices, or explicitly reject full buffers under TP.
        if tensor.shape != param.shape:
            raise NotImplementedError(
                f"{self.label}: {name} resolves to the fused expert buffer "
                f"{tuple(param.shape)} but arrived as {tuple(tensor.shape)}. "
                f"Re-establishing the layout works on whole expert slices, so "

atom/rollout/weight_updater.py:358

  • This docstring says the ATOM-named route skips _check_expert_sync_supported and therefore does not refuse quantized or expert-parallel MoEs, but the very next implementation line invokes that check. The stale contract can mislead callers and reviewers about which combinations are supported. Author must update the description to match the check.
        ``_check_expert_sync_supported``, so a quantized or expert-parallel MoE
        is not refused on this route either.

atom/rollout/weight_updater.py:356

  • The docstring also describes the pre-change behavior: it says this path uses a plain row-major copy, leaves _pending_expert_relayout empty, and returns without a layout step. The implementation now registers every expert in _pending_expert_relayout at lines 376-380 so the fused buffer is shuffled before success is reported. Author must update this description to reflect the current copy-and-relayout flow.
        reaches ``_apply_expert_weight``. Down the plain dispatch that is a
        row-major ``copy_`` into a buffer the kernel reads through aiter's
        16x16 expert permutation, with ``_pending_expert_relayout`` left empty
        so ``_finalize_expert_weight_sync`` returns at ``if not pending`` --
        and ``updated`` counting it as a success. It also skips

atom/rollout/weight_updater.py:289

  • This docstring claims the 3D full-load path still narrows the intermediate dimension by tp_rank, but FusedMoE._load_w13/_load_w2 explicitly skip the TP slicing branch when load_full=True and only narrow the destination from offset zero. That false claim obscures the TP failure in this new route. Author must document the actual rank-local-input requirement or update the implementation to slice before the full-load call.
        a 3D ``loaded_weight`` puts the loader on its full-load path, where the
        expert dimension is written whole and the intermediate dimension is
        still narrowed by TP rank.

tests/test_weight_sync_shuffle_layout.py:98

  • This test changes ATOM_USE_TRITON_GEMM after atom.model_ops.linear has already been imported, but linear.py imports gemm_a8w8_triton only inside the import-time use_triton_gemm() branch. When the module was collected with the default disabled, gemm_a8w8_triton remains None, so the triton_gemm parameterization never exercises the available-Triton branch that this layout decision is meant to cover. Author must load/reload the module with the env enabled, or inject a non-None Triton sentinel, so both available and unavailable cases are actually tested.
    monkeypatch.setenv("ATOM_USE_TRITON_GEMM", "1")
    expected = linear_mod.gemm_a8w8_triton is None
    assert weight_is_stored_preshuffled(QuantType.per_Token, dtypes.fp8) is expected
  • Files reviewed: 21/21 changed files
  • Comments generated: 0 new
  • Review effort level: Lite

xysheng-AMD and others added 7 commits September 15, 2026 05:46
The offload-resume admission branch sized its chunk with
`seq.num_prompt_tokens - seq.num_cached_tokens`, fifty lines above a sibling
that already carries the reason for using `num_tokens` and below a Phase 1
that already uses it.

A sequence re-admitted after `preempt` owes KV for the tokens it had already
generated, and `_mark_offload_load_ready` sets `num_cached_tokens` to whatever
the tier returned, which is not bounded by the prompt. Before the preemption
fixes in this series a sequence could not be `is_partial_prefill` with
`num_cached_tokens >= num_prompt_tokens`; after them that is the normal state
of one.

Both failure modes are reachable with chunked prefill on, one token apart.
Above the boundary the width is negative, `_prefill_chunk_for_budget` passes
it straight through, and `_assert_positive_prefill_chunk` raises in the engine
loop. Exactly on the boundary it is zero, which that helper reports as None,
so the sequence returns to the head of `waiting` and the admission loop
`break`s -- every tick, forever, starving every request queued behind it.
Phase 1 never rescues it either, because it scans `running` and the sequence
is in `waiting`.

All four added tests fail without this change, including the one that asserts
a request behind the resume still gets admitted.

Co-authored-by: Cursor <cursoragent@cursor.com>
…ight

`_await_readers_of` was added to two call sites and needed to be at six. On
the FP8 path it fired one line *after* the `param.data.copy_` it fences,
because `update_weights` calls `_post_process_fp8_weight` -- which carried the
wait -- after the overwrite rather than before it. The first write of every
sync therefore still raced the decode replays still in flight from the
generation step that had just ended.

Entirely unfenced: the bf16 `param.data.copy_` on all three entry points,
`_try_shard_weight`, the `weight_loader(param, tensor)` fallback, the packed
shard loader, and the whole routed-expert path. That last one matters most:
`_check_expert_sync_supported` requires an unquantized MoE, so the experts can
never reach `_post_process_fp8_weight` and the headline feature of this series
could not be fenced at all.

The fix is not four more calls. "Remember to call the fence" is what failed,
so every in-place write to a live parameter now goes through
`_copy_into_param` or `_load_into_param`, which wait and then write. No bare
`param.data.copy_` or `weight_loader(param, ...)` remains on any entry point;
the one bare loader call left writes into a local float32 accumulation buffer,
which has no readers. The layout shuffle in `_post_process_fp8_weight` gets
one too: it is the second in-place write to a weight the caller has just
overwritten, and the one the PR measured a decode graph catching half done.

Measured over an 8-step FP8 DAPO smoke, which contains seven weight syncs: no
wait loses 7 of 7, one wait per update at the entry points loses 5 of 7, and
this loses 0 of 7. Three steps cannot tell those apart -- it samples the event
twice, and the first version of this fix was clean over three.

`test_weight_sync_inplace_ordering.py` is rewritten around that: its old tests
asserted a wait *happened*, which the shipped bug satisfied. Each test now
records the buffer's contents at the moment its fence fires and asserts it
still held the old bytes, per write path. 9 of the 13 fail without this
change.

It carries one more thing, because it cannot be separated: `w13_weight` and
`w2_weight` are real parameters of the FusedMoE, so a trainer that mirrors
ATOM's state dict rather than the checkpoint's resolves in
`_get_param_to_module_mapping` and never reaches the routed-expert sync behind
`_apply_unmatched_weight`. Down the plain dispatch it took the row-major
`copy_` into a buffer the kernel reads through aiter's 16x16 expert
permutation, with `_pending_expert_relayout` left empty so
`_finalize_expert_weight_sync` returned at `if not pending` -- and `updated`
counting it as a success. `_check_expert_sync_supported` never ran on that
route either, so a quantized or expert-parallel MoE was not refused. Named
expert buffers are recognised in the dispatch now and handled like any other
expert write: refused if unsupported or partial, written through the fence, and
every slice registered for the relayout. The three dispatch blocks are where
both changes land, line for line, which is why they are one commit.

The three `except Exception as e:` around the loader fallback pick up a
`# noqa: BLE001` here. They are untouched pre-existing lines, but CI runs ruff
through reviewdog with `-filter-mode=diff_context`, so a finding within three
lines of a change is reported and `-fail-on-error=true` fails the job. The
catch-all is the intent -- a loader is model code and can raise anything, and
one tensor failing must not abandon the rest of the sync -- so it is annotated
with a reason, in the form the ten existing sites in `atom/` use, rather than
narrowed to placate the rule.

Co-authored-by: Cursor <cursoragent@cursor.com>
`_post_process_fp8_weight` re-ran `normalize_e4m3fn_to_e4m3fnuz` on every
weight update. `need_normalize_e4m3fn_to_e4m3fnuz` is a static property of the
layer -- `params_dtype == torch.float8_e4m3fnuz`, set once in
`create_weights` -- not a to-do list, and nothing clears it after the
load-time conversion, so the conversion re-ran on an already-converted
parameter. It is not idempotent in either buffer:

  * `weight_scale` is rebuilt as `scale * 2.0`, so it doubles again every
    sync. Measured: 3.0 -> 6.0 -> 12.0. After N syncs the dequantized weight
    is 2**N too large. The multiply also returns a freshly allocated tensor,
    which moves the scale's address out from under a captured decode graph --
    the same hazard the `shuffle_weights` fix in this series removed, on the
    buffer it did not cover.
  * the weight needs no conversion at all here. `_requantize_fp8_weight`
    quantizes into `param.dtype` against `finfo(e4m3fnuz).max`, and the
    direct-copy path is handed bytes already in `param.dtype`, so both arrive
    in the target convention.

Gated on `param.dtype == torch.float8_e4m3fn` so it converts what is
genuinely unconverted and nothing else, and the scale is written in place when
it does. The weight's own rebind stays: `normalize_e4m3fn_to_e4m3fnuz` fixes
its bytes through an int8 view of the same storage and hands back that storage
with a reinterpreted dtype, so `data_ptr()` does not move and no captured
graph can see it.

The function's own fence moves with it, out of the top of the function and
onto the two writes it makes -- the conversion's byte fixup and the layout
shuffle -- so a call that decides to write nothing no longer pays for one.

Gated on the e4m3fnuz per-token path, which is the normal ROCm FP8
configuration; inert on gfx950 with the smoke's config, where the flag is
False, so the unit tests are what cover it.

Co-authored-by: Cursor <cursoragent@cursor.com>
… way

`_finalize_expert_weight_sync` cleared `_pending_expert_relayout` after its
loop, so any raise inside the loop skipped the clear and left behind the
entries for buffers it had already shuffled. The next successful sync shuffled
those a second time, and per this function's own docstring that does not undo
the first shuffle, it produces a third layout. Silent, with `updated=N` logged
as success.

Three ways in: the half-rewritten-expert `RuntimeError` raised for a later
buffer in the same loop, a `NotImplementedError` out of
`_check_expert_sync_supported` mid-sync, and a bucketed SHM/IPC sync that
aborts before `is_last`.

Each entry is now dropped as its own shuffle completes, inside a `finally`.
Not a blanket clear: an entry the loop never reached describes a buffer that is
row-major right now, and a later sync re-establishing its layout is the only
thing that fixes it, so that one has to survive. Anything still pending when
the loop unwinds is logged at error, naming the buffers, rather than left to be
inferred from a later wrong answer.

The relayout's own write gets the fence the rest of the path gained, since it
is the second in-place write this sync makes to those slices and the one a
graph is most likely to catch half done -- a half-permuted expert reads as
plausible garbage rather than as an error.

Co-authored-by: Cursor <cursoragent@cursor.com>
Two loose ends around the graph release this series added to `_release_kv_cache`.

`ModelRunner.exit()` clears `runner.graphs` and, under TBO,
`model.tbo_graphs`; `release_cudagraphs` cleared only the first. The
replayable handle is in `runner.graphs` for a TBO graph too, so dropping that
is enough to stop the replay -- `tbo_graphs` is a parallel store holding the
graph, its per-ubatch contexts and the output tensor it captured. Left behind,
that entry pins the graph's private memory pool: exactly the footprint the
caller went to sleep to reclaim, held until a recapture happens to overwrite
the same key. It is cleared in the same place now, so the two release paths
agree, and the early return no longer skips it when `runner.graphs` is empty.

The second is that `sleep_keeps_memory_resident` defaults to False, so a
default-configured rollout under `PYTORCH_CUDA_ALLOC_CONF=expandable_segments`
now takes a recapture on every level-1 sleep/wake where before it took none.
That is the intended trade -- the alternative is replaying graphs against a
pool that has been freed and reallocated, which is wrong unconditionally --
but the operator only finds out when the recapture faults, from a handler that
by then has set `enforce_eager = True` permanently. And since
`sleep_keeps_memory_resident()` reads `enforce_eager` first, from that point
the option stops taking effect at all: a self-consistent state, since there
are no graphs left to keep valid, but not the one it was set for.

So the release site says it, where it can still be acted on: releasing graphs
with expandable segments configured and the option off warns and names the
option. And the failure handler says which of the two states it has left
behind, because from there the option's own log lines never appear again.

Co-authored-by: Cursor <cursoragent@cursor.com>
`weight_is_stored_preshuffled` unified the load-time and sync-time shuffle
decisions, and took `==` from the load side. `QuantType` is a pybind enum out
of the compiled `aiter.jit.module_aiter_core`, so a module whose `quant_type`
came from a differently-identified aiter import -- a plugin process, a
re-import, `atom.quant_spec`'s lazy proxy -- does not compare equal to these
members under `==`. Every branch then falls through to `return False`: the
sync writes the FP8 weight and never re-shuffles it, and the preshuffle GEMM
reads a row-major weight. Silent, and only after the first weight update.

`.value` is what this file already uses for every comparison that runs after
the load, what `weight_updater` used throughout before this function existed,
and what `attention_residual.py` documents the reason for. Unifying the two
call sites on `==` took the sync side backwards.

Behaviour-neutral where the enum has one module identity, which is the case in
the validation image: the function's full truth table is identical before and
after, and the load/sync agreement tests still pass.

Co-authored-by: Cursor <cursoragent@cursor.com>
… dtype

`shuffle_weights` writes the shuffled 2D weight through the existing storage
so an online update does not move an address a captured decode graph holds.
That needs a `copy_` kernel for the weight's dtype on its device, where the
rebind it replaced needed nothing: MXFP4's `Float4_e2m1fn_x2` has no CPU copy
kernel before torch 2.10, and `linear.py`'s online-quant path shuffles exactly
that dtype -- reachable from CPU weight prep, a CPU unit test, or a
`param_device="meta"` stream.

It falls back to the rebind on `NotImplementedError`. Nothing is lost by
doing so: a weight being shuffled on a device with no copy kernel for it is
not one a captured decode graph is replaying against, so there is no address
to preserve. The in-place write stays the path everything on device takes.

Co-authored-by: Cursor <cursoragent@cursor.com>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Graph recapture misses PIECEWISE and speculative-drafter stores, and FP8 synchronization has unresolved padding, TP, and aborted-transfer cases.

Get a fresh assessment by requesting another Copilot review.

Review details

Suppressed comments (4)

atom/rollout/memory_manager.py:77

  • [verified] PIECEWISE captures are stored in each CUDAGraphWrapper.concrete_cudagraph_entries (and their pools), not in runner.graphs. A PIECEWISE runner can therefore take this early return with stale wrapper graphs still holding the old weight/KV addresses, and _graphs_backup_keys is never recorded, so wake skips recapture and the next replay can use freed storage. Author must include the PIECEWISE wrapper state in release/recapture bookkeeping, and add a regression test for a runner with no ordinary graphs.
    if not getattr(runner, "graphs", None):
        return

atom/rollout/weight_updater.py:828

  • [verified] _maybe_pad_a8w8_preshuffle_output() can make this parameter wider than the logical trainer tensor (for example, N=4097 is stored as N=4224). _requantize_fp8_weight returns without writing on that shape mismatch, but this branch still increments updated, so an online FP8 sync reports success while that layer continues serving its old weight. Author must handle the known logical-versus-padded shape, including the scale rows, before requantization or propagate the failure instead of counting it as updated.
            elif self._is_fp8_param(module, param) and tensor.dtype != param.dtype:
                self._requantize_fp8_weight(module, param_name, param, tensor)
                updated += 1

atom/rollout/weight_updater.py:832

  • The update API documents full, unsharded tensors, but this same-dtype FP8 branch copies directly into the rank-local parameter and bypasses _try_shard_weight. On TP>1, a full FP8 tensor therefore has the wrong shape (or cannot be laid out correctly) and the sync fails instead of updating the rank's shard. Route this case through the same TP-sharding path as the non-FP8 fallback, then run _post_process_fp8_weight after the successful in-place write.
            elif self._is_fp8_param(module, param) and tensor.dtype == param.dtype:
                tensor = tensor.to(device=self.device)
                self._copy_into_param(param, tensor)
                self._post_process_fp8_weight(module, param)

atom/rollout/weight_updater.py:218

  • The pending relayout state survives an aborted non-final SHM/IPC bucket. If a bucket writes one expert shard and then raises, the next transfer reuses that stale shard and can combine it with the next update's other shard before shuffling, producing a weight assembled from two trainer states. Treat each bucket sequence as a transaction by clearing partial expert/paked-update state on abort, or associate it with an explicit transfer id.
        if not hasattr(self, "_expert_relayout_pending"):
            self._expert_relayout_pending = {}
        return self._expert_relayout_pending
  • Files reviewed: 21/21 changed files
  • Comments generated: 1
  • Review effort level: Lite

Comment thread atom/rollout/memory_manager.py Outdated
@valarLip

Copy link
Copy Markdown
Collaborator

Second pass — reviewed at 948e863d5

Thanks for turning the earlier round around so fast. Three of the new commits answer findings from my previous comment, so this pass leads with those: each of the three fixes has a defect, and one of them misses the default configuration entirely. After that, what is still open from last time, then what is new.

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 full chain.


1. The three fixes

1.1 f9a9eb661 — the graph release does not cover the default compilation level [verified]

release_cudagraphs returns early when runner.graphs is empty:

# atom/rollout/memory_manager.py:76
if not getattr(runner, "graphs", None):
    return

Under PIECEWISE capture, runner.graphs is always empty:

# atom/model_engine/model_runner.py:3880-3906
if _piecewise:
    ...
    self._piecewise_captured_tokens.add(num_tokens)
    ...
    continue                      # never reaches self.graphs[(bs, max_q_len)] = graph  (:3944)

And nothing else releases the per-piece graphs — grep -rn _piecewise_captured_tokens atom/ returns hits only inside model_runner.py, and grep -rn -i piecewise atom/rollout/ returns nothing at all.

PIECEWISE is --level 3, the default. So on a default-configured rollout, sleep(level=1) frees and reallocates the KV pool while the per-piece graphs still hold the old base — the Memory access fault by GPU node-2 this PR reproduces, unchanged. _graphs_backup_keys is never set, so _recapture_cudagraphs_if_needed also returns at its own not _graphs_backup_keys check.

Two more things sit after that early return and are therefore skipped on the same path: the graph_logits clear (the leak it was added for survives) and _warn_if_recapture_will_fault, so the operator gets no warning either.

The gate needs to key on "was anything captured", not on runner.graphs specifically — _piecewise_captured_tokens is the piecewise-side equivalent and needs its own release.

1.2 ec671eabe — the sweep missed the one function that must agree with the new one [verified]

weight_is_stored_preshuffled compares by .value, and its docstring states exactly why:

linear.py:454-460 — "QuantType is a pybind enum out of the compiled aiter.jit.module_aiter_core, so a module whose quant_type came from a differently-identified aiter import — a plugin process, a re-import, atom.quant_spec's lazy proxy — does not compare equal to these members under ==. Every branch would then fall through to return False."

But _maybe_pad_a8w8_preshuffle_output still compares by identity:

# linear.py:935
if not (self.quant_type == QuantType.per_Token and self.params_dtype == dtypes.fp8):
    return False

These two have to agree: the first decides whether the weight is stored preshuffled, the second decides whether to pad N so the preshuffle GEMM can consume it. Under precisely the condition the new docstring names, the first returns True (values match) and the second returns False (identity does not) — the weight is shuffled but never padded, and the RuntimeError at :946, whose stated purpose is to "fail loudly here rather than let shuffle_weights hit its cryptic x.shape[-1] % 32 == 0 assertion", is skipped along with it.

1.3 948e863d5 — the in-place write introduces a new aliasing hazard [mechanism verified, chain reported]

shuffle_weights now writes through the existing storage:

# atom/model_ops/utils.py:163-165
if shuffled.shape == weight.shape and shuffled.dtype == weight.dtype:
    weight.copy_(shuffled)

That is the right call for CUDA-graph address stability, but it changes the aliasing contract for everything holding a view of the parameter, and there are at least two such holders:

  • deepseek_v4.py:2511 caches self._wo_a_w_fp8 = w.data.view(G, N, K), read later at :2796. Before this PR the rebind left that cache on the old storage; now it aliases the live parameter, so a subsequent shuffle of the same parameter changes the bytes the cached BMM weight reads. ATOM_FP8_BLOCKSCALE_WEIGHT_PRESHUFFLE defaults to "1" (envs.py:506), so this is the default V4 path, and the reported consequence is silent wrong logits rather than an error. [reported] for the loader-ordering half of the chain.
  • linear.py:1242 — row_view() builds a Parameter over self.weight.data.narrow(0, start, length). The comment two lines above says the view gets "its own parameter dict so rebinding weight/scale cannot disturb self" — it is written against the rebind semantics this commit removed.

This is a fix-then-sweep: every consumer that retains a view or a .data alias of a shuffled parameter needs checking, not just these two.


2. Still open from the previous round

  • preempt()'s strip width. Sharper statement than I gave last time: preempt at :2449 strips seq.num_placeholder_tokens, while postprocess's own definition of the non-real tail at :2815 is num_tokens - num_placeholder_width - num_rejected. The two disagree by mtp_k on every spec-decode verify step, so mtp_k + num_rejected stale slots survive into the recomputed context. No new test drives a step where fwd_output actually carries rows for the sequence; all four exercise the no-verify deferred-prefill path, where the two formulas coincide.
  • _await_readers_of coverage. Four writes still bypass both wrappers: weight_scale.data.copy_ at :678, :690, :752 and weight_scale.data.fill_ at :684.
  • _pending_expert_relayout is not cleared on the failure path, so one half-delivered expert poisons every later sync for the life of the process.

3. New, ranked by blast radius [reported]

  • _apply_fused_expert_weight's TP contract is inverted (weight_updater.py:325). The 3D path never sets load_full, so _load_w13/_load_w2 re-slice the source by tp_rank and a rank-local tensor is only partially written — no shape error, updated counts success, and _finalize_expert_weight_sync bakes the result into the kernel layout. TP=1 hides it completely, which is the dangerous part.
  • add_request switched to preprocess_fanout (llm_engine.py:254), dropping the SamplingParams.n > 1 guard that exists "so that callers which expect exactly one Sequence back cannot silently drop the other siblings." generate() still returns a flat list, so every caller that zips prompts with outputs pairs prompt i with a sibling of prompt i//n. A clean error became a silent mispairing.
  • _cached_expert_mapping / _param_to_module are hasattr-cached and never invalidated (weight_updater.py:187), while self.model is rebound to UBatchWrapper under TBO and to torch.compile(...) at level 1. The name prefixes then never match, every sync is counted skipped at debug level — the silent no-op this PR exists to fix, reintroduced under two supported configs, and unrecoverable because the caches cannot be rebuilt.
  • if seq.return_logprobs and len(seq.logprobs) >= strip: (scheduler.py:2468) skips the logprob deletion on a shortfall while the token arrays are truncated unconditionally. Sequence.append_token never touches logprobs, and the new _schedule_first_decode_after_remote_kv appends drafts through it, so the shortfall is reachable and the desync is permanent. A min() plus a warning is the right shape, not a silent skip.

Also worth a look: _EXPERT_BUFFER_SHARDS lists only w13_weight/w2_weight while FusedMoE registers w13_bias and three scale params on the same modules; the vocab-tail mask is applied only in RLHFModelRunner.postprocess, so prefill_forward's sampler, compute_argmax_token and the drafters all bypass it; use_model_sensitive_rmsnorm is threaded into 2 of ~7 RMSNorm dispatch paths and the removal of the old try/except TypeError makes it an unguarded hard dependency on a recent aiter with no version floor recorded.


4. CI

Four of the new test files — test_shuffle_weights_storage.py, test_rollout_expert_weight_sync.py, test_rollout_vocab_mask.py, test_weight_sync_shuffle_layout.py — call pytest.skip(..., allow_module_level=True) without CUDA. CI has no GPU, so the only test of the central data_ptr()-stability fix never runs there. test_weight_sync_inplace_ordering.py uses per-test skipif and keeps its CPU-checkable assertions live; that is the shape the other four want.


Of the above, 1.1 is the one I would not merge without: the PR would land labelled as fixing the fault while the fault remains on the default --level 3.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Unresolved critical and moderate findings affect vocabulary masking, graph cleanup, TP expert synchronization, scheduler logprob alignment, and fan-out sampling.

Get a fresh assessment by requesting another Copilot review.

Review details

Suppressed comments (2)

atom/model_engine/llm_engine.py:274

  • [verified] This enables n > 1 for disaggregated configurations too, but ModelRunner.prefill_forward receives _needs_independent_noise from prepare_model and drops it when calling the sampler (atom/model_engine/model_runner.py:4572-4577). The first token for fan-out siblings can therefore use the shared-noise path, producing identical first tokens and reducing the requested diversity. Thread the flag through the prefill sampler and add a regression test for that path.
            fanout = self.io_processor.preprocess_fanout(
                prompt,
                sampling_param,
                stream_callback=callback,
                multimodal_data=mm_data,

atom/rollout/weight_updater.py:414

  • The fused route passes each full intermediate-dimension chunk into weight_loader as a 3D loaded_weight. On TP ranks, FusedMoE.weight_loader treats 3D inputs as full_load, and _load_w13/_load_w2 do not slice the source by tp_rank; the destination is only a local shard. Thus a trainer tensor in the advertised (E, 2I, H) / (E, H, I) global layout will fail with a shape or out-of-range copy when tp_size > 1, unlike the per-expert 2D route. Slice each chunk to the rank-local TP range before this call, or reject global fused tensors explicitly; the current path silently does not support the fused format under TP.
        gpu = tensor.to(device=self.device)
        # Split w13's gate and up halves along the intermediate dim, the way
        # the buffer stacks them. w2 arrives whole. Views, not copies: the
        # loader's copy handles a strided source, and materialising these
        # would double the largest tensor in the sync.
        for shard_id, chunk in zip(shard_ids, gpu.chunk(len(shard_ids), dim=1)):
  • Files reviewed: 27/27 changed files
  • Comments generated: 4
  • Review effort level: Lite

Comment thread atom/rollout/memory_manager.py
Comment thread atom/rollout/model_runner_ext.py
Comment thread atom/model_engine/scheduler.py
Comment thread atom/rollout/memory_manager.py
xysheng-AMD and others added 7 commits September 16, 2026 05:24
`release_cudagraphs` returned early when `runner.graphs` was empty, and under
PIECEWISE -- `--level 3`, the default -- that dict stays empty however much was
captured: the capture loop moves on before the assignment, each compiled dense
piece self-capturing into its own `CUDAGraphWrapper`. So the release was a
no-op, and the fault it exists to prevent -- a graph replaying the base of a KV
pool that `sleep(level=1)` has since freed and reallocated -- was untouched on
the configuration that walks into it. For DeepSeek-V4 those pieces hold the KV
scatter itself (the narrow split leaves it inside them, `_attn_pre`), so it is
the fault and not only the leak.

Two things this PR had already added sat behind the same early return and were
skipped with it: the `graph_logits` clear and `_warn_if_recapture_will_fault`.

Every store is now asked separately and the answers OR-ed, because "was
anything captured" is not a question `runner.graphs` can answer:

* the per-piece wrappers, reached through a registry (`graph_holders.py`) --
  the compile backend installs them with `module.__dict__[target] = ...` on a
  submodule of a split graph module only Dynamo holds, which defeats a walk
  from either end;
* `_piecewise_captured_tokens`, whose clearing is what stops the next step
  dispatching PIECEWISE and either replaying a dropped graph or recording a
  replacement mid-serve, uncoordinated, into the first collective;
* the drafter's per-batch recordings, walked because they are reachable. A
  draft pass writes the KV it attends, so these hold the pool's base the way a
  decode graph does, and `ATOM_DRAFT_CUDAGRAPH` is on by default;
* the graph pool handles, since starting a capture on a pool whose last graph
  has just gone away trips an allocator assertion rather than making a new one.

Recapture on wake keys on a flag set by whatever was released, not on
`_graphs_backup_keys` -- that list belongs to the manual store and is empty
here, so keying on it would have left the same configuration recapturing
nothing once the release was fixed.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>
… sync

`_finalize_expert_weight_sync` raised on an expert whose shards had not all
arrived, and the entry stayed in `_pending_expert_relayout` -- so it raised
again on the next sync, and the one after that, taking every other buffer's
relayout down with it each time. One malformed update disabled expert weight
sync for the life of the process.

Nothing completes such an entry: a sync sends every shard of an expert or the
write is refused outright, so keeping it bought nothing. The bytes are
recoverable without it -- the next update carrying the whole expert overwrites
both halves row-major, and that entry relays out normally, which is now what
the error tells the operator to do.

Split per expert rather than per buffer while here: an expert missing a shard
cannot have its layout re-established, and that is no reason for the buffer's
complete experts to go on being read row-major through the permutation. The
raise stays, once, after the cleanup: those slices are half new and half old, in
two layouts, and the kernel reads them as plausible garbage rather than failing.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>
…pers

`ModelRunner` rebinds `self.model` twice after the model is built: to a
`UBatchWrapper` under TBO, and to `torch.compile(...)` at compilation level 1.
Both hold the real model as a CHILD, so `named_modules()` on either prefixes
every parameter with the wrapper's own attribute name -- `model.` or
`_orig_mod.` -- and the mapping a weight sync resolves against then matches
nothing the trainer sends. Every weight is counted `skipped` at debug level and
the rollout goes on serving the weights it had: the silent no-op this path
exists to remove, under two supported configurations.

The four lookups now go through one unwrapped model, peeled by TYPE rather than
by attribute name -- nearly every HF-derived model has a submodule literally
called `model`, and peeling that would drop a prefix the trainer does send.

Each cache is keyed on the object it was built from instead of on `hasattr`,
which nothing can invalidate and which this class has no hook to invalidate
from. (`get_expert_mapping` and `packed_modules_mapping` were reachable through
both wrappers, which forward plain attribute lookups; asking the unwrapped
model is so that all four lookups describe one module rather than depending on
that forwarding.)

And a sync that matched NOTHING now says so above debug level. It stays
reachable however many name conventions are covered -- a wrapper was one, the
next trainer is another -- and `updated=0, skipped=N` on an info line reads
exactly like a bucket that legitimately held nothing.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>
…ecision

`weight_is_stored_preshuffled` compares `quant_type` by `.value` and its
docstring says why: the enum is a pybind type out of the compiled
`aiter.jit.module_aiter_core`, so a module whose `quant_type` arrived through a
differently-identified aiter import does not compare equal to these members
under `==`.

`process_weights_after_loading` asks that helper whether to shuffle and then
asks ITSELF whether to pad N -- by identity. Under exactly the condition the
docstring names, the first answers True and the second False: the weight is
shuffled and its N left unpadded, in a layout the preshuffle GEMM cannot
consume, and the RuntimeError whose stated purpose is to catch that is skipped
along with the padding. The sweep that fixed the first function missed the one
function that has to agree with it.

Every `quant_type` test in that method now compares by value, not just the two
that decide a layout: the scale shuffle at the tail and the per_Tensor
requantization are the same comparison in the same method, and a multi-partition
per_Tensor weight that skips requantization keeps one scale per partition where
the kernel reads a single one.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>
`preempt` strips `seq.num_placeholder_tokens`, and on a spec-decode verify step
that width named only part of what was there. The in-place overwrite is given a
window `mtp_k` wider than the run appended for it -- `required_placeholders =
num_placeholder + offset` -- and the step hands back `mtp_k - num_rejected + 1`
tokens for it, so `mtp_k + num_rejected` slots of the window are still
`eos_token_id` when the next run is appended behind them.

Measured, deferred output with drafts every step: the trailing `eos` run settles
at `2 * mtp_k + 1` whatever the acceptance rate, while the recorded width was
`mtp_k + 1 - num_rejected`. A preemption there left `mtp_k + num_rejected`
placeholders in the recomputed context -- which is the fault the width exists to
prevent, since a surviving placeholder stops being one: the recompute prefills a
context ending in `<|endoftext|>` and the model starts a new document.

So postprocess records what the overwrite did not consume, and the run appended
after it adds to that rather than replacing it. One writer, one reader, no
formula for `preempt` to re-derive. `=` also undercounted a sequence that went
two steps without a row: nothing consumed the first run, and the second appends
beyond it.

The tests that came with the width only drove the no-verify deferred-prefill
path, where the two formulas coincide, so none of them could see this. The new
ones carry rows.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>
`add_request` fans out, so a prompt with `SamplingParams.n > 1` becomes `n`
sibling sequences and `generate` hands back `n` outputs per prompt in ONE flat
list. Nothing said so. Before this PR that call reached `preprocess`, which
returns a single sequence and refuses n > 1, so offline n > 1 raised before a
token was generated -- and the guard replaced by the fan-out was the only thing
telling a caller about the shape of what it gets.

The ordering is prompt-major, because sequence ids are assigned in fan-out order
and `generate` sorts on them: a caller pairs prompts with outputs by expanding
its own list by `n`. That is what `Lumen-RL`'s ATOM server does -- it groups
equal prompts, sets `n` to the group size, and zips against the same expansion
-- so it is a contract, and now a tested one rather than an accident of the
sort. `preprocess` keeps its guard for callers that expect exactly one sequence.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>
… place

`shuffle_weights` writes through the existing storage rather than rebinding it,
which is what keeps a captured graph's address valid and also changes what a
retained view of the parameter sees. The sweep for consumers that keep one found
two, and both are comments that now describe the wrong semantics rather than code
that needs changing:

* `deepseek_v4`'s `_wo_a_w_fp8` is a view of the live parameter, taken AFTER the
  preshuffle above it -- built before, it would have named the same bytes in the
  wrong layout. Nothing shuffles that weight again: `quant_type` is set to `No`
  on the line below, which is also what a weight sync reads.
* `ColumnParallelLinear.make_row_view`'s comment claimed the view's own parameter
  dict isolates it from a rebind. It isolates the parent from the VIEW's
  bindings, never the bytes, and its only caller rebuilds the view on every call
  regardless -- the `is not` guard there compares two `param.data` objects, and
  that attribute hands back a fresh one on each access.

Also narrow what the vocab-tail mask claims: it covers the decode sampling path,
which is every token a colocated rollout generates, and not
`prefill_forward`'s own sampler or the TP-sharded `compute_argmax_token`, whose
masking needs the shard offset and belongs with that reduction.

Reported-by: valarLip
Co-authored-by: Cursor <cursoragent@cursor.com>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Unresolved critical FP8/MoE synchronization defects and moderate graph and vocabulary-handling issues remain.

Get a fresh assessment by requesting another Copilot review.

Review details

Suppressed comments (6)

atom/models/deepseek_v4.py:2516

  • This file contains the @support_torch_compile model boundary (DeepseekV4Model at atom/models/deepseek_v4.py:4298), and the repository rule is not to modify such model files. Keep this explanatory text outside the compile-boundary file (or revert this comment-only hunk) so future model-file changes do not bypass that boundary.
            # A VIEW of the live parameter, and `shuffle_weights` now writes
            # through the storage rather than rebinding it, so this tracks the
            # bytes rather than pinning the ones it was built from. That is why
            # it is taken AFTER the shuffle above: built before, it would have
            # named the same bytes in the wrong layout. Nothing shuffles this

atom/rollout/memory_manager.py:458

  • The failure handler only resets the manual self.graphs store. capture_cudagraph() can already have populated graph_logits, TBO entries, draft recordings, or registered PIECEWISE holders before a later capture fails; those invalid graph references then survive the eager fallback and continue pinning memory. Clear every graph store/holder in this path before returning to eager mode.
        except Exception:
            logger.exception(f"{self.label}: CUDA graph recapture failed")
            # Fall back to eager mode rather than crashing
            self.enforce_eager = True
            self.graphs = {}

atom/rollout/memory_manager.py:186

  • This substring check treats PYTORCH_CUDA_ALLOC_CONF=expandable_segments:False as enabled and emits a warning recommending resident memory even though the allocator has explicitly disabled expandable segments. Parse the option's boolean value instead of checking only for the option name.
    if "expandable_segments" not in os.environ.get("PYTORCH_CUDA_ALLOC_CONF", ""):

atom/rollout/model_runner_ext.py:108

  • This override does not cover the disaggregated prefill path: ModelRunner.prefill_forward samples directly at model_runner.py:4577 and never calls this postprocess. A RapidServe prefill worker can therefore still return a padded vocabulary ID as its first token before decode postprocessing, so the configured mask is not enforced for every rollout token. Apply the same mask before that sampler (with the appropriate sharded-vocabulary handling) or explicitly reject this configuration for disaggregated prefill.
        This covers the decode sampling path, which is every token a colocated
        rollout generates. Two other places reach a sampler without coming
        through here, and are NOT covered:

        * ``ModelRunner.prefill_forward`` samples the first token itself, for

atom/rollout/weight_updater.py:940

  • Expert scale/metadata names that resolve through get_expert_mapping() still enter the ordinary direct-parameter dispatch here. For example, ...experts.0.gate_proj.weight_scale resolves to w13_weight_scale, which misses the exact-name check above: a shape-matching tensor is copied into the scale, while a mismatched one falls into the loader/catch path instead of raising. This can pair a newly synced scale with old expert bytes and violates the unsupported expert-scale contract; reject all w13_*/w2_* expert metadata before the ordinary branches in each transport.
            if param_name in _EXPERT_BUFFER_SHARDS:
                self._apply_named_expert_buffer(name, param_name, module, param, tensor)
                updated += 1
            elif self._is_fp8_param(module, param) and tensor.dtype != param.dtype:

atom/rollout/weight_updater.py:109

  • skipped also counts known parameters with shape mismatches and loader failures in the update loops. Consequently a known parameter can produce updated=0, skipped=1 and this warning incorrectly says it resolved to no module, which sends operators to the wrong diagnosis. Track unmatched names separately (or pass that count) before emitting this message.
        if updated == 0 and skipped > 0:
            logger.warning(
                f"{self.label}: weight update matched NOTHING -- {skipped} "
                f"parameter(s) resolved to no module, so nothing was written and "
                f"the rollout is still serving the weights it had. Compare the "
  • Files reviewed: 27/27 changed files
  • Comments generated: 2
  • Review effort level: Lite

Comment thread atom/rollout/weight_updater.py
Comment thread atom/rollout/weight_updater.py

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Unresolved correctness issues remain in configuration compatibility, weight synchronization, graph cleanup, vocabulary masking, and scheduler bookkeeping.

Get a fresh assessment by requesting another Copilot review.

Review details

Suppressed comments (6)

atom/model_engine/scheduler.py:2796

  • [verified] The new placeholder count includes the missing_placeholders repaired by the loop immediately above, but that loop calls append_token without adding the corresponding 0.0 logprob. When return_logprobs is enabled, preempt now assumes every counted placeholder has a logprob and can delete a real token's entry, leaving token_ids/output_tokens and logprobs out of sync. Author must add the zero logprob for each repaired placeholder before recording the count.
                seq.num_placeholder_tokens = max(
                    0, required_placeholders - len(token_ids)
                )

atom/rollout/model_runner_ext.py:70

  • [verified] A positive true_vocab_size is accepted when hf_config.vocab_size is missing or zero, even though the postprocess guard later only masks when logits.shape[-1] > true_vocab_size. On a supported config with no top-level vocab_size, setting a value larger than the actual embedding/logit width therefore silently masks nothing—the invalid-value case this option is meant to reject. Author must validate against the actual embedding/logit width or reject positive overrides when that width cannot be determined.
        # Not every PretrainedConfig subclass carries vocab_size at the top
        # level; when it is missing there is nothing to check against.
        padded = getattr(config.hf_config, "vocab_size", 0)
        if padded and self._true_vocab_size > padded:

atom/rollout/model_runner_ext.py:117

  • [verified] true_vocab_size is only applied in this override's postprocess, but the inherited disaggregated ModelRunner.prefill_forward samples logits directly via self.sampler (atom/model_engine/model_runner.py:4551-4583) and returns that token to decode. With a padded checkpoint and true_vocab_size=151665, the first token can still be a tail id, so the rollout can emit the exact undecodable id this option is meant to prevent. Author must share the mask with prefill_forward (or reject this configuration for disaggregated prefill) before sampling.
        if self._true_vocab_size > 0 and logits.shape[-1] > self._true_vocab_size:
            logits[..., self._true_vocab_size :] = float("-inf")

atom/rollout/weight_updater.py:425

  • [verified] The fused 3D route passes rank-local chunks to weight_loader, but FusedMoE.weight_loader's 3D load_full flag is dropped by _copy_expert_shard before _load_w13/_load_w2; those loaders therefore apply the normal tp_rank source narrow for TP>1. For example, with TP=2 a local (E, 2I_local, H) gate/up tensor is narrowed again to half-width and the copy fails or loads the wrong slice. Author must propagate the full-load mode through the MoE loader (or change the input contract) and add a TP>1 fused-expert test.
        for shard_id, chunk in zip(shard_ids, gpu.chunk(len(shard_ids), dim=1)):
            self._load_into_param(
                param,
                weight_loader,
                chunk,
                # _copy_expert_shard dispatches on the name containing
                # "weight"; the fused leaf names do not, so hand it the
                # resolved ATOM name.
                weight_name=atom_name,
                shard_id=shard_id,
                expert_id=0,
            )

atom/rollout/weight_updater.py:891

  • [verified] The initial-load path applies fp4_utils.e8m0_shuffle() to weight_scale for per_1x32 weights (atom/model_ops/linear.py:940-943), but this sync post-process only reshuffles param. For a same-dtype MXFP8 update, the weight is moved to the kernel layout while its new scale remains row-major, so the GEMM reads mismatched scales after the first sync. Author must apply the corresponding scale transform with the same synchronization or reject this update mode.
            self._await_readers_of(param)
            shuffle_weights(param)

atom/rollout/weight_updater.py:804

  • [verified] This per_Token branch still computes q with direct division, whereas the initial/online load path uses AITER's get_hip_quant(QuantType.per_Token). The repository's ROCm quantizer documents that division can be one ULP off AITER and that all-zero rows must retain a zero scale, so an online sync can produce different FP8 codes/scale from a freshly loaded equivalent weight even after the layout fix. Author must reuse the shared quantizer or reproduce its reciprocal and zero-row semantics, then add a bitwise regression test.
        elif quant_type is not None and quant_type.value == _QT.per_Token.value:
            row_amax = tensor_gpu.abs().amax(dim=-1, keepdim=True)
            scale = (row_amax / fp8_max).clamp(min=1e-12)
            self._copy_into_param(param, (tensor_gpu / scale).to(fp8_dtype))
            weight_scale.data.copy_(scale.to(weight_scale.dtype))
  • Files reviewed: 27/27 changed files
  • Comments generated: 3
  • Review effort level: Lite

Comment thread atom/config.py
Comment thread atom/rollout/weight_updater.py
Comment thread atom/rollout/memory_manager.py
@xysheng-AMD

Copy link
Copy Markdown
Contributor

@valarLip Seven commits on top of 948e863d5, one per finding, so the delta is
948e863d5..HEAD.

Fixed

1.1 The graph release misses the default compilation level.
release_cudagraphs returned early when runner.graphs was empty. Under
PIECEWISE (--level 3, the default) that dict is always empty: the capture loop
continues before the assignment and each graph lives in its own
CUDAGraphWrapper. On DeepSeek-V4 this is the fault rather than a leak — the
narrow piecewise split leaves the KV scatter inside a graphed dense piece
(deepseek_v4.py:1611), so those graphs hold the base of the pool.

Every store is now asked separately and the answers OR-ed: runner.graphs +
graph_logits and UBatchWrapper.tbo_graphs as before, plus the per-piece
CUDAGraphWrappers, _piecewise_captured_tokens, DraftGraph._cuda_graphs and
the graph pool handles.

  • The wrappers are reached through a registry, not a walk: the backend installs
    one with module.__dict__[target] = ... on a submodule of a split
    fx.GraphModule, so it is outside _modules and there is no path from the
    runner. The new module imports weakref and nothing else, so the release path
    stays importable without aiter.
  • Lazy recapture in the wrapper is not a substitute: it only happens when
    PIECEWISE is dispatched, and a recapture mid-serve is uncoordinated and hangs
    on the first collective. So the dispatch record has to be cleared and
    capture_cudagraph() has to run on wake.
  • The wake gate keyed on _graphs_backup_keys, which records the manual store
    and is empty under PIECEWISE. It now keys on whether anything was released.
  • The drafter's recordings are released too: a draft pass writes the KV it
    attends, so they hold the pool's base, ATOM_DRAFT_CUDAGRAPH is on by
    default, and nothing released them before — not this PR and not exit().

Cost: a PIECEWISE runner now recaptures on wake, as a FULL one already did. The
graph_logits clear and the recapture warning, both sitting after the same early
return, are fixed with it.

1.2 The .value sweep missed the function that has to agree with it.
process_weights_after_loading asks the helper whether to shuffle, then asks
itself by identity whether to pad N. Under the condition the helper's docstring
names, the first answers True and the second False: the weight is shuffled, N is
left unpadded, and the RuntimeError meant to catch that is skipped. Every
quant_type test in that method now compares by value, including the scale
shuffle at :928 and the per_Tensor requantization at :860 — skipping the
latter leaves one scale per partition where the kernel reads a single one.

1.4 A half-delivered expert failed every later sync. The try/finally
discards only what it had already relaid out; the entry that raises is not in
that list and therefore survived. The rule I wrote last round to keep unreached
entries alive kept this one alive too. A missing shard can never be completed, so
the old code raised on it at every finalize. Now: complete-but-unshuffled
survives, half-delivered is dropped with the raise kept once after the cleanup,
and the split is per expert rather than per buffer. The bytes are recoverable
without the entry — the next update carrying the whole expert overwrites both
halves row-major — and the error says so.

1.5 The caches are never invalidated. The name-prefix half holds and is worse
than reported: under TBO named_modules() prefixes every parameter with
model., so no name the trainer sends matches, every weight is counted skipped
at debug level, and the whole update is a silent no-op. Level 1's torch.compile
does the same with _orig_mod.. All four lookups now go through one unwrapped
model, peeled by type (most HF-derived models have a submodule named model, and
peeling by name would drop a prefix the trainer does send), and each cache is
keyed on the object it was built from instead of hasattr. A sync that matches
nothing now warns.

§2 preempt()'s strip width. The trailing eos run settles at
2 * mtp_k + 1 (measured 3/5/7 for mtp_k 1/2/3), independent of the acceptance
rate: the overwrite window is mtp_k wider than the run appended for it, and the
step hands back mtp_k - num_rejected + 1 tokens. The recorded width was
mtp_k + 1 - r, so mtp_k + r slots survived. The offset you pointed at is
correct — it covers the previous step's unconsumed window; the recorded width was
wrong in counting only the appended part. postprocess now records what the
overwrite did not consume and the appended run adds to it. One more case this
fixes: a sequence that goes two steps without a row had its first run overwritten
in the count.

§3 add_request fan-out. The change is this PR's, and offline n > 1 raised
before it. The fan-out is load-bearing for this branch's consumer: Lumen-RL's
ATOM server groups equal prompts, sets n to the group size, calls generate,
and zips the flat result against the same prompt-major expansion. Restoring the
guard would break G rollouts per prompt. add_request and generate now state
the fan-out and the prompt-major ordering, with two tests pinning it;
preprocess keeps its guard. Making generate return a list of lists is a larger
break and needs the owner.

Disagreements and open items

load_full's contract runs the other way. moe.py:4136 sets
full_load = len(loaded_weight.shape) == 3, so 3D sets it, and _load_w13 with
load_full true narrows only the destination (:3696-3701), skipping the
tp_rank slicing of the source at :3710-3712. The source is not re-sliced and
the caller must send rank-local. A full tensor gives load_size > expert_shard_size, where expert_data.narrow goes out of bounds and raises
rather than writing part of the slice. If you are reading a different path into
_load_w13, point at it.

The get_expert_mapping half does not hold. UBatchWrapper defines
__getattr__ forwarding to self.model (ubatch_wrapper.py:603, in main
since #515) and OptimizedModule forwards too, so the mapping was always
reachable. Asking the unwrapped model is so that all four lookups describe one
module, not a repair.

1.3 aliasing: the mechanism holds, these two sites do not. Comments only.
deepseek_v4.py:2511 takes the cache after the preshuffle three lines above
it, so it names the bytes batched_gemm_a8w8_mxscale_bpreshuffle wants; a second
shuffle does not exist because the next line sets quant_type to No.
linear.py:1242's comment claims the view's own _parameters isolates it from a
rebind — it isolates the parent from the view's bindings and never the bytes, and
it does not matter: _local_q_proj is the only caller and rebuilds the view on
every call, because its is not guard compares two param.data objects and that
attribute returns a fresh one each access (p.data is p.data is False).

The four weight_scale writes: unchanged. They do bypass both wrappers, but
each immediately follows a fenced weight write (:678←:677, :684←:683,
:690←:689, :752←:743), and _await_readers_of is a device-level
torch.cuda.synchronize: on return everything previously issued on that device,
including any reader of the scale, has completed. The residual window exists for
the weight write too and is inherent to the design. 0/7 corrupt on the 8-step
smoke, twice. If you want the wrapper there regardless it costs nothing.

scheduler.py:2468 logprob shortfall: not addressed this round. The desync
is real. It is reachable only with return_logprobs; say so and I will add the
min() and the warning.

§4 CI: the conclusion holds, the remedy does not.
tests/model_ops/test_shuffle_weights_storage.py does not run in CI. In all four
files the aiter-backed imports sit at module scope below the skip, so per-test
skipif leaves them executing at collection, where the CPU runner raises
cannot import name 'logger' from 'aiter' (unknown location) — CI resolves
aiter to a namespace package. Moving the imports into the tests does not help
either: every assertion in those four needs shuffle_weight, a QuantType
member, a real FusedMoE or ModelRunner. The way to run them is a GPU CI job,
or making aiter's dtypes lazy upstream. The new code here stays importable
without aiter, so 30 of the 38 tests added this round run on the CPU runner.

Two findings, neither this PR's, neither fixed

V4 wo_a's fp8 preshuffle is defeated by a weight sync. wo_a sets
quant_type to No after its own preshuffle to stop LinearBase shuffling
again, and a weight sync reads the same field: it writes row-major FP8 into the
parameter in place and never re-shuffles, while _wo_a_w_fp8 and
batched_gemm_a8w8_mxscale_bpreshuffle keep reading it as 16x16-shuffled. Silent
wrong logits after the first sync, on gfx950 with
ATOM_FP8_BLOCKSCALE_WEIGHT_PRESHUFFLE=1 (the default). Not caused by the
in-place shuffle_weights: the sync's param.data.copy_ has always been in
place. A fix needs a module to declare "stored preshuffled" independently of
quant_type, which is a design call, and I have no V4 FP8 checkpoint here.

CUDAGraphWrapper's input-address contract makes PIECEWISE unusable in the RL
deployment.
The per-piece capture completes and then all eight rollout replicas
abort on the first replay with Input addresses for cudagraphs are different during replay. Pre-existing and untouched by this PR; it is why Lumen-RL pins
cudagraph_mode=FULL, and its release/versions.env records the same abort.
It is also why 1.1 went unnoticed: the configuration that fills the piecewise
store cannot complete a step there.

Data

Unit tests and static checks. 617 → 655 in the scheduler and weight-sync
groups; each of the seven commits green on its own with the count rising
monotonically (634, 636, 642, 646, 653, 655, 655 — the last is comment-only). 25
of the 38 new tests fail on the tree they were written for; the rest pin the
"release must not fire" side. black clean over 778 files; ruff's
diff_context gate 0 blocking findings against the PR's merge base;
Check Pre Checkin Signal green on the pushed head.

1.1 on device. The RL deployment cannot reach 1.1 (Lumen-RL pins
cudagraph_mode=FULL and sleep_keeps_memory_resident=true), so I drove ATOM's
engine directly: level=3, PIECEWISE, releasing sleep, two sleep/wake cycles
after init, nothing decoded so the replay assertion never enters. Same probe,
only the ATOM revision differs. On d6b9e147c: 0 graphs released, device free
203.93 → 203.93GB, and the wake dies with available_for_kv=-1435.69MB
(peak_torch=77.53GB still counts those graphs). On this branch: 1332 graphs
released (36 buckets x 37 pieces), 203.93 → 207.80GB, and both cycles recapture
and continue. The piecewise pool is 1.34GB for this model; cuda_graph.py
records V4 at 1.62GB after the granularity split and 8.37GB before it.

A releasing sleep under FULL was already correct before this PR. The MoE
example with sleep_keeps_memory_resident=false completes and passes on both
revisions, 24/24 release and recapture each, no negative pool. Sampling the eight
cards, each sleep window takes the total from ~991GB to 196-569GB. What this PR
adds is the stores a FULL capture does not fill.

RL examples on this branch. Both ATOM examples pass the launcher's --check:
MoE at k3_kl 0.00132 against a 0.00138 reference, 8B FP8 at 0.00316 against
0.00287. The MoE run reports expert layout re-established for 12288 expert slices across 96 fused buffers 24 times, identical to the same example on the
pinned revision, with every weight-sync bucket at skipped=0 (2352 and 336 of
them) and zero occurrences of the three new warnings. A third run, 8B with FULL
and a releasing sleep, logs the release and the recapture on 16/16 replica-steps
at k3_kl 0.00107258 against 0.00107182 for the same example with the pool kept
resident.

Accuracy (Kimi-K2.7-Code-MXFP4) is not this PR's: it went red at
44054b3d3 on the rebase alone, from #2170's flydsl gather_kv_b_proj backend
meeting a model that quantizes kv_b_proj with ptpc_fp8. #2221 disabled it
per-model and #2224 reverted that; neither is in this tree.

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.

7 participants