Repository navigation
FSDP x TP Saving / Loading, Planner / Writer based solution - #48830
michaelbenayoun wants to merge 26 commits into
Conversation
|
The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update. |
| ) | ||
| # Wait for rank 0 to finish writing the HF safetensors so other | ||
| # ranks don't return (and hit `from_pretrained`) before the files exist. | ||
| dcp.save(state_dict, storage_writer=writer, planner=HuggingFaceSavePlanner()) |
There was a problem hiding this comment.
what is the motivation to have our own writer ? The following works out of the box already iirc
from torch.distributed.checkpoint.hf_storage import HuggingFaceStorageWriter
dcp.save(
state_dict,
storage_writer=HuggingFaceStorageWriter(
path=checkpoint_dir,
save_distributed=True,
enable_consolidation=True,
),
)
There was a problem hiding this comment.
The core issue for saving a DTensor with a _StridedShard placement is that the DTensor's default checkpoint hooks create incorrect WriteItem. Check the docstring here to understand what is happening.
So our _CheckpointView creates multiple WriteItems associated with actual contiguous regions in the global tensor.
When the time comes to actual write data the HuggingFaceStorageWriter from PyTorch uses the name of the parameter as key, which in the case of _StridedShard is not unique anymore: we have multiple chunks of a global tensor referring to the same parameter name.
The proposed writer uses unique keys instead: chunk_{idx}, and stores a mapping chunk_{idx} -> parameter name to be able to resolve who is who at loading time.
CI recapDashboard: View test results in Grafana |
|
Closing this PR in favor of #48802 which offers the same features. |
What does this PR do?
This PR handles a few thing related to distributed training.
1.
_StridedShardplacement issueProvides the same functionality as #48802, while avoiding the
_StridedShard -> Shardredistribution used for checkpointing.It extends DCP with:
HuggingFaceSavePlanner: describes each_StridedShardas non-overlapping global regions and resolves each write request to the corresponding local tensor view.HuggingFaceStorageWriter: stores chunks under unique physical keys, preserving their original parameter names and global offsets in metadata. This prevents multiple chunks of the same parameter from overwriting each other within one safetensors file.HuggingFaceLoadPlannerandHuggingFaceStorageReader: map saved chunks into the destination tensor layout, supporting the same mesh, a different mesh, or ordinary tensors.Consolidation produces standard Hugging Face checkpoints.
2. Optimizer checkpoints
Optimizer checkpoints use the same workaround, then restore the original placements after loading. Optimizer state is flattened by parameter name so it can be restored when parameter groups change across FSDP × TP → TP.
DTensors and ordinary tensors use separate optimizer groups, allowing
foreach=Truefor both.3.
clip_grad_norm_alignmentThere were two independent implementation for
clip_grad_norm_to work with a mixture of ordinary and distributed tensors:Current implementation is more aligned with the general
torch.nn.utils.clip_grad_norm_as it supports multiple norm types and tries to get the best of the 2 implementations.A PR to align with the current implementation has been opened in Accelerate as well: huggingface/accelerate#4266
4. Distributed checkpointing / Consolidation
Load and save distributed and consolidated checkpoints for both the model and the optimizer.
Script
And the output: