Skip to content

[BugFix] Preserve replay epoch boundaries when resuming prefetch - #4507

Open
cyanseek wants to merge 2 commits into
pytorch:mainfrom
cyanseek:codex/replay-prefetch-epoch-resume
Open

cyanseek wants to merge 2 commits into
pytorch:mainfrom
cyanseek:codex/replay-prefetch-epoch-resume

Conversation

@cyanseek

Copy link
Copy Markdown

Description

Preserve the current replay epoch when restarting iteration with an exhausted sampler and unconsumed prefetched batches. ReplayBuffer.__iter__ now resets ran_out only after the prefetch queue has drained.

For example, with six items, batch size three and prefetch=1, consuming [0, 1, 2] and allowing the last batch to finish prefetching should leave only [3, 4, 5]. Previously, restarting iteration returned that batch followed by an additional complete epoch. The same duplication occurred after checkpoint restoration.

Motivation and Context

Fixes #4506. This prevents replay samples from being repeated when resuming a partially consumed epoch. Sampler behavior, checkpoint formats and public APIs are unchanged.

Types of changes

  • Bug fix (non-breaking change which fixes an issue).

Checklist

  • I have read the CONTRIBUTING guide.
  • My change requires a change to the documentation.
  • I have updated the tests accordingly.
  • I have updated the documentation accordingly.

Validation

The new public-API regression covers both drop_last settings and live, state-dict, disk and pickle continuation. It checks the exact remaining batch and the next complete epoch, using synchronize() to make the boundary deterministic without mocks or sleeps. All eight cases fail on unmodified main 8f92a50c on each of CPU and CUDA. With the fix, all 16 CPU/CUDA cases pass, including with warnings treated as errors. CUDA parameters carry the gpu marker so GPU CI collects them.

  • python -m pytest test/rb/test_rb_core.py test/rb/test_samplers.py -m 'not gpu' -q --tb=short --timeout=90: 918 passed, 6 GPU cases deselected.
  • CPU environment: Python 3.12.3, PyTorch 2.11.0+cpu, TensorDict 0.14.2, locally built TorchRL C++ extension. The wider suite emitted NumPy deprecation and multiprocessing warnings; the new regression passes with -W error.
  • NVIDIA H200 NVL, driver 580.159.03, PyTorch 2.7.1+cu128, Python 3.12, TensorDict 0.14.2: python -m pytest test/rb/test_rb_core.py test/rb/test_samplers.py -m gpu -q --tb=short --timeout=90 gives 14 passed, 918 deselected. This includes the 8 new CUDA regressions and all 6 existing GPU cases in these files. Four existing weight-updater deprecation warnings were emitted during collection.
  • On the same GPU host, python -m pytest test/rb/test_rb_core.py -k resume_prefetched_epoch -q -W error --tb=short --timeout=90 gives 16 passed.
  • Pre-commit checks on the changed files and git diff --check pass.

The full repository suite and end-to-end training were not run. This is a host-side iteration correctness fix with no performance claim.

@pytorch-bot

pytorch-bot Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/rl/4507

Note: Links to docs will display an error until the docs builds have been completed.

⚠️ 16 Awaiting Approval

As of commit bf41b10 with merge base 8f92a50 (image):

AWAITING APPROVAL - The following workflows need approval before CI can run:

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Sep 30, 2026

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

BugFix CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. ReplayBuffers

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[BUG] Resuming a prefetched replay epoch repeats the next epoch

1 participant