Meta's AI Training Fix Makes Big Models Run Smoother
This technical fix could mean faster AI updates and fewer crashes for everyone.
FSDPParamGroup.finalize_backward may clean up an all-gather launched by backward prefetch but never consumed. It recorded a device-specific event and waited through torch.accelerator.current_stream(), whose generic Stream.wait_event does not handle that device-specific event correctly during CUDA graph capture.
The all-gather stream remained unjoined, so capture failed with cudaErrorStreamCaptureUnjoined. The fix is to use the FSDP device handle for the wait, matching every other FSDP wait_event call. A FakePG CUDA graph gradient accumulation test was added that leaves a prefetched all-gather pending at finalization and replays the graph.
- Meta fixed a bug that caused crashes when training very large AI models.
- The fix ensures memory operations finish properly, making training more reliable.
- This could lead to faster AI improvements and fewer disruptions in services you use.
Why It Matters
More reliable AI training means quicker, better AI tools for everyday use.