accelerate
422a5292 - FSDP2: per-rank torch.save/load for SHARDED_STATE_DICT to fix 2800+ NPU checkpoint timeout (#4105)

Commit
73 days ago
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>
Author
Parents
Loading