Skip to content

[RLlib] Skip updates if some learners don't have batches, prevent deadlocks - #66030

Open
ArturNiederfahrenhorst wants to merge 27 commits into
ray-project:masterfrom
ArturNiederfahrenhorst:rllib-skip-empty-batch
Open

ArturNiederfahrenhorst wants to merge 27 commits into
ray-project:masterfrom
ArturNiederfahrenhorst:rllib-skip-empty-batch

Conversation

@ArturNiederfahrenhorst

@ArturNiederfahrenhorst ArturNiederfahrenhorst commented Sep 9, 2026 •

Copy link
Copy Markdown
Contributor

Description

With learners, we can arrive in situations where a) we have multiple learners and b) not enough data is produced to feed all learners at a given DDP step. Today, this creates an empty batch with an unbound loss_per_module.

The naive solution is to skip empty batches but that creates deadlocks in torch DDP as workers wait for each other's steps.
I looked into different solutions of how to solve this, such as having blueprint of valid batches to step with if an empty batch arrives but this solution is messy in practice as each learner needs to have the correct bluebrint at all times (even before the first valid training batch arrives).

This PR proposes as solution where, prior to a DDP step, all learners agree that they have data available and how often they will step on any available minibatches.

Known limitation: differing module sets across Learners

The agreement added here is per Learner, not per module. Each submodule of a MultiRLModule is wrapped in its own TorchDDPRLModule (TorchLearner._make_modules_ddp_if_necessary), so every module has its own all-reduce, and MultiRLModule._forward_train only forwards the modules that are present in the batch it is given.

If two Learners receive batches with different sets of module IDs — for example in multi-agent setups where an agent does not appear in every shard of the sampled episodes — the Learner that is missing a module never runs that module's backward pass, while its peers wait on that module's all-reduce, and the group deadlocks. _should_skip_update cannot express this case: it is a single boolean per Learner.

This is pre-existing behaviour that this PR neither introduces nor makes worse; it was confirmed on two CPU-DDP Learners fed {p0, p1} and {p0} respectively.

A smaller residual of the same limitation is reachable through sharding. ShardBatchIterator now spreads the remainder over the leading shards, so a shard can only be left without rows for a module that holds fewer rows than there are Learners. When that happens the group skips the whole update rather than dropping the module on one Learner, for the reason above. With a single Learner there is no group to stay in step with, so the module is dropped instead and the rest of the batch trains.

A follow-up can close it: self.module.keys() is identical on every Learner, so the Learners can agree on a per-module bitmask (AND'ed across the group) and train the intersection, filtering the batch to it. The all-or-nothing skip in this PR then becomes the special case of an empty intersection, and the sharding residual above disappears with it.

Behavior changes

  • A skipped update no longer runs the gradient-based update's hooks. before_gradient_based_update and after_gradient_based_update bracket the minibatch loop, so neither fires when there is no loop: PPO would otherwise read back a KL this update never measured (the metric keeps its key and peeks as NaN, so it warns about a divergence that did not happen), and DQN, whose target sync is gated on sampled timesteps, could take a Polyak step toward an online network that has not moved.
  • num_epochs > 1 without minibatch_size now takes the widest module's rows as the minibatch, not batch.count. That field is env steps -- a different unit, equal to the rows only for a single module, and one a shard cannot report faithfully. Multi-agent runs on that path train the same rows in num_epochs full passes instead of more, smaller cycles whose size depended on how many agents shared an env step.
  • learner_env_steps_dropped_on_skip_lifetime is now learner_module_steps_dropped_on_skip_lifetime and counts rows. Under the old name a skip that discarded 128 rows reported 1, because a shard takes its env steps from whichever module was sliced last.
  • An update() from no episodes at all (episodes=[]; a Learner's shard of a short list of episode refs, or every episode it was sent lost with its EnvRunner) is skipped like an empty batch. Previously the learner connector pipeline ran on the empty list, and pieces such as AddOneTsToEpisodesAndTruncate raised IndexError -- in a group before the skip agreement, leaving the peers waiting in it.
  • never_skip_update=True now only turns the skip into an error, raised on every Learner of the group after the plan agreement: _should_skip_update is not consulted, and a train batch without timesteps for a module raises, naming the module on the Learner that received it. Raising before the agreement would leave the peers waiting in it. The Learners still agree on the number of minibatches. The collective the flag used to save is also what keeps unequal shards from stepping a different number of times, a deadlock that master avoided with the driver-side count this PR removes (4001 timesteps over two Learners shard as 2000 and 2001; over 500-row minibatches that is 4 steps against 5). Both never_skip_update and _should_skip_update are marked experimental.
  • TorchDifferentiableLearner and TorchMetaLearner forward, differentiate and step only the modules that have data in the minibatch; a module dropped for having no rows passes its parameters through unchanged (the functional update used to walk every module and raise KeyError). While at it, a pre-existing bug in TorchDifferentiableLearner.compute_gradients is fixed: torch.autograd.grad returns one flat tuple, and the old mapping zipped every module's parameters against it from its start, so the second and later modules were stepped by the first module's gradients.
  • TorchMetaLearner.update calls before_gradient_based_update after the skip decision, so a skipped meta update runs neither hook, as Learner.update does.
  • ShardBatchIterator shards a batch without any module into empty batches instead of raising UnboundLocalError.
  • A skipped update no longer advances the weights' sequence number, so the off-policyness metric and the EnvRunners' weight sync count only real updates.

Regression test overview

Every regression test added here was run against the commit before its fix and fails, or hangs.
The following is a table that compares behavior before/after to make it easer to review @pseudo-rnd-thoughts

Test Without its fix
test_minibatch_count_is_fixed_without_minibatch_size the plan proposes num_minibatches=0, leaving each Learner to derive its own
test_epochs_survive_a_dropped_module 129 module steps instead of 258 -- one pass where two were asked for
test_skipped_update_leaves_the_kl_coeff_alone KL divergence for Module default_policy is non-finite
test_module_without_rows_is_dropped, test_update_is_skipped_when_no_module_has_rows ValueError: One of the module batches is empty!
test_shard_batch_iterator_spreads_the_remainder 5 rows over 4 shards split 2/2/1/0
test_update_without_episodes_is_skipped IndexError: list index out of range in add_one_ts_to_episodes_and_truncate.py
test_learner_group_skips_when_a_learner_receives_no_episodes hangs: the starved Learner raises in its connector, its peer waits in the agreement
test_never_skip_update (count assertion) _sync_update_plan is never called
test_never_skip_update_in_a_group hangs: 8 vs 2 minibatches; with one empty shard, the empty Learner raises before the agreement and its peer waits in it
test_never_skip_update (agreement assertion) the empty batch raises without taking part in the agreement
test_update_empty_batch_is_skipped (sequence number) each skipped update advances the weights' sequence number
test_module_without_rows_is_dropped (the update() part) KeyError: 'm2'
test_each_module_is_stepped_by_its_own_gradient m2 is stepped by m1's gradient
test_hooks_do_not_run_for_a_skipped_update before_gradient_based_update runs once, after_gradient_based_update never
test_shard_batch_iterator_shards_a_batch_without_modules, test_learners_agree_on_the_update_plan (whole-group empty batch) UnboundLocalError on end

Manual Multi-GPU tests (not in CI)

I ran APPO Cartpole and multi agent PPO Cartpole to catch sneaky side effects.
Validation was on 4x L4 with ray 2.56.0 plus this PR's diff:

  • NCCL, two GPU Learners. Every scenario behaved as expected. Normal and skipped updates leave both Learners bit-identical. Shards of 96 and 32 rows step the same 8 minibatches on both Learners, also with never_skip_update=True. Twelve async rounds with two empty shards return all 24 results with exact skip counts. APPO's Learner thread skips episode-ref shards that come up empty, with its simple queue and with its circular buffer. With the agreement patched out, the same scenarios desync or hang, so they can fail. With never_skip_update=True and one empty shard, the group hung because the empty Learner raised before the agreement (Bugbot's finding); a43d3bb4f3 fixes that after this run, covered by test_never_skip_update_in_a_group on CPU.
  • Learning, 2.56.0 vs this PR, two GPU Learners and four EnvRunners, both arms side by side on each seed:
Experiment Reached threshold Env steps/s
APPO CartPole 3/3 vs 3/3 4,309 vs 4,318
Multi-agent PPO CartPole 3/5 vs 4/5 2,232 vs 2,195

No training run skipped an update. A sixth multi-agent seed stalls in its first update on both arms, from a pre-existing very rare sharding bug that #66530 fixes separately.

APPO throughput comparison (3 arms x 3 seeds: unpatched 2.56.0 / this PR / this PR with the plan collective disabled, which is what never_skip_update=True does) showed no regression: 4753 vs 4757 env steps trained/s,

MiniBatchCyclicIterator had two termination modes: an explicit
`num_total_minibatches` count, and a data-dependent one that stopped once every
module's data had been covered `num_epochs` times. The two are equivalent -- the
data-dependent mode stops at the smallest k for which every module m satisfies
floor(k * step_m / n_m) >= num_epochs, that is k = max_m ceil(num_epochs * n_m /
step_m) -- but only the count is knowable before iterating.

Keep the count as the only termination condition. Constructed without
`num_total_minibatches`, the iterator now derives it up front through the new
classmethod `num_minibatches()`, which callers that need the number in advance can
use as well. Cursors, wrap-arounds and per-epoch shuffles are untouched, so the
sequence of yielded minibatches is unchanged.

Extract `_len_and_step()` while here: the single rule for how much of a module's
batch one minibatch consumes -- timesteps, or sequences for batches sliced in the
B dimension. Both the count and the iteration loop use it, which removes a
duplicated SEQ_LENS check and turns the pathological case of a minibatch smaller
than one sequence, previously an infinite loop, into an error.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
An empty train batch -- no timesteps for any module, e.g. because sampled episodes
were lost to EnvRunner or node failures, or because `policies_to_train` excludes
every module -- crashed `Learner.update()` with an UnboundLocalError: the minibatch
loop ran zero times, leaving `loss_per_module` unbound. Skipping the update is not
sufficient either. With `num_learners > 1` the Learners run one all-reduce per
minibatch, so a Learner that skips on its own, or that steps through a different
number of minibatches than its peers, drops out of the collective sequence and
deadlocks them.

The Learners therefore agree on how to carry out each update, in one collective:

- `Learner._should_skip_update(batch)` (new, overridable) returns whether this
  Learner wants to skip; by default, when its batch holds no module data.
- `Learner._sync_update_plan(plan)` reconciles that verdict together with the number
  of minibatches to step through (`UpdatePlan`). The base implementation returns the
  plan unchanged; `TorchLearner` implements it as a single all-reduce over the DDP
  process group: skip if any Learner wants to, and step the average of the proposed
  minibatch counts. All Learners then either skip together or step the same number of
  times -- the latter also covers shards of unequal size, which can desync a group on
  its own.
- `_create_iterator_if_necessary()` returns None to signal "skip this update", warns
  once and counts it. Overrides of `update()` only have to check for None and return
  early.
- `AlgorithmConfig.learners(never_skip_update=True)` opts out of all of it: no
  reconciliation, no per-update collective, and an empty train batch raises instead.

Because every Learner now derives the minibatch count from its own shard,
`TrainingData.shard()` no longer estimates one for the episodes path.

`DifferentiableLearner` skips locally, without agreement: it computes gradients
through `torch.autograd.grad`, which engages no collectives.

Three lifetime metrics make starved Learners diagnosable:
`learner_update_skipped_empty_batch_lifetime` (this Learner's own batch was empty),
`learner_update_skipped_for_peer_lifetime` (it had data but skipped to stay in sync)
and `learner_env_steps_dropped_on_skip_lifetime` (timesteps discarded by skips).

Also fixes `logging.getLogger("__name__")` in torch_meta_learner.py and
torch_differentiable_learner.py.

Tested with pytest on rllib/core/learner/tests/test_learner.py (12 passed),
rllib/utils/tests/test_minibatch_utils.py (3 passed),
rllib/algorithms/appo/tests/test_appo_multi_agent_data_balance.py (2 passed) and the
CPU scaling modes of rllib/core/learner/tests/test_learner_group.py (2 passed).
Multi-Learner behaviour -- no deadlock, identical weights across ranks, synchronous
and asynchronous -- was validated separately on a two-GPU NCCL setup.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Changes in this file are not related to the PR

Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
@@ -34,7 +34,7 @@
TensorType,
)

logger = logging.getLogger("__name__")
logger = logging.getLogger(__name__)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Unrelated bug fix.

@@ -407,7 +415,7 @@ def _create_iterator_if_necessary(
minibatch_size: Optional[int] = None,
shuffle_batch_per_epoch: bool = False,
**kwargs,
) -> Iterable:
) -> Optional[Iterable]:

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Before, we'd always return an empty dict.
That's not idiomatic in python.
If we are not able create some iterable, let's log that issue >here< and return None.

Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
`shard()` returned `(TrainingData, kwargs)` pairs so that sharding could pass extra
per-shard keyword arguments on to each Learner's `update()` call. The only value
ever sent that way was `num_total_minibatches`, which the Learners now derive
themselves, so every branch returned an empty dict and the single caller splatted
nothing.

Return a plain list of `TrainingData` instead, and drop the `**kwargs` the method
no longer reads. Those were only ever used to look up `minibatch_size` and
`num_epochs` by name, which coupled the data-splitting code to the signature of
`Learner.update()`.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
training_data.shard(
num_shards=len(self),
len_lookback_buffer=self.config.episode_lookback_horizon,
**kwargs,

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

We don't need that anymore as now the number of minibatch steps (which was the only kwarg we used here) is computed in the learners.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Code Review

This pull request introduces synchronization of update plans across multi-learner groups in RLlib to prevent deadlocks in distributed setups, along with a never_skip_update configuration option to opt out of skipping empty batches. The review identified three critical issues: a missing import of ALL_MODULES in differentiable_learner.py that will cause a NameError, a potential TypeError in _sync_update_plan when num_total_minibatches is None, and a potential ZeroDivisionError in minibatch_utils.py if a module batch is empty.

Comment thread rllib/core/learner/differentiable_learner.py
Comment thread rllib/core/learner/learner.py Outdated
Comment thread rllib/utils/minibatch_utils.py

@cursor cursor Bot left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Stale Bugbot comment from a previous run.

Comment thread rllib/core/learner/learner.py Outdated
@ArturNiederfahrenhorst

Copy link
Copy Markdown
Contributor Author

Did 9 runs with 2 GPUs each, 3 on master, 3 with the learner sync on this PR and 3 on this PR, but never_skip_updates enabled (so no learner sync). Used CartPole with a small network to minimize update time and maximize the possible impact of the sync. Could not measure the impact so I think we are good from a performance standpoint.

@ray-gardener ray-gardener Bot added the rllib RLlib related issues label Sep 14, 2026
@@ -49,43 +47,32 @@ def shard(
self,
num_shards: int,
len_lookback_buffer: Optional[int] = None,
**kwargs,

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

We don't need that anymore as now the number of minibatch steps (which was the only kwarg we used here) is computed in the learners.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

The main point of the changes in this file is that we want to know how many minibatches we produce up front. This enables us to agree on a shared plan among the learners.

Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

@cursor cursor Bot left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Stale Bugbot comment from a previous run.

Comment thread rllib/utils/minibatch_utils.py
@pseudo-rnd-thoughts pseudo-rnd-thoughts added the go add ONLY when ready to merge, run all tests label Sep 16, 2026

@pseudo-rnd-thoughts pseudo-rnd-thoughts left a comment •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I'll finish my review tomorrow but I thought that Simon fixed the known limitation where in multi-agent settings if all learners don't get all module-ids in a batch then a deadlock can happen.

Also, could you send me the code to replicate this issue, I want to understand if my Ray Train hang detector can help for this

@cursor cursor Bot left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Stale Bugbot comment from a previous run.

Comment thread rllib/core/learner/learner.py Outdated
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

@pseudo-rnd-thoughts pseudo-rnd-thoughts left a comment •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Ok, checking out the PR, claude spotted the following things

update(episodes=[]) still raises before the skip check. AddOneTsToEpisodesAndTruncate (rllib/connectors/learner/add_one_ts_to_episodes_and_truncate.py:120) does episodes[0], which raises IndexError in the PPO/IMPALA/APPO/BC/MARWIL/IQL learner pipelines before _should_skip_update runs. The PR's empty-episodes test passes only because the testing connector doesn't include that piece. So the IMPALA case where a package has fewer episode refs than learners (ShardObjectRefIterator gives one learner []) is not covered yet. One fix: return an empty MultiAgentBatch from _make_batch_if_necessary when episodes == [], without calling the connector.

Overall: the sync point and the count-only MiniBatchCyclicIterator look right to me.

My understanding is that if you have missing data on one of the learners, then we skip that training step for all learners. While this is a weird case edge and I think the right decision, we need to document this for users.

Claude is suggesting narrowing it to a per-module payload (rows per module on each rank), with learners taking part with zero gradients for modules they don't have.

Also, could _should_skip_update and never_skip_update be marked experimental, so they don't become public API we later have to deprecate?

Comment thread rllib/core/learner/learner.py
Comment thread rllib/core/learner/learner.py Outdated
Comment thread rllib/core/learner/learner.py Outdated
Comment thread rllib/core/learner/learner.py Outdated
… count module steps

- `ShardBatchIterator` spreads the remainder over the leading shards, so no
  shard is left empty once a module has at least `num_shards` rows.
- With `num_learners <= 1` a module without rows is dropped and the rest of
  the batch trains; only a group has to skip, because its Learners cannot
  train different module sets (one all-reduce per module).
- Rename `LEARNER_ENV_STEPS_DROPPED_ON_SKIP_LIFETIME` to
  `LEARNER_MODULE_STEPS_DROPPED_ON_SKIP_LIFETIME` and log module steps: a
  shard's `env_steps()` is the row count of whichever module was sliced last.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…batch_size`

`MiniBatchCyclicIterator` also cycles the batch when only `num_epochs` is set,
using `batch.count` rows per module. That fallback was resolved *after* the
Learners agreed on a plan, so every Learner proposed 0 and then worked the
count out from its own shard. In multi-agent runs the shards disagree --
`ShardBatchIterator` gives a shard the row count of whichever module it sliced
last -- so one extra row in the largest module makes one Learner take an extra
step and deadlock the group on its next all-reduce.

Resolve the fallback before the agreement, and propose the count whenever there
is minibatching at all: a lone Learner proposes exactly what it would have
derived, so only the group's behavior changes. The `elif num_epochs > 1` branch
that picked the iterator is now unreachable and goes away.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
`after_gradient_based_update` runs for a skipped update too, so that schedulers
and target-network syncs keep following the sampled timesteps. Metrics are a
different matter: a windowed metric keeps its key from earlier updates and peeks
as NaN once its window is empty, which is indistinguishable from an update that
really did diverge. PPO read the KL back that way and warned "KL divergence ...
is non-finite" on every skipped update, sending users after a model problem that
is not there (the coefficient itself was safe, since every comparison against
NaN is False).

Record whether the update in flight was skipped in `Learner._update_skipped` and
have `PPOLearner` consult it. No change to the hook's signature, which subclasses
outside RLlib override.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

@cursor cursor Bot left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Stale Bugbot comment from a previous run.

Comment thread rllib/core/learner/learner.py
A shard takes its env steps from whichever module `ShardBatchIterator` sliced
last, so dropping that module for having no rows can leave `batch.count` at 0
while the rest of the batch still holds data. Read back as the minibatch size,
that 0 is falsy and quietly turns `num_epochs` passes into a single one.

Fall back to the widest module's rows, which is what "one minibatch is the whole
batch" means once `count` is unusable. `batch.count` is still preferred wherever
it says something, so nothing changes for batches that carry real env steps.

Reported by Cursor Bugbot on 58a8942.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…s env steps

`num_epochs` > 1 without `minibatch_size` means "one minibatch is the whole
batch", and `MiniBatchCyclicIterator` reads that size as rows taken from every
module. `batch.count` is env steps: a different unit, equal to the rows only
when there is a single module, and one a shard cannot report faithfully at all
-- `ShardBatchIterator` splits each module separately, so its shards have no env
steps of their own and take the row count of whichever module was sliced last.

Take the widest module's rows instead. The count then works out to exactly
`num_epochs` on every Learner whatever the shards look like, since the governing
module's ratio cancels, which removes the disagreement that made this path
deadlock rather than only settling it beforehand. Multi-agent runs that reach
this path change behavior: the same rows are trained in `num_epochs` full passes
instead of more, smaller cycles whose size depended on how many agents shared an
env step.

`DifferentiableLearner` carried the same line and the same bug.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…date

`after_gradient_based_update` is an extension point, and calling it for an update
that never happened asks every subclass to know that. Worse, the ones in RLlib
already act on it: PPO reads back a KL this update never measured (the metric
keeps its key and peeks as NaN, so it warns about a divergence that did not
happen), and DQN, whose target sync is gated on sampled timesteps, can take a
Polyak step toward an online network that has not moved since the last one.

So skip the hooks with the update instead of flagging it: `Learner._update_skipped`
and PPO's extra guard are gone again. `before_gradient_based_update` moves below
the point where the batch is built and the decision is made, so the two always run
as a pair -- IMPALA sets its entropy coefficient there and still does so before
the loss. `TorchMetaLearner`'s skip path drops its hook call too, which is what
`DifferentiableLearner` already did.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

@cursor cursor Bot left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Stale Bugbot comment from a previous run.

Comment thread rllib/core/learner/torch/torch_meta_learner.py
…ut rows

Its `_create_iterator_if_necessary` is a copy of `Learner`'s and drifted: it
skipped only for a batch with no modules at all, so a shard carrying a ModuleID
with zero rows -- which is what `ShardBatchIterator` hands out, and what the
meta-learner passes straight through to its inner learners -- still reached
`MiniBatchCyclicIterator` and raised. That is the crash this PR removes from
`Learner`, surviving in its sibling.

It computes gradients with `torch.autograd.grad` and takes part in no collective,
so it is in the position of a `Learner` without peers and can make the same
decision: drop the modules that have no rows, skip only when nothing is left.

First tests for the class, covering both outcomes. They need a concrete subclass
for the abstract loss, which is never reached.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

@cursor cursor Bot left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Stale Bugbot comment from a previous run.

Comment thread rllib/core/learner/differentiable_learner.py
…tor on nothing

A Learner receives `episodes=[]` when its shard of a short list of episodes or
episode refs is empty -- `ShardObjectRefIterator` hands out `[]` whenever an
update carries fewer episode refs than there are Learners, which is routine in
IMPALA/APPO when one EnvRunner's sample already fills an update's package --
or when every episode it was sent was lost with its EnvRunner.

`_make_batch_if_necessary` ran the learner connector pipeline on that empty
list. PPO, IMPALA, APPO and MARWIL start their pipelines with
`AddOneTsToEpisodesAndTruncate`, the off-policy algorithms with
`AddNextObservationsFromEpisodesToTrainBatch`; both index `episodes[0]` and
raise before `_should_skip_update` gets to see the batch. In a group that
leaves the peers waiting in `_sync_update_plan` -- the deadlock this change
set is meant to remove. The existing `update(episodes=[])` test only passed
because the testing pipeline has neither piece.

Hand back an empty `MultiAgentBatch` without calling the pipeline; the skip
logic takes it from there. Same for `DifferentiableLearner`.

Tests: a PPO Learner updated from `episodes=[]` now skips (fails with the
`IndexError` above without the fix), and a two-Learner PPO group fed a single
episode ref skips together, the starved Learner as empty and its peer for it.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
`never_skip_update=True` bypassed the whole plan-agreement branch, so it also
opted the group out of settling on one number of minibatches -- and nothing
replaced the driver-side count that `TrainingData.shard()` used to compute for
the episodes path before this change set removed it. A group whose shards
differ in size then steps a different number of times per Learner and
deadlocks, skip or no skip: 4001 timesteps over two Learners shard as 2000 and
2001, which over 500-row minibatches is 4 steps against 5.

The flag now only turns the skip into an error: `_should_skip_update` is not
consulted, a train batch without timesteps for a module raises (naming the
module), and the Learners still propose and settle their minibatch count
through `_sync_update_plan`. The collective this saved was measured at
0.1-0.2 ms per update; the deadlock it now prevents has no timeout.

Tests: the single-Learner test asserts the count is still proposed with the
flag on; a two-Learner group with the flag and shards of 256 and 64 rows
settles on 5 minibatches of 32 (hangs without the fix).

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…odules with data

`_create_iterator_if_necessary` drops a module without rows from the batch
(like one excluded via `policies_to_train`), but the functional update still
walked every module of the `MultiRLModule`: `_make_functional_call` indexed
`batch[mid]` for all of them and raised `KeyError` on the first real
`update()` -- the new test only exercised the iterator. `TorchMetaLearner`
had the same call.

The functional forward now covers the modules in the batch, `compute_gradients`
differentiates the parameters of the modules that contributed to the loss, and
`apply_gradients` passes every other module's parameters -- and any parameter
without a gradient -- through unchanged, so the returned dict still holds all
modules for the meta-learner's next call.

While at it: the gradients come back from `torch.autograd.grad` as one flat
tuple, and the old mapping zipped every module's parameters against that
tuple from its start, handing the second and later modules the first
module's gradients. The mapping now consumes the tuple in order.

Tests: the dropped-module test runs `update()` end to end (`m1` trains, `m2`
passes through), and a new test checks that each of two modules with data
takes one step along the gradient of its own loss; both fail without the fix.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
Review request: neither should become public API that later needs a
deprecation cycle. The hook carries `@ExperimentalAPI` and both docstrings
say so.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…en there is an update

`TorchMetaLearner.update` still called `before_gradient_based_update` ahead of
the skip decision and returned without `after_gradient_based_update` on a skip,
while `Learner.update` runs both hooks as a pair after that decision. Move the
call below the skip check, so a skipped meta update runs neither hook.

Test: a meta learner over one inner learner, updated from an empty batch, runs
neither hook, and a real update runs both once (fails without the fix:
`before_gradient_based_update` runs once on the skip).

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…ising

`ShardBatchIterator` took the shard's env steps from the loop variables of the
last module it sliced, so a batch without any module left them unbound and
`LearnerGroup.update(batch=...)` raised `UnboundLocalError` on the driver.
Such a batch is a valid input now that empty batches are skipped: it shards
into one empty batch per Learner, and every Learner skips it.

Tests: the sharder on a module-less batch, and a two-Learner group updated
from one (both fail with the `UnboundLocalError` without the fix).

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

@cursor cursor Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Cursor Bugbot has reviewed your changes using default effort and found 2 potential issues.

Fix All in Cursor

Reviewed by Cursor Bugbot for commit 61abb96. Configure here.

Comment thread rllib/core/learner/learner.py
Comment thread rllib/core/learner/torch/torch_differentiable_learner.py
…he agreement

With `never_skip_update=True`, a Learner whose train batch had no timesteps
for a module raised before `_sync_update_plan`, while its peers entered that
collective and waited for it forever. So in a group the flag's error never
surfaced: the update hung (reproduced on 2 GPU Learners, as Bugbot reported).

The starved Learner now votes to skip in the agreement like the default path
does, and when the group's plan says skip, every Learner raises -- the starved
one naming its modules, its peers naming the cause. All Learners leave
`update()` after the same collective, so the error reaches the driver and the
group's collectives stay aligned.

Tests: a two-Learner group with the flag and one empty shard raises the flag's
error and trains in sync on the next update (hangs without the fix); the
single-Learner test checks the empty batch took part in the agreement before
raising (it never called it before).

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…e docstrings

`Learner.update`, `DifferentiableLearner.update` and `TorchMetaLearner.update`
advanced the weights' sequence number before deciding to skip, so a skipped
update counted as a gradient update in the off-policyness metric and in the
number EnvRunners compare when they sync weights. The increment now follows the
skip check; a real update sees the same number as before.

Docstrings:
- `_should_skip_update`: with a single Learner, modules without timesteps are
  dropped before the hook runs, so a lone Learner only skips a batch in which
  no module has timesteps left.
- `DifferentiableLearner.compute_gradients` / `apply_gradients`: gradients
  cover only the modules in the loss; `apply_gradients` returns every module,
  passing the others through.

Tests: skipped updates leave the sequence number unchanged on a `Learner`, a
`TorchMetaLearner` and a `DifferentiableLearner`; a real update advances it by
one.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

# Call `before_gradient_based_update` to allow for non-gradient based
# preparations-, logging-, and update logic to happen.
self.before_gradient_based_update(timesteps=timesteps or {})

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

If there is no update, we don't want to call this hook

# parameters for each module. Implement this via `foreach_module`.
grads = torch.autograd.grad(
total_loss,
sum((list(param.values()) for mid, param in params.items()), []),

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

This worked only if there were gradients for all modules.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

This is basically a copy of the changes in TorchLearner. They share most of their logic that we had to change here which is its own problem that we'll have to tackle in another PR.

if self.config.num_learners <= 1:
for module_id in list(batch.policy_batches.keys()):
if len(batch.policy_batches[module_id]) == 0:
del batch.policy_batches[module_id]

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I think it's a bit awkward that this changes the learning behavior between learning on 1 and 2 GPUs in some multi agent cases. But it's better than any alternative it seems.

if not minibatch_size and num_epochs > 1:
minibatch_size = max(
(len(b) for b in batch.policy_batches.values()), default=0
)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

This is essentially a hack around the fact that batch.count is unreliable and we don't yet have a good way to resolve it.

Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>
…ner.py

The previous commit deleted `rllib/core/learner/tests/test_differentiable_learner.py`
but left its `py_test` target in `rllib/BUILD.bazel`, so `lint: pytest_format`
failed with "path is missing!" for that file.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Signed-off-by: Artur Niederfahrenhorst <artur@anyscale.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

go add ONLY when ready to merge, run all tests rllib RLlib related issues

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants