Repository navigation
Conversation
🔗 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.
|
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
Preserve the current replay epoch when restarting iteration with an exhausted sampler and unconsumed prefetched batches.
ReplayBuffer.__iter__now resetsran_outonly 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
Checklist
Validation
The new public-API regression covers both
drop_lastsettings and live, state-dict, disk and pickle continuation. It checks the exact remaining batch and the next complete epoch, usingsynchronize()to make the boundary deterministic without mocks or sleeps. All eight cases fail on unmodified main8f92a50con each of CPU and CUDA. With the fix, all 16 CPU/CUDA cases pass, including with warnings treated as errors. CUDA parameters carry thegpumarker 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.-W error.python -m pytest test/rb/test_rb_core.py test/rb/test_samplers.py -m gpu -q --tb=short --timeout=90gives 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.python -m pytest test/rb/test_rb_core.py -k resume_prefetched_epoch -q -W error --tb=short --timeout=90gives 16 passed.git diff --checkpass.The full repository suite and end-to-end training were not run. This is a host-side iteration correctness fix with no performance claim.