Skip to content

fix(storage): retain rollout remainders in mini-batch generators - #244

Draft
Afloat16 wants to merge 1 commit into
leggedrobotics:mainfrom
Afloat16:fix-rollout-minibatch-remainders-20261009
Draft

Afloat16 wants to merge 1 commit into
leggedrobotics:mainfrom
Afloat16:fix-rollout-minibatch-remainders-20261009

Conversation

@Afloat16

@Afloat16 Afloat16 commented Oct 9, 2026

Copy link
Copy Markdown
Contributor

When the rollout size is not divisible by num_mini_batches, the feedforward generator shuffles only num_mini_batches * (batch_size // num_mini_batches) transitions. The recurrent generator similarly slices only num_mini_batches * (num_envs // num_mini_batches) environments. This permanently excludes the tail from every training epoch. For five environments, three rollout steps and two mini-batches, the feedforward path drops one transition and the recurrent path drops all three transitions from the final environment.

Partition the complete rollout/environment range using integer boundaries i * size // num_mini_batches. All samples are included exactly once per epoch, the number of optimizer updates stays the same, and batch sizes differ by at most one. Recurrent observations, masks and initial GRU/LSTM states continue to use the matching trajectory boundaries. Feedforward shuffling still happens once per generator, preserving the existing order across epochs. Invalid counts outside 1..size now raise a clear ValueError instead of producing empty batches or dividing by zero.

The regression tests check complete per-epoch coverage, alignment of every batch field, recurrent episode-start states, real GRU/LSTM backward passes and a real PPO update whose only learnable cue occurs in the discarded tail. Divisible partitions retain their previous batch contents, order and Torch RNG state exactly in 24 additional profiles (1,144 recorded arrays).

Validation with Python 3.12.14, Torch 2.7.1+cpu, TensorDict 0.7.2 and NumPy 2.3.5:

  • New regression module: 30 passed. On the unchanged base, 19 valid-rollout checks fail and six invalid-count checks expose the older error behavior; five controls pass.
  • pytest -q tests -k 'not TestDistributedAlgorithms': 216 passed, four deselected. The same selected scope on the unchanged base gives 25 failed and 191 passed, with those same four profiles deselected.
  • pre-commit run --all-files: all configured checks passed.

The four actual Gloo multiprocess tests were not run in this local environment; the complete upstream distributed matrix remains necessary. GPU/simulator training, long-run convergence and other Python-version matrices were not tested. This draft makes no training-performance claim.

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant