Repository navigation
FSDP x TP Saving / Loading - #48802
michaelbenayoun wants to merge 47 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. |
| # being emitted if the function is not used | ||
| import torch.distributed.checkpoint as dcp | ||
| from torch.distributed.checkpoint.state_dict import get_optimizer_state_dict | ||
| from torch.distributed.checkpoint.state_dict import StateDictOptions, get_optimizer_state_dict |
There was a problem hiding this comment.
same import top + guarding distributed as well. Make sure as well to check torch version availabilities for the func
There was a problem hiding this comment.
But the comment at the top mentions that it is imported here on purpose to avoid warnings.
There was a problem hiding this comment.
okay then add guarding. We have to make sure the user downloaded a torch version that has DCP
There was a problem hiding this comment.
The guard is already there, line 514.
michaelbenayoun
left a comment
There was a problem hiding this comment.
I've addressed your comments. WDYT?
0e1cc55 to
7b88913
Compare
| transformers_explicit_filename=getattr(config, "transformers_weights", None), | ||
| tqdm_class=tqdm_class, | ||
| ) | ||
| # Local checkpoints saved with `save_pretrained(..., distributed_checkpoint=True)` hold rank-local shards |
There was a problem hiding this comment.
At this point, the model is not even created so it is not even distributed, hence not a DistributedMixin.
I could, of course, create functions or static methods, but is it worth it?
We have 2 blocks for this logic in from_pretrained.
There was a problem hiding this comment.
I was thinking to have something like:
distributed_checkpoint_dir = get_distributed_checkpoint_dir(pretrained_model_name_or_path, subfolder, gguf_file, state_dict)
if distributed_checkpoint_dir is not None:
checkpoint_files, sharded_metadata = _get_resolved_distributed_checkpoint_files(distributed_checkpoint_dir, hf_quantizer)
else:
checkpoint_files, sharded_metadata = _get_resolved_checkpoint_files(
pretrained_model_name_or_path=pretrained_model_name_or_path,
...
)| ) | ||
| loading_info, disk_offload_index = cls._load_pretrained_model(model, state_dict, checkpoint_files, load_config) | ||
| loading_info = cls._finalize_model_loading(model, load_config, loading_info) | ||
| if distributed_checkpoint_dir is not None: |
There was a problem hiding this comment.
same here we can do like above
if distributed_checkpoint_dir is not None:
loading_info, disk_offload_index = cls._load_and_finalize_distributed_pretrained_model(
model, distributed_checkpoint_dir, load_config
)
else:
loading_info, disk_offload_index = cls._load_pretrained_model(
model, state_dict, checkpoint_files, load_config
)
loading_info = cls._finalize_model_loading(model, load_config, loading_info)| _load_sharded_checkpoint_in_distributed_model(model, checkpoint_dir, strict=strict) | ||
| elif is_sharded_checkpoint(os.path.join(checkpoint_dir, "sharded")): | ||
| _load_sharded_checkpoint_in_distributed_model(model, os.path.join(checkpoint_dir, "sharded"), strict=strict) | ||
| elif os.path.isfile(safe_index_file): |
There was a problem hiding this comment.
why is there 2 paths when checkpoint is consolidate ?
There was a problem hiding this comment.
It depends if the file is one safetensors file or a bunch of shards.
Refactored to separate the logic of getting the filenames and calling _load_sharded_checkpoint_in_distributed_model.
|
Sorry a bit late, reviewing now! |
| import torch | ||
|
|
||
|
|
||
| def _check_distributed_checkpointing_available(raise_if_not: bool = True) -> bool: |
There was a problem hiding this comment.
i prefer not having the raise_if_not boolean and have the error raise happening where this is called. If we keep it, we find ourselves with pattern in the code like
if not _check_distributed_checkpointing_available(raise_if_not=False):
raise OSError("Loading a distributed checkpoint requires torch>=2.7.")
which is not consistent (as there is a double raise but one was turn off)
| from transformers.distributed import DistributedConfig | ||
| from transformers.distributed.checkpoint import load_model_checkpoint_distributed | ||
|
|
||
| if dist.is_available(): |
There was a problem hiding this comment.
use is_torch_distributed_available() as it uses across the codebase
CI recapDashboard: View test results in Grafana |
| token=token, | ||
| **hub_kwargs, | ||
| # Native DCP checkpoints remain sharded, use `distributed_checkpoint=False` for an interoperable checkpoint. | ||
| consolidate=False, |
There was a problem hiding this comment.
we should have an utils function called consolidate_distributed_checkpoint() in transformers.distributed.checkpoint
|
|
||
|
|
||
| def _prepare_state_dict_for_dcp(state_dict): | ||
| """ |
There was a problem hiding this comment.
the docstring example doesnt make it obvious to understand what's going on (especially with __create_write_items__, __create_chunk_list__ and __get_tensor_shard__)
| transformers_explicit_filename=getattr(config, "transformers_weights", None), | ||
| tqdm_class=tqdm_class, | ||
| ) | ||
| # Local checkpoints saved with `save_pretrained(..., distributed_checkpoint=True)` hold rank-local shards |
There was a problem hiding this comment.
I was thinking to have something like:
distributed_checkpoint_dir = get_distributed_checkpoint_dir(pretrained_model_name_or_path, subfolder, gguf_file, state_dict)
if distributed_checkpoint_dir is not None:
checkpoint_files, sharded_metadata = _get_resolved_distributed_checkpoint_files(distributed_checkpoint_dir, hf_quantizer)
else:
checkpoint_files, sharded_metadata = _get_resolved_checkpoint_files(
pretrained_model_name_or_path=pretrained_model_name_or_path,
...
)| ) | ||
| loading_info, disk_offload_index = cls._load_pretrained_model(model, state_dict, checkpoint_files, load_config) | ||
| loading_info = cls._finalize_model_loading(model, load_config, loading_info) | ||
| if distributed_checkpoint_dir is not None: |
There was a problem hiding this comment.
same here we can do like above
if distributed_checkpoint_dir is not None:
loading_info, disk_offload_index = cls._load_and_finalize_distributed_pretrained_model(
model, distributed_checkpoint_dir, load_config
)
else:
loading_info, disk_offload_index = cls._load_pretrained_model(
model, state_dict, checkpoint_files, load_config
)
loading_info = cls._finalize_model_loading(model, load_config, loading_info)| safe_index_file = os.path.join(checkpoint_dir, SAFE_WEIGHTS_INDEX_NAME) | ||
| safe_weights_file = os.path.join(checkpoint_dir, SAFE_WEIGHTS_NAME) | ||
|
|
||
| for dcp_dir in (checkpoint_dir, os.path.join(checkpoint_dir, "sharded")): |
There was a problem hiding this comment.
since we hardcode the save_distributed_checkpoint path to consolidate=False, we only have 1 case which is the non-consolidated DCP checkpoint:
self.save_distributed_checkpoint(
model_to_save,
save_directory,
# Native DCP checkpoints remain sharded, use `distributed_checkpoint=False` for an interoperable checkpoint.
consolidate=False,
)
no need the for loop then
| return | ||
|
|
||
| checkpoint_files = [] | ||
| if os.path.isfile(safe_index_file): |
There was a problem hiding this comment.
do we really need this part ? i dont see how this part can be triggered with the current save_pretrained(distributed_checkpoint=True|False)/from_pretrained() API .
| set_model_state_dict(model, state) | ||
|
|
||
|
|
||
| def load_model_checkpoint_distributed(model, checkpoint_dir: str | os.PathLike, strict: bool = True) -> None: |
There was a problem hiding this comment.
I think we want load_model_checkpoint_distributed and save_model_checkpoint_distributed to be private function so that users dont call those functions and rely on save_pretrained(distributed_checkpoint=True|False)/from_pretrained() API only
|
We are getting there ! Main points:
|
What does this PR do?
This PR handles a few thing related to distributed training.
1.
_StridedShardplacement issueThe
_StridedShardplacement own multiple disjoint regions of a tensor.It produces overlapping checkpoint metadata, causing failure because the default PyTorch checkpoint adapter was designed with the idea that each rank's shard is a single contiguous chunk.
The propose solution consists in redistributing
_StridedShard -> Shard. While it forces additional communication and the temporary allocation of the full tensor, it has the advantage to be a super lightweight solution.2. Distributed checkpointing / Consolidation
Load and save distributed checkpoints for both the model.
3. Optimizer checkpointing
Handled in #49256
Script
And the output is: