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.

Source: PyTorch release

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 behaviorReported result
Isolated H100 packing testFused mixed BF16/FP32-to-FP32 packing measured 0.1888 ms versus 0.3582 ms for pre-casting and homogeneous packing.
Temporary FP32 storageThe fused approach avoided roughly 192 MiB in the isolated H100 test.
All-gather copy-inComplete copy-in was about 2.4–2.5% slower in layouts with roughly 500 parameters.
KDA-style FSDP2 testAll-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.