[realfill] Fix the LR scheduler when num_train_epochs is passed in a distributed training env - #14885
Conversation
…distributed training env
|
Hi @825pranav, thanks for the PR! It does not appear to link an issue it fixes. If this PR addresses an existing issue, please add a closing keyword (e.g. Please note that PRs without a linked issue are likely to be automatically closed 10 days after this notice. Once the PR links an issue (or gets the |
|
The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update. |
|
Thanks @sayakpaul for reviewing and merging! Happy to take another script from #8384 if that's useful. |
What does this PR do?
Part of #8384, for
examples/research_projects/realfill/train_realfill.py(claimed here).The LR scheduler was sized from the dataloader length before
accelerator.prepare(), so with multiple processes it was built for the unsharded length and the LR never finished decaying. This computes the step count from the per-process length, same approach as #8312.One difference from the other scripts: realfill doesn't pass
lr_schedulertoaccelerator.prepare(), so it steps once per batch on each process. I kept that and scale bygradient_accumulation_stepsinstead ofnum_processes.Checked with a toy model using the same loop (unprepared scheduler,
accelerator.accumulate, 2 epochs, GA=2, 24 samples), run with 2 CPU processes (gloo):Minimal command:
ruff checkandruff format --checkpass.Before submitting
Who can review?
@sayakpaul @geniuspatrick