PyTorch has added a per-parameter mixed-precision policy to FSDP2, allowing distributed-training workloads to use differing parameter compute dtypes without necessarily adding all-gather operations. For developers using these layouts, FSDP groups staging copies by source and target dtype while retaining one all-gather per fully_shard.
On reduce-scatter, gradients from mixed compute dtypes still need to enter a buffer with one dtype. When parameters share a reduction dtype, PyTorch says it preserves the existing single-collective behavior by extending the CUDA _chunk_cat.out path to read BF16 and FP32 inputs and cast them directly into the reduction buffer during chunking and padding. When effective reduction dtypes genuinely differ, FSDP creates a bucket and reduce-scatter for each dtype.
PyTorch’s isolated H100 test, using a 50.3-million-element transformer-like workload with small FP32 regions, measured 0.1888 ms for fused mixed BF16/FP32-to-FP32 packing versus 0.3582 ms for pre-casting followed by homogeneous packing—about 47% lower packing latency—and avoided roughly 192 MiB of temporary FP32 storage.
The trade-off was slower all-gather packing: complete copy-in was about 2.4–2.5% slower in layouts with roughly 500 parameters. In an eight-block, 1.75-billion-parameter KDA-style FSDP2 test, all-gather copy-in was 4.2–4.9% slower, but synchronized forward-plus-backward time stayed near 53 ms in both configurations. These are PyTorch’s reported results for the tested workloads, not a guarantee for every model or dtype layout.
Definition. FSDP2 per-parameter mixed precision lets distributed-training parameters use different compute dtypes while managing their all-gather and reduce-scatter behavior.
| Test or behavior | Reported result |
|---|---|
| Isolated H100 packing test | Fused mixed BF16/FP32-to-FP32 packing measured 0.1888 ms versus 0.3582 ms for pre-casting and homogeneous packing. |
| Temporary FP32 storage | The fused approach avoided roughly 192 MiB in the isolated H100 test. |
| All-gather copy-in | Complete copy-in was about 2.4–2.5% slower in layouts with roughly 500 parameters. |
| KDA-style FSDP2 test | All-gather copy-in was 4.2–4.9% slower, while synchronized forward-plus-backward time stayed near 53 ms in both configurations. |
Key takeaways
- FSDP groups staging copies by source and target dtype while retaining one all-gather per fully_shard.
- When parameters share a reduction dtype, PyTorch preserves single-collective behavior by casting BF16 and FP32 inputs directly into the reduction buffer.
- When effective reduction dtypes differ, FSDP creates a bucket and reduce-scatter for each dtype.
- PyTorch reported 0.1888 ms fused packing versus 0.3582 ms for pre-casting and homogeneous packing in an isolated H100 test.
- The fused approach avoided roughly 192 MiB of temporary FP32 storage in that test.
- PyTorch reported slower all-gather copy-in in tested layouts, while synchronized forward-plus-backward time remained near 53 ms in its KDA-style test.
FAQ
What does the new FSDP2 policy allow?
It allows distributed-training workloads to use different parameter compute dtypes without necessarily adding all-gather operations.
When does FSDP use multiple reduce-scatter operations?
FSDP creates a bucket and reduce-scatter for each dtype when effective reduction dtypes genuinely differ.
What performance trade-off did PyTorch report?
PyTorch reported lower mixed-dtype packing latency but slower all-gather copy-in in the tested layouts.
Do these benchmark results apply to every model?
No. PyTorch described the results as applying to the tested workloads, not as a guarantee for every model or dtype layout.