nemo_rl.algorithms.single_controller#
SingleController: asyncio orchestrator for the RL training loop.
CPU-only Ray actor that runs two concurrent pumps plus a watchdog, and coordinates the other actors via lightweight RPCs. SC sends control signals and reads metadata only — model tensors still move through DataPlane or NCCL.
Data flow: _rollout_pump → gen.generate_and_push(prompt, dp_client) ← RPC to GenWorker GenWorker → dp_client.put_samples(…) _train_pump → sampler.evict/select against TQReplayBuffer → _value_stage(meta) (PPO only) → value.get_values_from_meta(…) Value → dp_client.get/put_samples(…) (via its own client) → _advantage_stage(meta) → dp_client.get_samples(…) → adv_estimator.compute_advantage(…) → dp_client.put_samples(…) → _value_train_epochs(meta) (PPO only) → value.train_from_meta(…) Value → dp_client.get_samples(…) (via its own client) → trainer.begin/train_microbatches/finish_train_step (split API, driver-side TQPolicy via asyncio.to_thread) Trainer → dp_client.get_samples(…) (via its own client) → dp_client.clear_samples(…) ← SC clears after train _sync_weights → WeightSynchronizer.sync_weights()
Module Contents#
Classes#
Outcome returned by one rollout checkpoint save attempt. |
|
Controller sidecars captured with one native TQ snapshot. |
|
CPU-only Ray actor that orchestrates the RL training loop. |
Functions#
Return the latest value keyed by generation data-parallel worker. |
|
Return a deterministic nearest-rank percentile for telemetry. |
|
Compute whole-step OPD metrics from exact pooled sufficient statistics. |
|
Return only the data-plane columns produced for this train step. |
Data#
API#
- nemo_rl.algorithms.single_controller.Generation#
None
- nemo_rl.algorithms.single_controller.log#
‘getLogger(…)’
- nemo_rl.algorithms.single_controller._SUPERVISOR_DRAIN_TIMEOUT_S#
30.0
- class nemo_rl.algorithms.single_controller._RolloutCheckpointSaveResult#
Outcome returned by one rollout checkpoint save attempt.
- saved: bool#
None
- class nemo_rl.algorithms.single_controller._RolloutCheckpointCut#
Controller sidecars captured with one native TQ snapshot.
- dataloader_state: dict[str, Any]#
None
- sampler_dispatch_index: int#
None
- replacement_reserve: list[nemo_rl.data.interfaces.DatumSpec]#
None
- replay_metadata: Optional[nemo_rl.algorithms.async_utils.replay_buffer.TQReplayMetadataState]#
None
- rollout_recovery_payload: Optional[bytes]#
None
- rollout_recovery_group_count: Optional[int]#
None
- replay_row_count: int#
None
- staging_row_count: int#
None
- rolled_back_train_group_count: int#
None
- mutation_version: int#
None
- tq_save_seconds: float#
None
- nemo_rl.algorithms.single_controller._latest_generation_values(
- metrics: dict[str, Any],
- metric_name: str,
Return the latest value keyed by generation data-parallel worker.
- nemo_rl.algorithms.single_controller._percentile(values: list[float], quantile: float) float#
Return a deterministic nearest-rank percentile for telemetry.
- nemo_rl.algorithms.single_controller._pooled_opd_metrics(
- stat_sum: float,
- stat_sumsq: float,
- count: int,
- gap_sum: float,
Compute whole-step OPD metrics from exact pooled sufficient statistics.
stat_*pool the advantage andgap_sumthe raw teacher-student gap over the same tokens. They differ once TROPD (proximal_teacher_alpha < 1) or subtract_global_baseline reshapes the advantage.
- nemo_rl.algorithms.single_controller._train_fields_for_step(
- *,
- policy_logprobs_required: bool,
- reference_logprobs_required: bool,
Return only the data-plane columns produced for this train step.
- class nemo_rl.algorithms.single_controller.SingleControllerActor(
- master_config: nemo_rl.algorithms.single_controller_utils.config.MasterConfig,
- actor_args: nemo_rl.algorithms.single_controller_utils.setup.SingleControllerActorArgs,
- setup_timing_metrics: nemo_rl.algorithms.metric_utils.SetupTimingMetrics,
CPU-only Ray actor that orchestrates the RL training loop.
Owns three concurrent asyncio tasks:
_rollout_pump: dispatches prompts to GenerationWorkerActor
_train_pump: claims DataPlane meta, trains, clears consumed rows, then runs _sync_weights (drain gate + weight synchronization) inline after each optimizer step
_stall_watchdog_pump: publishes rollout counters and reports stalls or unhealthy environments, which are the failures that otherwise produce no signal at all
Plus _gen_fleet_probe_pump when fleet health is enabled, which probes generation shard liveness on its own, much shorter clock.
All other actors are passive — they expose methods and wait to be called.
Initialization
Initialize the SingleController actor.
- Parameters:
master_config – SC MasterConfig.
actor_args – Pre-built actor args from setup_single_controller.
setup_timing_metrics – Driver-side setup timings; logged here (Logger isn’t cloudpickleable).
- _recovering_from_refit: bool#
False
- _tracer: Any#
None
- async run() dict[str, Any]#
Main entry point. Runs until max_train_steps is reached.
- async _run_pumps() dict[str, Any]#
Start the rollout / train / watchdog pumps and run until one finishes.
- async ping() dict[str, Any]#
Liveness check — returns immediately if event loop is running.
- _log_telemetry_metrics(
- metrics: dict[str, float],
- *,
- step: int,
- prefix: str,
Log benchmark telemetry on an axis independent of trainer steps.
- _record_rollout_timing(
- *,
- work_started: float,
- dispatch_started: float,
Record one committed group’s queue and execution durations.
- _log_rollout_restore_metrics(
- *,
- replay_metadata_load_seconds: float,
- recovery_prepare_seconds: float,
- restored_replay_groups: int,
Log rollout restore phases after the controller is ready to dispatch.
- async _maybe_restore_replay_buffer() int#
Restore the local replay index for the native TQ checkpoint.
Recovery is authoritative only for samplers that explicitly support buffered-group restoration. The native snapshot and replay metadata file must both be present and agree on their manifest and group count.
- async _maybe_restore_rollout_recovery(
- *,
- restored_replay_groups: int,
Restore unfinished ownership for prioritized rollout-pump redispatch.
- async _queue_restored_replay_groups_for_regeneration( ) None#
Convert discarded canonical replay groups into fresh prompt work.
- async _load_recovery_prompt(
- *,
- group_id: str,
- sample_id: str,
Resolve and collate one stable dataset prompt reference.
- _validate_restored_sampler_cursor() None#
Require the sampler cursor to cover every restored target step.
- _require_unoccupied_target_step(
- target_step: Optional[int],
Reject a sampler admission that collides with existing buffer ownership.
- async _rehydrate_rollout_recovery_prompts( ) None#
Resolve positional prompt references against a stable map-style dataset.
This assumes the dataset exposes integer
__getitem__and retains the same ordering across checkpoint and restart.
- async _admit_reserved_prompt_groups(
- group_ids: list[str],
Commit one admission and mark every reserved group as admitted.
- Returns:
The target-step stamp assigned to every group.
- async _redispatch_restored_rollouts(
- launch: Callable[[nemo_rl.data.interfaces.DatumSpec, Optional[int], str], Awaitable[None]],
Prioritize durable unfinished groups while the train pump drains TQ.
Launching happens inside the ordinary rollout pump so restored groups use the same in-flight and replay-capacity semaphores as new work. The train pump runs concurrently and releases replay capacity as it consumes canonical or newly recovered groups; therefore recovery cannot deadlock merely because the checkpoint contained more unfinished ownership records than free replay slots.
- async _validate_replay_inventory(
- replay_metadata: nemo_rl.algorithms.async_utils.replay_buffer.TQReplayMetadataState,
Require the canonical TQ keys to match the SC replay index exactly.
Live checkpoint callers must hold the exclusive data-plane barrier so commits and clears cannot race this inventory read. Restore calls are also safe before the rollout and train pumps start any live writers.
- async _validate_rollout_recovery_inventory(
- cut: nemo_rl.algorithms.async_utils.replay_buffer.DataPlaneMutationCut,
- *,
- replay_metadata: Optional[nemo_rl.algorithms.async_utils.replay_buffer.TQReplayMetadataState],
- clear_unreferenced: bool,
Validate staging ownership while the caller holds a stable cut.
- async _maybe_restore_replacement_reserve() None#
Restore spare prompts diverted before the previous run’s checkpoint.
These were pulled from the dataloader and never dispatched, so the restored dataloader resumes past them. Nothing else in the checkpoint holds them, and without this they are simply gone: one batch of the dataset per divert, plus the training step
_clamp_max_num_stepshad budgeted for it.No sampler-name guard, unlike the buffer restore. Spares carry no stamp – they are prompts that never reached
admit– so nothing about them depends on which sampler wrote the checkpoint. They are restored even into a run that has since switched toon_dropped_prompt="shrink", where the pool is never drawn on but is still drained back into training at the end of the dataloader.
- async _ray_get(obj_ref: Any) Any#
Await a Ray ObjectRef without blocking the asyncio event loop.
- async _call_dp(method_name: str, **kwargs) Any#
Call a DataPlaneClient method or a Ray actor exposing that method.
- async _save_data_plane_checkpoint(
- checkpoint_path: nemo_rl.utils.checkpoint.PathLike,
- *,
- train_steps: int,
- trainer_version: int,
- current_epoch: int,
- replay_metadata: Optional[nemo_rl.algorithms.async_utils.replay_buffer.TQReplayMetadataState] = None,
- rollout_recovery_payload_sha256: Optional[str] = None,
- rollout_recovery_group_count: Optional[int] = None,
Save a required TQ snapshot inside an SC checkpoint bundle.
A sampler with replay-buffer recovery writes an authoritative native TQ snapshot bound to its metadata-only replay index by a digest. Other samplers retain shadow-mode snapshots until their recovery contract is defined. Failures propagate so a finalized bundle never silently omits the advertised data-plane component.
- static _request_staging_keys( ) list[str]#
Return the full receipt-manifest staging ownership for a request.
- async _cleanup_known_finalization_request_unlocked(
- cut: nemo_rl.algorithms.async_utils.replay_buffer.DataPlaneMutationCut,
- request: nemo_rl.experience.rollout_reassembler_actor.ReassemblyRequest,
Clear known request ownership while holding a barrier mutation slot.
- async _cleanup_known_finalization_request( ) None#
Clear a known request outcome without racing a native TQ snapshot.
- async _finalize_with_actor( ) Optional[nemo_rl.experience.rollout_reassembler.FinalizedGroup]#
Finalize and index one group atomically with respect to TQ saves.
Returns the committed FinalizedGroup once the group is committed to the replay buffer (callers may read valid_row_count/total_row_count off it to decide whether the group is worth keeping), or None when the finalizer itself dropped it as a structural outcome (ownership already cleaned up; the caller credits the step short). A low valid-row fraction is no longer a finalizer-side drop – the caller decides that, since only the caller can source a replacement.
- async _cleanup_consumed_metas_unlocked(
- cut: nemo_rl.algorithms.async_utils.replay_buffer.DataPlaneMutationCut,
- metas: list[nemo_rl.data_plane.KVBatchMeta],
Clear consumed ownership while holding a barrier mutation slot.
- _log_data_plane_metrics(total_step_time: float) None#
Log this step’s data-plane cost. Never raises.
On by default, so this runs every step of every recipe. Mirrors
grpo_sync._log_data_plane_metrics.
- _log_data_plane_metrics_impl(total_step_time: float) None#
Log this step’s data-plane cost. No-op unless observability is enabled.
The synchronous loop logs these series from
_log_data_plane_metricsingrpo_sync. Without the same call here the single-controller path builds the metrics client, pays for its counters on every op, and emits nothing – the failure is silent, because an empty dashboard looks the same as a data plane that cost nothing.Driver scope only, and the prefix says so. This client issues the advantage stage’s get, the put that writes the advantages back, and the post-train clear; the bulk traffic is the trainer and generation workers’ own clients, in their own processes with their own counters, so
comm_volume_mbhere is well under what the job actually moved.grpo_syncgets a cluster view by fanning out over its policy worker group; this loop has no such group to fan out over, so driver scope is all there is here.
- static _group_ids_from_meta(
- meta: nemo_rl.data_plane.KVBatchMeta,
Return stable prompt-group IDs in canonical sample order.
Reads the same tag the advantage stage keys its baseline on rather than parsing the
_g{i}suffix back off the sample ids, so the two do not disagree about what a group is. The stage would raise on a batch whose tag is missing anyway; doing it here fails one stage earlier in the same iteration.
- async _rollout_pump() None#
Continuously dispatch rollout tasks until cancellation.
Per batch: 0. Under on_dropped_prompt=”replace”, divert the batch into the spare pool if the pool is below its low-water mark, and skip admission entirely. Otherwise await sampler.admit(…) to wait until the batch may dispatch and obtain its target_step stamp.
Per prompt:
Acquire _buffer_capacity slot (backpressure)
Acquire sem (cap concurrent in-flight rollouts)
Wait for _rollout_permitted (paused during weight sync)
Call rollout_manager.generate_and_push(prompt) — local async RolloutManager reserves a slot, runs the rollout, then commits the group via TQReplayBuffer (→ dp_client.put_samples + mark ready). In token-capture mode (finalizer actors present) the rollout is run via generate_for_finalization instead and its metadata-only request is submitted to the finalizer actor pool.
If the prompt was dropped, substitute a spare and repeat step 4 – for this step, or for whichever step lends this one a finished group in its place (see _take_replacement, _promote_into_step) – or credit the step short so the train pump can close it
Decrement _inflight_rollouts
Once every epoch is done, whatever is left in the spare pool is dispatched as ordinary steps rather than discarded (see _drain_reserve_into_steps).
- _divert_batch_to_reserve(
- prompt_batch: nemo_rl.distributed.batched_data_dict.BatchedDataDict[nemo_rl.data.interfaces.DatumSpec],
Consume a whole batch as spare prompts instead of admitting it as a step.
Returns whether the batch was taken, in which case the caller must not admit it. Diverting before
admitis what keeps the stamp sequence honest: admitting a batch and then dispatching nothing for it would leave a target step that no group is ever generated for, which is exactly the hang the shortfall accounting exists to prevent.A whole batch at a time because the dataloader only yields batches. The spares that go unused are not wasted work – nothing has been generated for them – and they stay in the pool for later steps.
Nothing is diverted until the sampler has actually stamped a batch, so a run whose sampler never stamps does not lose a batch of prompts to a pool it can never draw on. The cost is that the first batch is always admitted rather than diverted; in practice the pool is filled while that first batch’s rollouts are still running, so it is available by the time any of them can be given up on.
- async _drain_reserve_into_steps(
- launch: Callable[[nemo_rl.data.interfaces.DatumSpec, Optional[int], Optional[str]], Awaitable[None]],
Train on the leftover spares once the dataloader has nothing more to give.
Spares were consumed from the dataset like any other prompt, so leaving them in the pool at the end of the last epoch throws away data the run already paid for.
It also restores the step count.
_clamp_max_num_stepsderivesmax_num_stepsfromlen(dataloader), and every diverted batch is one fewer batch the loop can admit – so without this a replace-mode run quietly finishes one step short of the budget it was configured with, per divert.Whole steps only. A partial pool dispatched as a step is short by construction, and
min_step_batch_fractionwould then reject it and fail a run that had otherwise completed cleanly. In the ordinary case the pool holds exactly one batch (the dataloader usesbatch_size=num_prompts_per_step), so the common outcome is that the whole thing is recovered.Not gated on
on_dropped_prompt: an empty pool makes this a no-op anyway, and only “replace” ever fills one, so the gate would buy nothing while stranding a pool restored from a checkpoint into a run that has since switched to “shrink”.
- _take_replacement(
- target_step: Optional[int],
- replacements_used: int,
A spare prompt to stand in for a dropped group, or None to shrink instead.
None covers the four ways a replacement can be unavailable: it was not asked for, the sampler did not stamp this prompt so no step is waiting on it, the per-slot budget is spent, or the pool is empty because the dataloader is exhausted. Every one of them falls back to
on_dropped_prompt="shrink"rather than waiting, because a step whose replacements keep failing still has to close.
- _promote_into_step(
- target_step: Optional[int],
Fill a dropped step from a later step’s finished work, and name the lender.
Where a replacement goes, rather than whether one happens. The lost step closes on generation that already exists instead of waiting out a rollout with the trainer idle, and the caller redirects its spare prompt to the lender, which is due a training step later and has the slack to absorb the wait. The same prompt is generated either way.
Only ever reached with a spare already in hand, which is what makes the borrow safe to take: an unrepaid loan is the same hole one step later.
Returns None – leaving the caller filling the dropped step directly – when nothing is stamped so no step is stranded, when the trainer has already moved past this step (a second drop can land after the first one closed it short, and a group stamped for a finished step would only be evicted), or when no later step has a finished group to lend. The last is always the case at
in_order.max_lookahead_versions=0, where the next batch is not dispatched until this step trains.- Returns:
The step that lent the group, which the caller now owes a rollout, or None.
- _credit_shortfall(target_step: Optional[int]) None#
Record that a stamped step will never receive a group it is waiting for.
- _target_groups_for_step(step: int) int#
How many prompt groups this step should train on, after dropped prompts.
num_prompts_per_stepis the target; groups stamped for this step that were given up on are subtracted, because they are never arriving and a sampler that matches batches to steps exactly cannot substitute another step’s groups for them. Without this the pump waits on a group no one is generating.The step trains on fewer samples than configured, which is the point: a smaller step beats a stalled run. The count is logged as
dropped_prompt_groupsso the batch size a step actually used is recoverable afterwards.How much smaller is bounded by
min_step_batch_fraction, and that bound has to live here because neither drop budget provides it. Both budgets are run-scoped – the consecutive counter is cleared by any commit, including commits for other steps – so drops landing on one step while other steps succeed can shrink it without ever tripping them.- Raises:
RuntimeError – The step fell below
min_step_batch_fractionofnum_prompts_per_step. Training a fraction of a batch is a silent change to the gradient estimate, so it is refused rather than absorbed.
- async _train_pump() None#
Per-prompt-group streaming train loop.
Per step, with 1-4 running per streaming chunk and 5-6 once the chunk loop closes: 1. Select the rollouts to train on. a. sampler.evict drops stale groups from the buffer and clears their TQ rows. b. sampler.select returns K prompt groups, or None, and removes them from the buffer. The DP rows survive, already training-shaped because the buffer wrote them that way at rollout time. c. One _buffer_capacity permit is released per group that left the buffer. 2. Prepare the batch. a. Policy and reference logprobs. b. Value model forward (PPO only), with the policy parked on CPU so the critic never shares the training GPUs with it. c. _advantage_stage. 3. Train on the chunk. a. Value model (PPO only): train_from_meta, which is a whole optimizer step. That is why a PPO step is pinned to a single chunk – the value workers have no split train API yet (#2625). b. Policy model: train_microbatches_from_meta, which only accumulates gradients. c. PPO only: all critic updates run before all policy updates. Their counts are ppo.critic_ppo_epochs and ppo.ppo_epochs, respectively. Each policy optimizer step closes here rather than in 5 – a PPO step is one chunk, so there is nothing to accumulate across chunks. 4. Close the chunk. Refresh min_sample_version and the dispatch tally. The consumed rows stay in TQ: staged capture deltas are read by the policy workers during 3, so nothing is cleared until the step closes. 5. Train the policy model (GRPO) – finish_train_step all_reduces the accumulated gradients, rescales, and runs optimizer.step. Then _cleanup_consumed_metas clears every consumed canonical row and its staged capture deltas. 6. Refit the model. Sync the new policy weights to generation.
PPO critic warmup (ppo.policy_training_start_step > 0) changes which of those run. For the first N steps 3a still trains the critic every step, but 3b and 6 are skipped, so the policy is neither trained nor refit. The trainer version still advances, and the sampler’s lookahead is widened while the policy is frozen.
- async _stall_watchdog_pump() None#
Report rollout health, and detect stalls nothing else catches.
Progress is the pair (committed groups, completed train steps) rather than a timestamp: both counters already exist, and “neither has moved” is the property that actually matters.
Deliberately not conditioned on rollouts being in flight. An earlier version required that, on the reasoning that an idle controller has legitimately no work – and a fault-injection run walked straight through the gap. Killing a generation worker wedged the loop with zero rollouts in flight and zero failures recorded: the rollout pump was blocked on backpressure behind a train pump that could no longer finish a step, so nothing was in flight to count. The watchdog observed six minutes of idleness and said nothing.
What separates a real stall from an idle gap is whether work remains, so that is what is checked instead.
- async _gen_fleet_probe_pump() None#
Probe the generation fleet on its own clock.
Separate from the watchdog because the two cadences answer different questions. The watchdog publishes counters and notices a stalled run, which is a minutes-scale concern; liveness detection is the input to every recovery decision and has to be seconds-scale.
Sharing the watchdog’s loop made
probe_interval_sdecorative – probes ran atwatchdog.interval_sand nothing read the configured value. With the shipped defaults that put detection at30s * unhealthy_threshold, i.e. 60-90s, which is longer than the refit deadline: by the time a hung refit aborted, the monitor still had the dead shard as SUSPECT, so the rebuild that abort exists to trigger saw an empty absent set and did nothing. Arithmetic, not a race – it could never have worked. Job 5925668.
- async _probe_generation_fleet() None#
Ask every serving generation shard whether it is still alive.
Ray actor liveness is the cheap authoritative signal for “the process is gone”, and it is what the probe uses. It does not catch every failure – a vLLM engine core can die while the worker process and its HTTP thread survive – which is why the routing adapters also report the failures they observe. The two signals feed the same counters.
Only serving shards are probed: a quarantined shard answering again says nothing about whether its weights are current, and the monitor ignores such probes anyway.
Shards are probed concurrently. Sequentially, a tick costs up to
probe_timeout_sper shard, so a fleet of four would take 8s to complete a round the config promises every 5s – and config validation only checksprobe_timeout_s < probe_interval_s, which silently assumes one probe per tick. Concurrent, a round is bounded byprobe_timeout_sat any fleet size.
- async _push_router_membership() None#
Tell the NeMo-Gym router which backends are currently serving.
Pushed as the full set rather than a delta, so a dropped or reordered update – or a restarted router, which comes up believing every backend serves – converges on the next tick without sequence numbers or replay.
Pushed unconditionally, not gated on the membership epoch moving. The gate looked free – an unchanged serving set costs nothing to skip – but it made the router’s own restart unrecoverable: a recreated actor rebuilds
_servingas every backend, while the epoch it was last pushed at has not moved, so the gate blocked every corrective push and Gym routed to a quarantined shard for the rest of the run. The payload is a short list of strings on a probe-interval timer; the gate bought nothing and cost the guarantee both docstrings advertised.It is also what makes the router’s reflex drop safe: dropping a failing backend locally is only correct because a later push puts it back.
- async _drain_router_failures() None#
Fold the router’s observed backend outcomes into the fleet ledger.
The router is the only component that sees a wedged engine: it answers
is_alivefrom a healthy worker process, so no probe can condemn it. The router holds no monitor reference by design – membership flows one way – so it counts per backend URL and this drains them here, on the tick that already talks to it.BOTH halves, and the successes first.
consecutive_reported_failuresis the only counter that can condemn a wedged engine, and it is a streak – but on the router path nothing ever cleared it.report_successhas exactly one caller, on the native adapter, so a router run made the streak monotonic and every shard reachedunhealthy_thresholdeventually, however healthy. Three unrelated blips days apart, with thousands of successes between them, condemned a shard.Successes first because they describe the same window: replaying failures onto a streak that a success in that window should already have cleared is what made the count monotonic in the first place. A genuinely wedged shard produces no successes, so its condemnation timing is unchanged.
Deliberately not a reset on a clean probe: the reported streak is kept separate from the probe streak precisely because a wedged engine still answers
is_alive.
- _stand_down_refit_deadline(shard_idx: int) None#
Tell every policy worker to cancel an in-flight refit deadline.
Fire-and-forget, and deliberately not awaited: this runs inside a probe whose job is to keep the ledger current, and a worker that cannot answer is already the larger problem. Reaching the worker at all depends on the refit running off its event loop – see await_off_loop.
Only ever called for a CONFIRMED actor death. A frozen rank never produces one, so its deadline still fires and still ends the run attributably.
- async _reconcile_refit_membership(
- force: bool = False,
Ask the weight transport to match the live fleet before the refit runs.
A no-op without fleet health: with no monitor there is no notion of a shard being gone, so the transport keeps the membership it was built with – which is the pre-existing behaviour, and why this is inert by default.
Returns whether the communicator was actually rebuilt, or None when the transport owns no membership at all. The recovery path needs all three: after an abort the old communicator is gone, so “nothing to reconcile” means there is nothing to retry with either – but “this transport has no membership” is a different refusal and deserves a different message.
forcesays the communicator is gone rather than merely unchanged. The synchronizers skip a rebuild when the absent set matches what they last built with, which is what stops a lost shard costing two full rebuilds on every subsequent step – but after an abort the membership is identical and the communicator is dead, so the recovery path has to override that or it retries over nothing.
- _recovery_window() collections.abc.Iterator[None]#
Mark the span where the serving set is deliberately empty.
_recover_from_failed_refit marks every serving shard partial, so they all go STALE and serving_shards() is empty until _record_refit_landed runs – after a rebuild and a full retry refit, both of which await and yield the event loop.
_stall_watchdog_pump is a task on that same loop and calls raise_if_exhausted() on every tick, defaulting to one every 30s. With min_healthy_shards=1 and zero serving, any tick landing in this window ends the run over an exhausted fleet while the retry that would have refilled it is still in flight – killing the recovery that was about to succeed, and blaming the wrong thing in the log.
A flag rather than deferring mark_weights_partial until after the rebuild: that would also close the window, but it would leave shards holding a mix of old and new weights in the serving set while it did.
- async _recover_from_failed_refit(failure: BaseException) None#
Drop whatever stopped participating, rebuild the communicator, allow a retry.
Two failures arrive here and they are not the same event:
RefitAborted– a rank went silent inside the collective and a worker’s watchdog broke it. Every engine that was receiving is left holding a mix of old and new weights, so none of them may serve until a refit completes.RayActorError– the collective finished and a shard died in the epilogue, before its RPC returned. Nothing is partial; the survivors have complete weights. Left uncaught this killed a run whose data transfer had already succeeded.Both need the same repair, because both leave a communicator that no longer matches the fleet, and in the abort case no communicator at all.
The probe here is the point. Waiting for the health monitor to reach its own conclusion is what failed before: its verdict is paced by probe rounds while this is an event, so the abort arrived first and the rebuild saw an empty absent set and did nothing. Asking now, on this thread, turns a race into a lookup.
- _refit_participants() set[int]#
Shards eligible to receive this refit’s weights, as of right now.
Captured at the moment membership settles rather than read at promotion time, because a restart finishing mid-transfer turns its shard STALE – which is not absent – and the communicator was already built without it.
Derived from the fleet rather than from the transport’s membership so it holds for backends that own no membership at all: a shard that is absent when the transfer starts receives nothing either way.
- _record_refit_landed(participants: set[int]) None#
Write down what each shard now holds, and return the STALE ones to service.
Two things, because they are the same fact seen from two sides: this refit reached these shards. The version is what they hold; promotion is what that entitles them to.
Promotion is the exit from STALE, and the reason marking partial weights is safe rather than terminal. An aborted refit leaves every engine that was receiving with a mix of old and new weights, so they are pulled out of service – but nothing else moves a shard out of STALE, so without this the recovery would succeed and then leave the fleet empty, which
raise_if_exhaustedwould end the run over. A worse failure than the one being recovered from, and reached only on the recovery path.Only STALE shards are promoted. A SUSPECT shard also took part in the refit, but it is failing probes for its own reasons and promoting it here would reset the failure count that is supposed to condemn it. It is still stamped: what weights an engine holds is not a verdict on how well it is serving them.
And only STALE shards that were IN the refit. Asking “is this shard STALE?” alone was correct until restart existed, because nothing could turn a shard STALE while a refit was in flight. A restart can: it takes minutes, nothing blocks it, and mark_loaded moves the shard DEAD -> STALE at whatever moment the reload lands. A shard absent when membership settled received no weights from this transfer, so promoting it would return it to service holding the checkpoint it read off disk – the outcome this module’s docstring exists to prevent. It stays STALE, is not absent, and the next refit picks it up.
The stamp used to live inside
report_refitalone, which meant it was only ever written by a promotion. Nothing turns a shard STALE on a refit that succeeds, so a fleet that has never lost a shard reports version 0 for the life of the run however many refits it received – and a metric that reads 0 on every healthy shard is one nobody watches, which is the part that matters: this is the reading that would catch the next bug of this shape.- Parameters:
participants – shards eligible for this refit, from :meth:
_refit_participantsat the point membership settled.
- async _check_env_health(timeout_s: float) list[str]#
Ask each environment actor that exposes a health check whether it is whole.
Returns the problems found, empty when everything is well. It reports rather than raises so the caller can route the verdict through
stall_action, the same way the stall path does. Raising here bypassedstall_actionentirely: under the documented default ("warn", which promises to “only report”), and withgym_subprocess_checkdefaulting to true, an unhealthy environment killed the run – a run-ending path switched on by default, in a feature whose whole posture is inert-by-default.Each probe is bounded.
NemoGymis an asyncio actor, so a wedged environment – precisely the case this check exists to catch – left the await hanging forever, the pump never ticked again, and stall detection was dead exactly when it was needed. A probe that does not answer within one tick IS the unhealthy signal; it is not a reason to stop watching.Environments without the method are skipped rather than treated as unhealthy; only NeMo-Gym has subprocess servers to lose.
- async _abort_stale_inflight() int#
Abort in-flight rollouts that the sampler can no longer select.
- async _capture_rollout_checkpoint_cut(
- cut: nemo_rl.algorithms.async_utils.replay_buffer.DataPlaneMutationCut,
- checkpoint_path: nemo_rl.utils.checkpoint.PathLike,
Save TQ and capture matching restart state under the barrier.
Groups selected by an unfinished streamed step are absent from the live replay index but remain in TQ. Re-index them only in this persisted cut; the live trainer keeps accumulating gradients without modification.
- async _write_rollout_checkpoint_sidecars(
- checkpoint_path: pathlib.Path,
- cut: nemo_rl.algorithms.single_controller._RolloutCheckpointCut,
Write controller state beside TQ and return its on-disk byte size.
- async _save_rollout_checkpoint(
- *,
- force: bool = False,
Publish one rollout-only snapshot anchored to durable trainer state.
- _log_rollout_checkpoint_outcome(
- *,
- outcome: nemo_rl.algorithms.single_controller_utils.rollout_checkpoint.RolloutCheckpointAttemptOutcome,
- reason: nemo_rl.algorithms.single_controller_utils.rollout_checkpoint.RolloutCheckpointAttemptReason,
- attempt_duration_seconds: float,
Record every scheduled checkpoint attempt, including no-op cuts.
- async _rollout_checkpoint_pump() None#
Persist rollout state periodically, including during streamed train.
- async _rollout_telemetry_pump() None#
Sample generation and publication throughput at a fixed cadence.
- async _log_rollout_throughput_metrics(*, emit: bool = True) None#
Serialize backend sampling so cumulative counters have one baseline.
- async _collect_and_log_rollout_throughput_metrics(
- *,
- emit: bool = True,
Compare generation-engine output with canonical TQ publication.
- async _save_checkpoint(
- step_metrics: dict[str, Any],
- *,
- is_policy_training_step: bool,
- is_final_checkpoint: bool,
Serialize full and rollout-only checkpoint publication.
- async _save_checkpoint_impl(
- step_metrics: dict[str, Any],
- *,
- is_policy_training_step: bool,
- is_final_checkpoint: bool,
Write a full checkpoint for the just-finished train step.
Everything except the (possibly async) policy weight write must be on disk before
begin_finalization. Non-colocated engines keep serving rollouts throughout. Colocated engines are stood down for the save. The policy optimizer is skipped during critic warmup: it has never stepped.
- _REFIT_UNWIND_GRACE_S#
60.0
- _refit_await_budget_s() Optional[float]#
How long to wait for the refit before giving up, or None to wait forever.
- async _sync_weights_within(kv_scales, what: str) None#
Run the refit off-loop, and stop waiting if it outlives the deadline.
WHY THIS IS NEEDED ON TOP OF EVERY WORKER-SIDE BOUND. A frozen-but-alive rank is a Ray actor that never answers, and Ray puts no timeout on an actor call. So the controller’s await never resolves however well the workers behave. Job 6508251 measured that end state: the deadline fired, the workers aborted, the trainers returned, every actor was idle – and the run still sat for 1800s because this await had no bound. It is the last unbounded wait on the refit path, and unlike the others it is not in NCCL or CUDA.
Timing out raises RefitAborted so this joins the existing recovery path rather than inventing a second one: the caller reconciles membership and retries once. By then the fleet probe has usually condemned the silent shard, so the rebuild can exclude it and the retry may genuinely succeed.
A DEDICATED DAEMON THREAD, not asyncio.to_thread. to_thread runs on the default ThreadPoolExecutor, whose workers are non-daemon and are joined at interpreter exit – so a thread still blocked on the frozen actor would hang shutdown, trading a wedge in the refit for a wedge on the way out. asyncio.wait_for cannot cancel a running thread either way; this only controls whether the orphan can block exit.
The orphan is a real consequence, not a free win. If the retry succeeds the run continues with a thread still parked inside the old sync_weights, and if that rank were ever resumed it would wake up holding a communicator that has since been replaced. Acceptable against a guaranteed 30-minute stall, and it is why this path is bounded-failure-first rather than resume-and-forget.
- async _sync_weights(
- *,
- calibration_data: Optional[nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]] = None,
Pause new rollout dispatches, synchronize weights, resume.
SC owns the pause gate. vLLM serves through the refit; it supports live weight updates. A colocated engine is instead already stood down: the synchronizer’s sync is the wake, and step 4 resumes dispatch.
Flow:
_rollout_permitted.clear() — no new dispatches
Optionally calibrate FP8 KV-cache scales.
Materialize deferred policy parameter all-gathers.
weight_synchronizer.sync_weights(kv_scales=…)
_rollout_permitted.set() — resume
- Parameters:
calibration_data – Optional data used to calibrate FP8 KV-cache scales before synchronizing weights.
- Returns:
The number of stale in-flight rollout groups aborted before the weight synchronization.
- async _value_stage(
- meta: nemo_rl.data_plane.KVBatchMeta,
Run the PPO value model’s forward pass over the selected chunk.
Tensors never touch SC: workers fetch the sequence columns from DataPlane and commit the per-token prediction back under
values, which the advantage stage then reads alongside the rewards. The value model is loaded and offloaded around the call, so it holds the training GPUs only for the duration of the forward.- Returns:
The batch metadata with the
valuescolumn recorded on it.
- async _value_train_epochs(
- meta: nemo_rl.data_plane.KVBatchMeta,
- *,
- num_epochs: int,
Run consecutive critic epochs under one model onload/offload cycle.
- Returns:
The final epoch’s
train_from_metaoutput; earlier epochs’ results are discarded.
- async _advantage_stage(
- meta: nemo_rl.data_plane.KVBatchMeta,
Fetch advantage inputs, compute advantages, and write them back.
SC owns the prompt-group-scoped advantage stage because the selected
KVBatchMetastill contains complete prompt groups before trainer DP sharding. Tensor payloads still move through DataPlane: SC fetches only the configured advantage input columns and writes the computedadvantagescolumn back under the samesample_ids.- Returns:
The updated batch metadata and whether the batch contains at least one valid training token.
- async _run_advantage_stage( ) nemo_rl.algorithms.single_controller_utils.advantage_stage.AdvantageOutcome#
Run one advantage stage on the pool, or in-process without one.
- _absorb_advantage_outcome( ) None#
Fold one call’s reduced results into this step’s accumulators.
- _retune_lookahead_versions() None#
Widen the sampler’s lookahead while the policy is frozen, then shrink it back.
Port of ppo.py’s _async_ppo_generation_lead_steps.