Model and Inference Engineering · Staff
FSDP trains with bf16. Can we export its current compute weights as a resumable checkpoint?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
Not safely as a drop-in replacement for the training state. In a common FSDP mixed-precision setup, forward and backward use lower-precision parameter copies while the sharded parameters kept outside those phases are full precision for optimizer updates. PyTorch's FSDP documentation says param_dtype controls forward and backward precision and its model checkpoints save parameters in full precision. A custom exporter that captures the temporary bf16 compute tensors and labels them as the training checkpoint has thrown away lower bits of the master parameters. Loading them into fp32 storage later cannot recover those bits.
There are actually two artifacts a team might want. A serving artifact can intentionally convert weights to bf16 or another supported inference format, then pass output-quality tests. A resumable training checkpoint should preserve the canonical model parameters, optimizer state, scheduler state, random state and data progress according to the recipe. Restoring fp32 optimizer moments next to rounded bf16-exported weights is not equivalent to restoring the original trajectory. The loss might look normal for a few batches while updates drift over a long continuation.
I would inspect dtypes and hashes at three points: the sharded parameter outside forward, the temporary compute representation, and the saved state dict. On a small deterministic run, save a reference checkpoint through the documented FSDP state-dict path and another through the proposed exporter. Restore both with matching optimizer and data state, perform one or more identical updates, and compare parameter deltas. This does not demand bit-identical output across every distributed kernel, but it will expose a large and systematic rounding error from the exporter.
A candidate may say that bf16 has a wide exponent range, so it should be fine. Range is not precision. The mantissa still discards detail when fp32 values are rounded. If the project deliberately trains with low-precision master parameters, its contract differs, so inspect that actual implementation. Rank 7 dies while the training checkpoint is being written. Which checkpoint can you resume? is about whether the checkpoint is complete after a rank dies. The optimizer checkpoint loads. Did its moments attach to the right parameters? is about matching optimizer state to parameters. Here all files can load and still represent a different numerical starting point.
Continue reading
Related questions
Read beyond the question
Explore more model and inference engineering
Follow another question in this area, or search the complete Question Library.
Browse this area →Browse Question Library →