[Lumen-RL] Improve FP8 rollout weight synchronization and CUDA Graph stability - #2028
Conversation
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
There was a problem hiding this comment.
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_eagerbranch of_release_weights(), the conditionif 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.
There was a problem hiding this comment.
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 checksnot 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
returnin 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_blockscaleis 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 meanslogger.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)
)
There was a problem hiding this comment.
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. Sinceresume_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_handlernow defaults to WARNING viaATOM_LOG_LEVEL, but the logger itself is still set to INFO. That means INFO log calls will still constructLogRecords and be filtered at the handler, which defeats the performance intent described in the comment above. Set the logger level to the sameATOM_LOG_LEVELyou apply to the console handler.
console_handler.setLevel(
getattr(logging, _envs.ATOM_LOG_LEVEL, logging.WARNING)
)
There was a problem hiding this comment.
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: ... returnguard, control flow implies enforce_eager is True, soif not self.enforce_eager and ...is always false and CUDA graphs will never be released in the eager sleep path. Drop the redundantnot self.enforce_eagercheck 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
There was a problem hiding this comment.
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_LEVELcurrently only changes the console handler level; theloggeritself is still hard-coded toINFO. That meanslogger.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 singlelevelfromATOM_LOG_LEVELand 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 whennot self.enforce_eagermakes the subsequent CUDA-graph release block unreachable, and in the eager path (self.enforce_eager is True) the current conditionif not self.enforce_eager and ...will never run. Ifself.graphsis populated whileenforce_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: addinguse_model_sensitive_rmsnormwill 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
)
f0ca2b7 to
bf820bd
Compare
There was a problem hiding this comment.
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. Ifself.graphsis 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")
There was a problem hiding this comment.
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-decoratedmodel (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)
)
bf820bd to
2b5294c
Compare
|
Reviewed at 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: The three headline fixes are each undone one or two functions away, in the same fileThat 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_tokens1. [verified]
|
`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>
There was a problem hiding this comment.
🟡 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_tokensis 0._prefill_chunk_for_budgetreturnsNonefor 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 inmodel.tbo_graphswhen 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. Clearmodel.tbo_graphsalong 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
|
@valarLip The central observation is correct and holds for all three headline fixes: a Seven items fixed, three answered. Two of the answered ones carry a mechanism Fixed1 — 2 — Four more call sites would not fix it, since a forgotten call is the failure 4 — offload-resume sizes with 6 — 7 — Behaviour-neutral where the enum has one module identity, as in the validation 8 — an expert tensor under ATOM's own fused name bypasses the relayout. Ships inside the fence commit: in all three dispatch blocks the new branch and 9 — The blanket The second suggestion — treating a non-empty Answered, mechanism corrected3 — Two defects there are worse than an address move. The weight needs no conversion on this path: Inert where the FP8 smoke runs — the flag is False on gfx950 — so the 5 — The failure mode is not use-after-free. Under TBO the replayable handle still 10 — the default sleep now takes a recapture. Agreed on both halves, The second half is now addressed where it can be acted on rather than where it Test coverageThe strongest item in the review. The two scheduler commits had no tests: Added, all runnable without a GPU, with the count that fails unfixed:
The ordering file was rewritten because its old tests asserted that a wait On the module-level skips: The "no collection-abort risk" note is correct. Worth recording how it can Verification
|
There was a problem hiding this comment.
🔵 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: eachCUDAGraphWrapperkeeps itsconcrete_cudagraph_entries, and the piecewise capture path can therefore leave this dictionary empty. This guard returns without setting_graphs_backup_keysor 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_logitsseparately fromrunner.graphs(seeModelRunner.capture_cudagraph). Clearing onlyrunner.graphsleaves these captured output buffers referenced during sleep, soempty_cache()cannot reclaim their graph-private memory and wake can recapture on top of it. Cleargraph_logitswhenever the graph stores are invalidated, including whenrunner.graphsis 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 setfull_load=True, and_load_w13()/_load_w2()only narrow the local destination; they do not slice the full source byself.tp_rank. For example, TP=2 with a full(E, 2I, H)gate_up_projproduces(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 forw1/w3and dim 2 forw2before 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_weightsdocumentsnamed_tensorsas full, unsharded weights, and_requantize_fp8_weightexplicitly narrows such tensors forworld_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_paramand 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_weightroute: 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 raisesNotImplementedErrorfor 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_supportedand 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_relayoutempty, and returns without a layout step. The implementation now registers every expert in_pending_expert_relayoutat 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, butFusedMoE._load_w13/_load_w2explicitly skip the TP slicing branch whenload_full=Trueand 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_GEMMafteratom.model_ops.linearhas already been imported, butlinear.pyimportsgemm_a8w8_tritononly inside the import-timeuse_triton_gemm()branch. When the module was collected with the default disabled,gemm_a8w8_tritonremainsNone, so thetriton_gemmparameterization 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-NoneTriton 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
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>
There was a problem hiding this comment.
🟡 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 inrunner.graphs. A PIECEWISE runner can therefore take this early return with stale wrapper graphs still holding the old weight/KV addresses, and_graphs_backup_keysis 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_weightreturns without writing on that shape mismatch, but this branch still incrementsupdated, 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_weightafter 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
Second pass — reviewed at
|
There was a problem hiding this comment.
🟡 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 > 1for disaggregated configurations too, butModelRunner.prefill_forwardreceives_needs_independent_noisefromprepare_modeland 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_loaderas a 3Dloaded_weight. On TP ranks,FusedMoE.weight_loadertreats 3D inputs asfull_load, and_load_w13/_load_w2do not slice the source bytp_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 whentp_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
`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>
There was a problem hiding this comment.
🟡 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_compilemodel boundary (DeepseekV4Modelatatom/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.graphsstore.capture_cudagraph()can already have populatedgraph_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:Falseas 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_forwardsamples directly atmodel_runner.py:4577and never calls thispostprocess. 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_scaleresolves tow13_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 allw13_*/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
skippedalso counts known parameters with shape mismatches and loader failures in the update loops. Consequently a known parameter can produceupdated=0, skipped=1and 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
There was a problem hiding this comment.
🟡 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_placeholdersrepaired by the loop immediately above, but that loop callsappend_tokenwithout adding the corresponding0.0logprob. Whenreturn_logprobsis enabled,preemptnow assumes every counted placeholder has a logprob and can delete a real token's entry, leavingtoken_ids/output_tokensandlogprobsout 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_sizeis accepted whenhf_config.vocab_sizeis missing or zero, even though the postprocess guard later only masks whenlogits.shape[-1] > true_vocab_size. On a supported config with no top-levelvocab_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_sizeis only applied in this override'spostprocess, but the inherited disaggregatedModelRunner.prefill_forwardsamples logits directly viaself.sampler(atom/model_engine/model_runner.py:4551-4583) and returns that token to decode. With a padded checkpoint andtrue_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 withprefill_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, butFusedMoE.weight_loader's 3Dload_fullflag is dropped by_copy_expert_shardbefore_load_w13/_load_w2; those loaders therefore apply the normaltp_ranksource 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()toweight_scaleforper_1x32weights (atom/model_ops/linear.py:940-943), but this sync post-process only reshufflesparam. 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_Tokenbranch still computesqwith direct division, whereas the initial/online load path uses AITER'sget_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
|
@valarLip Seven commits on top of Fixed1.1 The graph release misses the default compilation level. Every store is now asked separately and the answers OR-ed:
Cost: a PIECEWISE runner now recaptures on wake, as a FULL one already did. The 1.2 The 1.4 A half-delivered expert failed every later sync. The 1.5 The caches are never invalidated. The name-prefix half holds and is worse §2 §3 Disagreements and open items
The 1.3 aliasing: the mechanism holds, these two sites do not. Comments only. The four
§4 CI: the conclusion holds, the remedy does not. Two findings, neither this PR's, neither fixedV4
DataUnit tests and static checks. 617 → 655 in the scheduler and weight-sync 1.1 on device. The RL deployment cannot reach 1.1 ( A releasing sleep under FULL was already correct before this PR. The MoE RL examples on this branch. Both ATOM examples pass the launcher's
|
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
mainat 82f0b24. Verified on Qwen3-8B-Base and Qwen3-30B-A3Bon 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_weightof the layer'sFusedMoE, which is neither theincoming name nor anything
packed_modules_mappingdescribes. They matchednothing, 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_loaderwrites plainrow-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
Parameterobjects while a captured graph and_param_to_modulestill point at the old ones, and several are not idempotent(
Fp8MoEMethod's per-tensor path collapsesw13_weight_scalefrom[E, 2]to[E]on its first call, so a second raisesIndexError).shuffle_expert_slices()is the layout step alone, applied in place to onlythe 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 ornumerically 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_projand(E, H, I)experts.down_proj. Same buffers, same dim order, w13's first halfalong 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_loaderonce per half: a 3Dloaded_weightputs the loader onits 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_loadingfor the initial load and inWeightUpdaterMixin._post_process_fp8_weightfor an online update. The twodisagreed wherever the rule was not simply the env var: the module-level
needs_preshuffled_weightexception that DeepSeek's fusedqkv_a_projsets,the triton a8w8
per_TokenGEMM that wants the unshuffled weight, thenon-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 syncside 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_weightsreboundtensor.datato aiter's return value for a 2Dweight, 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.
self.weightandself.weight_scalewith freshParameters, which carry none of the attributes__init__hung on theoriginals.
weight_loader()readsweight_loader_processoff the parameterit 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 andnothing else. The decode graphs captured the base address of that pool, but
only
_release_weightscleared them and recorded_graphs_backup_keys, so alevel-1 sleep left them in place and
_recapture_cudagraphs_if_neededreturnedearly on wake with nothing to recapture.
_resume_kv_cachethen allocated apool 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:
The graph release moves into
release_cudagraphs()and is called from bothrelease 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;
mainhas the same shape.Graph recapture on wake faults under
expandable_segmentsSleep 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, soConfig.sleep_keeps_memory_residentoffers 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 definesneither
enforce_eagernor the config field gets the behaviour every callerhad 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 andcan 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 environmentvariable named after a downstream project. It is a property of the checkpoint,
not of the deployment, so it belongs in
Config, and a downstream-prefixedname has no place in ATOM.
Config.true_vocab_sizedefaults 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.pyguardedfrom aiter.ops.shuffle import moe_shuffle_scalewith anImportErrorfallback toshuffle_scale. The two are not aliases:moe_shuffle_scaledispatches on the chip and callsshuffle_scale_n32k4ongfx1250, so the fallback would have quietly laid out MoE scales for the wrong
kernel there. AITER has exported
moe_shuffle_scalesince #3756 (2026-06-20)and
mainimports it directly, so that is restored.layernorm.pycalled AITER's RMSNorm withuse_model_sensitive_rmsnorm=1inside a
try/except TypeErrorthat parsed the exception message and cachedthe verdict in two module-level globals. That parameter has been part of
rmsnorm2d_fwdandrmsnorm2d_fwd_with_addsince #647 (2025-07-17), and bothare
@torch_compile_guardcustom ops where probing from inside the tracedregion 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_RMSNORMunconditionally leaves the default pathexactly as
mainhas it, and there is nothing for a branch to choose between.Verified bit-identical to
main's call on MI355X.add_requestfan-out had no testLLMEngine.add_requestroutes throughpreprocess_fanoutso thatSamplingParams.n > 1produces n sequences instead of one. Nothing assertedthat, 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
preemptonly stripped the trailing placeholder token when speculativedecoding was on. Without speculation it left one behind, and the placeholder is
eos_token_id.The placeholder is not a speculation-only artifact:
postprocessappends onewhenever
need_placeholderholds, which includesis_deferred_out— definedas
pipeline_parallel_size == 1(model_runner.py:189), i.e. true for everyTP-only engine. The scheme is safe because the next step's
postprocessoverwrites the placeholder in place. A sequence preempted on this step never
reaches that next step, so the overwrite never happens:
<|endoftext|>, so itstarts a fresh document instead of continuing the answer;
ignore_eos=Falsethe 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, thecounter
postprocesssets where it appends andpreemptclears just below.The matching
0.0entries the placeholders pushed ontoseq.logprobsgo withthem —
postprocesspatches those in place too, so removing tokens withoutthem would leave the two lists describing different positions.
A chunked recompute was declared final at the prompt boundary
is_final_chunkcompared progress againstseq.num_prompt_tokens. Correct fora first admission, wrong for a sequence re-admitted after
preempt, which mustrecompute 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 thenext_token_idsloop directly below, whichcarries a comment saying exactly that.
So a 337-token recompute split into 256 + 81 was declared complete after the
first chunk, because
256 >= 128held against the prompt length. The remaining81 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.
postprocessre-derived the same predicate asseq.num_cached_tokens < seq.num_prompt_tokens. It now reads theis_final_chunkthe scheduler already froze at pre-advance offsets, ratherthan re-deriving it: by the time
postprocessruns,num_tokensmay havegrown 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_seqssequencesat 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=8192fits 9 perbatch.
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_weightsused torebind
tensor.datato aiter's return value, which moves the address out fromunder 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
Test Result
Python compilation passed. Part 1 + FP8 ordering:
239 passed. Scheduler:301 passed. Black clean across the repo; Ruff reports nothing on the linesthis 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 twosleep/wake cycles, because the point of
sleep_keeps_memory_residentis thatnothing moves — something a test on
updated/releasedcounters 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:Corruption is
mismatch/k3_kl— the trainer's KL between the logprobs therollout 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 forceddirectly so the only variable is pool size.
max_num_batched_tokensAttribution, from the scheduler's own preemption and re-prefill logs:
In the
800 / 8192 / neitherrun, 45 of 64 requests were preempted and 38broke. Every broken request was one of the preempted ones; not a single
request that was never preempted broke.
The
first fix onlyrow isolates the second bug. All 32 broken requeststhere were chunked recomputes, and 9 continued another request's count with
no
<|endoftext|>anywhere near the break:1200000, 1200001, 1200002, ...— req 2's range.2200003, 2200004, 2200005, ...— req 12's.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
SchedulerandBlockManagerwith a fake forward, checkingthat 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:
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_kl0.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:
gpu_memory_utilization=0.30, no scheduler fixesPreemptions per replica per step, from
get_metrics_statistics(): 2265-2596 at0.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
mainhas since grown them independently: the 128x128 block-scale FP8 online
quantization in
atom/quantization/quark/utils.pyandatom/model_ops/linear.py, andATOM_LOG_LEVEL. There is nothing to reviewfor 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
mainat0d4e96605are squashed into the first commit. They were written against a much older
mainand several only fix up the commit before them — including several thatlater commits in this series revert outright — so replaying them individually
onto today's
mainproduces intermediate trees that never existed andconflicts against code
mainhas since evolved on its own. The resulting treeis 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