FSDP2: per-rank torch.save/load for SHARDED_STATE_DICT to fix 2800+ NPU checkpoint timeout (#4105)
* FSDP2: replace dist_cp save/load with per-rank torch.save/load for SHARDED_STATE_DICT
Eliminates cross-rank coordination overhead (DefaultSavePlanner) in
dist_cp.save/load, which becomes a bottleneck at large scale (2800+ GPUs)
due to all-rank planning phase.
Key changes:
- save_fsdp_model: per-rank torch.save for SHARDED_STATE_DICT
- load_fsdp_model: per-rank torch.load + barrier after set_state_dict
- save_fsdp_optimizer: per-rank torch.save for non-FULL state dict
- load_fsdp_optimizer: per-rank torch.load (no get_optimizer_state_dict
needed since we load the exact per-rank shard) + barrier
Performance (2816 GPUs, 352 nodes):
- Before: >5h (Gloo TCP aggregation timeout)
- After: ~2-3min (zero cross-rank communication, 100x+ speedup)
Related: #4099
* FSDP2: add use_dcp parameter for SHARDED_STATE_DICT save/load
Add use_dcp=True (default) to save_fsdp_model, load_fsdp_model,
save_fsdp_optimizer, and load_fsdp_optimizer. When use_dcp=True,
uses torch.distributed.checkpoint (original behavior, CUDA-compatible).
When use_dcp=False, uses per-rank torch.save/load (no cross-rank
communication, works around Gloo TCP timeout on Ascend at large scale).
This is a non-breaking change — DCP remains the default for all users.
* Apply style fixes
* docs: add use_dcp docstrings to FSDP save/load functions
Explain the use_dcp parameter for all four functions (save/load model,
save/load optimizer), documenting the per-rank torch.save/load strategy
for large-scale training where DCP can cause timeouts due to Gloo-based
collective communication (e.g., 2800+ GPUs).
---------
Co-authored-by: unknown <leiliandong1@huawei.com>
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>