diff --git a/src/transformers/distributed/checkpoint.py b/src/transformers/distributed/checkpoint.py new file mode 100644 index 000000000000..2e8ec63bcdcd --- /dev/null +++ b/src/transformers/distributed/checkpoint.py @@ -0,0 +1,199 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""DCP planners for DTensors whose local storage contains disjoint global regions.""" + +from itertools import product + +import torch +from torch.distributed.checkpoint import DefaultLoadPlanner, DefaultSavePlanner +from torch.distributed.checkpoint.metadata import ChunkStorageMetadata, MetadataIndex, TensorProperties +from torch.distributed.checkpoint.planner import TensorWriteData, WriteItem, WriteItemType +from torch.distributed.tensor import DTensor, Shard +from torch.distributed.tensor.placement_types import _StridedShard + + +def _slice_regions(regions, start, length): + """Slice a concatenation of (global offset, length) intervals without allocating tensor data.""" + result = [] + local_offset = 0 + for offset, size in regions: + left, right = max(start, local_offset), min(start + length, local_offset + size) + if left < right: + result.append((offset + left - local_offset, right - left)) + local_offset += size + return result + + +class _CheckpointView: + """ + A checkpoint view of a DTensor for saving and loading that exposes its local storage as a set of disjoint global + regions when the DTensor is sharded with `_StridedShard` placements. + + Example: + Context: + - Global tensor: [10, 11, 12, 13, 14, 15, 16, 17] + - Placement: _StridedShard(dim=0, split_factor=2) + - Mesh size: 2 + - Rank 0 local: [10, 11, 14, 15] + + The DTensor hooks would normally expose the local storage as a single contiguous region, which would be saved + as a single chunk: + ``` + __create_write_items__: + [WriteItem(name="weight", global_shape=[8], offset=[0], size=[4])] + + __create_chunk_list__: + [ChunkStorageMetadata(offsets=[0], sizes=[4])] + + __get_tensor_shard__(MetadataIndex("weight", offset=[0])): + [10, 11, 14, 15] + ``` + While the data is correct, the chunk metadata is misleading because it implies that the local storage + corresponds to a single contiguous region of the global tensor, which is not the case. + + The `_CheckpointView` exposes the local storage as two disjoint regions, which will be saved as two chunks: + ``` + __create_write_items__: + [WriteItem(name="weight", global_shape=[8], offset=[0], size=[2]), + WriteItem(name="weight", global_shape=[8], offset=[4], size=[2])] + __create_chunk_list__: + [ChunkStorageMetadata(offsets=[0], sizes=[2]), + ChunkStorageMetadata(offsets=[4], sizes=[2])] + __get_tensor_shard__(MetadataIndex("weight", offset=[0])): + [10, 11] + __get_tensor_shard__(MetadataIndex("weight", offset=[4])): + [14, 15] + ``` + + The view preserves the tensor's placements and exposes slices of its existing local storage. + """ + + def __init__(self, tensor): + self.tensor = tensor + self.chunks = [] + self.views = {} + coordinate = tensor.device_mesh.get_coordinate() + if coordinate is None: + return + # Compute the global regions of the tensor that are stored in the local storage. + # The regions are represented as a list of (global offset, length) intervals for each dimension. + regions = [[(0, size)] for size in tensor.shape] + for axis, placement in enumerate(tensor.placements): + if placement.is_partial(): + raise ValueError("Checkpointing DTensors with Partial placements is unsupported.") + if not isinstance(placement, (Shard, _StridedShard)): + continue + dim = placement.dim + size = sum(length for _, length in regions[dim]) + splits = placement.split_factor if isinstance(placement, _StridedShard) else 1 + split_size = (size + splits - 1) // splits + selected = [] + for split in range(splits): + split_start = split * split_size + length = max(0, min(split_size, size - split_start)) + shard_size = (length + tensor.device_mesh.size(axis) - 1) // tensor.device_mesh.size(axis) + start = min(coordinate[axis] * shard_size, length) + selected.extend(_slice_regions(regions[dim], split_start + start, min(shard_size, length - start))) + regions[dim] = selected + + local = tensor.to_local() + expected_shape = tuple(sum(length for _, length in intervals) for intervals in regions) + if tuple(local.shape) != expected_shape: + raise ValueError(f"Unsupported DTensor layout: expected local shape {expected_shape}, got {local.shape}.") + dimensions = [] + for intervals in regions: + local_offset = 0 + dimension = [] + for offset, length in intervals: + dimension.append((offset, length, local_offset)) + local_offset += length + dimensions.append(dimension) + for region in product(*dimensions): + offsets = torch.Size(item[0] for item in region) + sizes = torch.Size(item[1] for item in region) + view = local[tuple(slice(start, start + length) for _, length, start in region)] + self.chunks.append(ChunkStorageMetadata(offsets, sizes)) + self.views[offsets] = view + + def size(self): + return self.tensor.size() + + def __create_write_items__(self, fqn, object): + return [ + WriteItem( + index=MetadataIndex(fqn, chunk.offsets), + type=WriteItemType.SHARD, + tensor_data=TensorWriteData( + chunk=chunk, properties=TensorProperties.create_from_tensor(self.tensor), size=self.tensor.size() + ), + ) + for chunk in self.chunks + ] + + def __create_chunk_list__(self): + return self.chunks + + def __get_tensor_shard__(self, index): + return self.views[index.offset] + + +def _checkpoint_views(state_dict): + return { + name: _CheckpointView(value) + if isinstance(value, DTensor) and any(isinstance(p, _StridedShard) for p in value.placements) + else value + for name, value in state_dict.items() + } + + +class HuggingFaceSavePlanner(DefaultSavePlanner): + """ + Extend the default DCP save planner with support for `_StridedShard` placements. + """ + + def create_local_plan(self): + original = self.state_dict + self._views = _checkpoint_views(original) + self.state_dict = self._views + try: + return super().create_local_plan() + finally: + self.state_dict = original + + def lookup_object(self, index): + value = self._views[index.fqn] + if isinstance(value, _CheckpointView): + return value.__get_tensor_shard__(index) + return super().lookup_object(index) + + +class HuggingFaceLoadPlanner(DefaultLoadPlanner): + """ + Extend the default DCP load planner with support for `_StridedShard` placements. + """ + + def create_local_plan(self): + original = self.state_dict + self._views = _checkpoint_views(original) + self.state_dict = self._views + try: + return super().create_local_plan() + finally: + self.state_dict = original + + def lookup_tensor(self, index): + value = self._views[index.fqn] + if isinstance(value, _CheckpointView): + return value.__get_tensor_shard__(index) + return super().lookup_tensor(index) diff --git a/src/transformers/distributed/hf_storage.py b/src/transformers/distributed/hf_storage.py new file mode 100644 index 000000000000..6780f6cf1dd5 --- /dev/null +++ b/src/transformers/distributed/hf_storage.py @@ -0,0 +1,172 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Safetensors DCP storage supporting multiple chunks of a parameter per file.""" + +import json +from dataclasses import dataclass + +import torch +from safetensors import safe_open +from safetensors.torch import _getdtype, save +from torch.distributed.checkpoint import DefaultLoadPlanner +from torch.distributed.checkpoint._hf_utils import ( + CUSTOM_METADATA_KEY, + SAVED_OFFSETS_KEY, + _gen_file_name, + _HFStorageInfo, + _metadata_fn, +) +from torch.distributed.checkpoint.hf_storage import ( + HuggingFaceStorageReader as TorchHuggingFaceStorageReader, +) +from torch.distributed.checkpoint.hf_storage import ( + HuggingFaceStorageWriter as TorchHuggingFaceStorageWriter, +) +from torch.distributed.checkpoint.metadata import ( + ChunkStorageMetadata, + Metadata, + MetadataIndex, + StorageMeta, + TensorProperties, + TensorStorageMetadata, +) +from torch.distributed.checkpoint.planner import WriteItemType +from torch.distributed.checkpoint.storage import WriteResult +from torch.futures import Future + + +_CHUNK_VERSION = "transformers_dcp_chunk_version" + + +@dataclass +class _ChunkStorageInfo(_HFStorageInfo): + tensor_key: str + + +class HuggingFaceStorageWriter(TorchHuggingFaceStorageWriter): + """ + Write unique physical chunk keys while retaining logical names in file metadata. + + Unconsolidated files require the matching Transformers reader. Consolidation produces + standard Hugging Face safetensors files and an index with the original parameter names. + """ + + def write_data(self, plan, planner): + storage_plan = plan.storage_data.get("fqn_to_index_mapping") + buckets = self._split_by_storage_plan(storage_plan, plan.items) + highest_index = max(storage_plan.values()) if storage_plan else 1 + results = [] + for file_index, items in buckets.items(): + if not items: + continue + file_name = _gen_file_name(file_index, highest_index, plan.storage_data.get("shard_index")) + tensors, chunks = {}, {} + for number, item in enumerate(items): + if item.type == WriteItemType.BYTE_IO: + raise ValueError("Safetensors checkpoints only support tensor values.") + key = f"chunk_{number}" + tensor = planner.resolve_data(item).detach().to("cpu").contiguous() + tensors[key] = tensor + chunks[key] = { + "fqn": item.index.fqn, + "shape": list(item.tensor_data.size), + SAVED_OFFSETS_KEY: list(item.tensor_data.chunk.offsets), + } + results.append( + WriteResult( + index=item.index, + size_in_bytes=tensor.numel() * tensor.element_size(), + storage_data=_ChunkStorageInfo(file_name, tensor.size(), tensor.dtype, key), + ) + ) + with self.fs.create_stream(self.fs.concat_path(self.path, file_name), "wb") as stream: + stream.write( + save( + tensors, + metadata={"format": "pt", _CHUNK_VERSION: "1", CUSTOM_METADATA_KEY: json.dumps(chunks)}, + ) + ) + future = Future() + future.set_result(results) + return future + + def finish(self, metadata, results): + if self.save_distributed and not self.enable_consolidation: + return + output_path = self.consolidated_output_path or str(self.path) + mapping = self.fqn_to_index_mapping or dict.fromkeys(metadata.state_dict_metadata, 1) + reader = HuggingFaceStorageReader(str(self.path)) + saved_metadata = reader.read_metadata() + reader.set_up_storage_reader(saved_metadata, is_coordinator=True) + weight_map, total_size = {}, 0 + # Read one output file's tensors at a time. Only the coordinator executes finish(). + for file_index in sorted(set(mapping.values())): + tensors = { + name: torch.empty(info.size, dtype=info.properties.dtype) + for name, info in metadata.state_dict_metadata.items() + if mapping[name] == file_index + } + planner = DefaultLoadPlanner() + planner.set_up_planner(tensors, saved_metadata, is_coordinator=True) + reader.read_data(planner.create_local_plan(), planner).wait() + file_name = _gen_file_name(file_index, max(mapping.values())) + with self.fs.create_stream(self.fs.concat_path(output_path, file_name), "wb") as stream: + stream.write(save(tensors, metadata={"format": "pt"})) + weight_map.update(dict.fromkeys(tensors, file_name)) + total_size += sum(t.numel() * t.element_size() for t in tensors.values()) + with self.fs.create_stream(self.fs.concat_path(output_path, _metadata_fn), "w") as stream: + json.dump({"metadata": {"total_size": total_size}, "weight_map": weight_map}, stream, indent=2) + + +class HuggingFaceStorageReader(TorchHuggingFaceStorageReader): + """Read chunk-key safetensors, PyTorch HF-writer safetensors, and ordinary safetensors.""" + + def read_metadata(self): + tensors, storage = {}, {} + for path in self.fs.ls(self.path): + if not path.endswith(".safetensors"): + continue + with safe_open(path, framework="pt") as file: + extra = file.metadata() or {} + version = extra.get(_CHUNK_VERSION) + if version is not None and version != "1": + raise ValueError(f"Unsupported safetensors chunk metadata version: {version}.") + chunks = json.loads(extra.get(CUSTOM_METADATA_KEY, "{}")) + for key in file.keys(): + view = file.get_slice(key) + shape, dtype = torch.Size(view.get_shape()), _getdtype(view.get_dtype()) + info = chunks.get(key, {}) + name = info["fqn"] if version else key + offset = torch.Size(info.get(SAVED_OFFSETS_KEY, [0] * len(shape))) + global_shape = torch.Size(info["shape"] if version else [o + s for o, s in zip(offset, shape)]) + chunk = ChunkStorageMetadata(offset, shape) + if name not in tensors: + tensors[name] = TensorStorageMetadata(TensorProperties(dtype=dtype), global_shape, []) + else: + tensors[name].size = torch.Size(max(a, b) for a, b in zip(tensors[name].size, global_shape)) + tensors[name].chunks.append(chunk) + storage[MetadataIndex(name, offset)] = _ChunkStorageInfo(path, shape, dtype, key) + return Metadata(tensors, storage_data=storage, storage_meta=StorageMeta(load_id=self.load_id)) + + def _process_read_request(self, file, request, planner): + info = self.storage_data[request.storage_index] + slices = tuple( + slice(offset, offset + length) for offset, length in zip(request.storage_offsets, request.lengths) + ) + tensor = file.get_slice(info.tensor_key)[slices] + target = planner.resolve_tensor(request).detach() + if target.shape != tensor.shape: + raise ValueError(f"Checkpoint slice shape {tensor.shape} does not match destination {target.shape}.") + target.copy_(tensor) + planner.commit_tensor(request, target) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 646b8ac102cc..35d9d792bce5 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -34,6 +34,7 @@ _is_torch_distributed_initialized, gather_full_state_dict, initialize_distributed_mesh, + load_model_checkpoint_distributed, save_model_checkpoint_distributed, ) @@ -199,11 +200,25 @@ def should_save_on_this_rank(self, is_main_process: bool) -> bool: save_on_this_rank = save_on_this_rank and _get_torch_distributed_rank() == 0 return save_on_this_rank + def load_distributed_checkpoint(self, checkpoint_dir: str | os.PathLike) -> None: + """Load model weights from a local safetensors checkpoint into this initialized model. + + Pass the directory containing the rank-local safetensors files, or the retained `sharded/` + directory after consolidation. + All ranks must call this method when using a distributed model. The destination may use the + original mesh, a different mesh, or ordinary tensors without a process group. + + This method preserves the destination's placements and requires materialized tensors, it does + not initialize the model architecture or load its configuration or optimizer state. + """ + load_model_checkpoint_distributed(self, checkpoint_dir) + def save_distributed_checkpoint( self, model_to_save, save_directory: str | os.PathLike, *, + consolidate: bool = True, push_to_hub: bool = False, save_on_this_rank: bool = True, repo_id: str | None = None, @@ -212,7 +227,7 @@ def save_distributed_checkpoint( token: str | bool | None = None, create_pr: bool = False, ) -> None: - """Save an FSDP-wrapped model via DCP and optionally push to the Hub.""" + """Save an FSDP-wrapped model as safetensors via DCP and optionally push to the Hub.""" if not is_torch_greater_or_equal("2.7"): raise OSError("save_pretrained(..., distributed_checkpoint=True) requires torch>=2.7.") if not is_fsdp_managed_module(model_to_save): @@ -224,7 +239,7 @@ def save_distributed_checkpoint( "save_pretrained(..., distributed_checkpoint=True) requires the model to have been " "initialized with a distributed_config (_device_mesh is None)." ) - save_model_checkpoint_distributed(model_to_save, save_directory) + save_model_checkpoint_distributed(model_to_save, save_directory, consolidate=consolidate) if push_to_hub and save_on_this_rank: model_card = create_and_tag_model_card(repo_id, self.model_tags, token=token) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 571e4239f087..bbedec1ceeff 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -15,6 +15,7 @@ import os import warnings +from collections import defaultdict from datetime import timedelta from typing import TYPE_CHECKING, TypeGuard @@ -290,15 +291,12 @@ def gather_full_state_dict(model) -> dict[str, torch.Tensor]: return {} -def save_model_checkpoint_distributed(model, checkpoint_dir: str) -> None: - """Save model parameters as standard HF-format sharded safetensors using - DCP + HuggingFaceStorageWriter with consolidation enabled. +def save_model_checkpoint_distributed(model, checkpoint_dir: str, *, consolidate: bool = True) -> None: + """Save rank-local model shards as safetensors with DCP, optionally consolidating them. - Every rank first writes its own shard in parallel under - `/sharded/`, then a consolidation pass reads those shards - and emits HF-compatible `model-*-of-N.safetensors` (+ index) at - `/`. The result is a directory `from_pretrained` reads - through its normal path — no special flag needed at load time. + With `consolidate=True`, rank-local files are kept in `sharded/` and complete weights + are written at the root. Otherwise, load the rank-local files with + `load_distributed_checkpoint`; they are not `from_pretrained` checkpoints. """ if not is_torch_greater_or_equal("2.7"): raise OSError("Distributed checkpointing requires `torch>=2.7`.") @@ -306,47 +304,171 @@ def save_model_checkpoint_distributed(model, checkpoint_dir: str) -> None: # Import here because otherwise it emits a warning every time it's imported on some hardware - this keeps the warning from # being emitted if the function is not used import torch.distributed.checkpoint as dcp - from torch.distributed.checkpoint.hf_storage import HuggingFaceStorageWriter from torch.distributed.checkpoint.state_dict import get_model_state_dict + from .checkpoint import HuggingFaceSavePlanner + from .hf_storage import HuggingFaceStorageWriter + state_dict = get_model_state_dict(model) - dcp.save( - state_dict, - storage_writer=HuggingFaceStorageWriter( - path=checkpoint_dir, - save_distributed=True, - enable_consolidation=True, - ), + writer = HuggingFaceStorageWriter( + path=checkpoint_dir, + save_distributed=True, + enable_consolidation=consolidate, ) - # 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()) + + # All ranks wait until consolidated weights are ready for loading. _distributed_barrier() -def save_optimizer_distributed(model, optimizer, checkpoint_dir: str) -> None: - """Save optimizer state via DCP.""" +def load_model_checkpoint_distributed(model, checkpoint_dir: str | os.PathLike) -> None: + """Load local safetensors weights into an initialized model, preserving its current mesh and placements. + + Pass the directory containing the rank-local safetensors files, or the retained `sharded/` directory after + consolidation. + """ + if not is_torch_greater_or_equal("2.7"): + raise OSError("Distributed checkpointing requires `torch>=2.7`.") + + import torch.distributed.checkpoint as dcp + from torch.distributed.checkpoint.state_dict import get_model_state_dict, set_model_state_dict + + from .checkpoint import HuggingFaceLoadPlanner + from .hf_storage import HuggingFaceStorageReader + + has_safetensors = any(name.endswith(".safetensors") for name in os.listdir(checkpoint_dir)) + if not has_safetensors: + raise ValueError(f"No safetensors files found in {checkpoint_dir}.") + + reader = HuggingFaceStorageReader(str(checkpoint_dir)) + + original_state = get_model_state_dict(model) + if any(value.is_meta for value in original_state.values() if isinstance(value, torch.Tensor)): + raise ValueError("Materialize the model's tensors before loading a distributed checkpoint.") + + dcp.load(original_state, storage_reader=reader, planner=HuggingFaceLoadPlanner()) + set_model_state_dict(model, original_state) + + +def save_optimizer_distributed(model, optimizer, checkpoint_dir: str, *, consolidate: bool = False) -> None: + """Save optimizer state via DCP, optionally also writing `optimizer.pt`. + + Native DCP files are retained in `checkpoint_dir` in both cases. Consolidation + materializes the full optimizer state in rank 0's CPU memory. All ranks must call. + """ if not is_torch_greater_or_equal("2.7"): raise OSError("Distributed checkpointing requires `torch>=2.7`.") # Import here because otherwise it emits a warning every time it's imported on some hardware - this keeps the warning from # 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 + + # Key group options by parameter name so regrouping after a mesh change remains loadable. + options = StateDictOptions(flatten_optimizer_state_dict=True) + from .checkpoint import HuggingFaceSavePlanner + + optimizer_state_dict = get_optimizer_state_dict(model, optimizer, options=options) + dcp.save({"optimizer": optimizer_state_dict}, checkpoint_id=checkpoint_dir, planner=HuggingFaceSavePlanner()) + if consolidate: + if _get_torch_distributed_rank() == 0: + from torch.distributed.checkpoint.format_utils import dcp_to_torch_save - optimizer_state_dict = get_optimizer_state_dict(model, optimizer) - dcp.save({"optimizer": optimizer_state_dict}, checkpoint_id=checkpoint_dir) + dcp_to_torch_save(checkpoint_dir, os.path.join(checkpoint_dir, "optimizer.pt")) + _distributed_barrier() -def load_optimizer_distributed(model, optimizer, checkpoint_dir: str) -> None: - """Load optimizer state via DCP.""" +def load_optimizer_distributed(model, optimizer, checkpoint_dir_or_file: str) -> None: + """Load optimizer state from a DCP directory or a consolidated `optimizer.pt` file. + + Passing a directory uses the retained DCP shards. Passing the file loads the + full optimizer state on each rank's CPU before distributing it into the current + layout. Prefer the directory when memory is limited. All ranks must call. + """ if not is_torch_greater_or_equal("2.7"): raise OSError("Distributed checkpointing requires `torch>=2.7`.") # Import here because otherwise it emits a warning every time it's imported on some hardware - this keeps the warning from # 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, set_optimizer_state_dict + from torch.distributed.checkpoint.state_dict import ( + StateDictOptions, + get_optimizer_state_dict, + set_optimizer_state_dict, + ) + + options = StateDictOptions(flatten_optimizer_state_dict=True) + optimizer_state_dict = get_optimizer_state_dict(model, optimizer, options=options) + if os.path.isfile(checkpoint_dir_or_file): + from torch.distributed.tensor import distribute_tensor + + loaded_state = torch.load(checkpoint_dir_or_file, map_location="cpu", weights_only=True)["optimizer"] + missing_keys = optimizer_state_dict.keys() - loaded_state.keys() + if missing_keys: + raise ValueError(f"Missing keys in optimizer checkpoint: {sorted(missing_keys)}") + for key, target in optimizer_state_dict.items(): + value = loaded_state[key] + if isinstance(target, torch.Tensor) and ( + not isinstance(value, torch.Tensor) or value.shape != target.shape + ): + raise ValueError(f"Optimizer checkpoint tensor {key!r} must have shape {tuple(target.shape)}.") + for key, target in optimizer_state_dict.items(): + value = loaded_state[key] + if is_dtensor(target): + value = distribute_tensor(value.to(target.device), target.device_mesh, target.placements) + elif isinstance(target, torch.Tensor): + value = value.to(target.device) + optimizer_state_dict[key] = value + else: + from .checkpoint import HuggingFaceLoadPlanner - optimizer_state_dict = get_optimizer_state_dict(model, optimizer) - dcp.load({"optimizer": optimizer_state_dict}, checkpoint_id=checkpoint_dir) + dcp.load( + {"optimizer": optimizer_state_dict}, + checkpoint_id=checkpoint_dir_or_file, + planner=HuggingFaceLoadPlanner(), + ) set_optimizer_state_dict(model, optimizer, optimizer_state_dict) + + +def clip_grad_norm_(parameters, max_norm, norm_type=2.0, error_if_nonfinite=False, foreach=None): + """ + Equivalent to torch.nn.utils.clip_grad_norm_ but supports a mixture of ordinary and DTensors parameters. + """ + from torch.nn.utils import clip_grads_with_norm_, get_total_norm + + parameters = [parameters] if isinstance(parameters, torch.Tensor) else list(parameters) + norm_type = float(norm_type) + max_norm = float(max_norm) + params_by_mesh = defaultdict(list) + for param in parameters: + if param.grad is not None: + params_by_mesh[param.grad.device_mesh if is_dtensor(param.grad) else None].append(param) + + if len(params_by_mesh) <= 1 and max_norm != float("inf"): + total_norm = torch.nn.utils.clip_grad_norm_(parameters, max_norm, norm_type, error_if_nonfinite, foreach) + return total_norm.full_tensor() if is_dtensor(total_norm) else total_norm + + if not params_by_mesh: + return torch.tensor(0.0) + + norms = [] + for params in params_by_mesh.values(): + norm = get_total_norm([param.grad for param in params], norm_type, foreach=foreach) + norms.append(norm.full_tensor() if is_dtensor(norm) else norm) + stacked_norms = torch.stack([norm.to(norms[0].device) for norm in norms]) + + # For order zero, each group norm counts nonzero tensor norms, combine those counts by summing. + total_norm = stacked_norms.sum() if norm_type == 0 else torch.linalg.vector_norm(stacked_norms, norm_type) + + if error_if_nonfinite and torch.logical_or(total_norm.isnan(), total_norm.isinf()): + raise RuntimeError( + f"The total norm of order {norm_type} for gradients from " + "`parameters` is non-finite, so it cannot be clipped. To disable " + "this error and scale the gradients by the non-finite norm anyway, " + "set `error_if_nonfinite=False`" + ) + + if max_norm != float("inf"): + for params in params_by_mesh.values(): + clip_grads_with_norm_(params, max_norm, total_norm, foreach) + return total_norm diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index 275e83ef268e..8904dfa20e44 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -3247,6 +3247,7 @@ def save_pretrained( save_peft_format: bool = True, save_original_format: bool = True, distributed_checkpoint: bool = False, + consolidate_distributed_checkpoint: bool = True, **kwargs, ): """ @@ -3293,13 +3294,20 @@ def save_pretrained( its reverse mapping. The reverse mapping needs to exists even if the model was loaded from a None legacy checkpoint. distributed_checkpoint (`bool`, *optional*, defaults to `False`): - When saving an FSDP-wrapped model, use the distributed checkpoint (DCP) path instead of gathering weights - to CPU first. Every rank must call this method; rank 0 writes the consolidated Hugging Face safetensors. + When saving an FSDP-wrapped model, write safetensors with distributed checkpointing (DCP) instead of + gathering weights to CPU first. Every rank must call this method. When `False`, FSDP weights are gathered to CPU on rank 0 via `gather_full_state_dict` before writing. Native FSDP requires `torch>=2.7`. + consolidate_distributed_checkpoint (`bool`, *optional*, defaults to `True`): + Consolidate rank-local safetensors files into complete model weights loadable with `from_pretrained()`. + Intermediate files are retained under `sharded/`. When `False`, only rank-local safetensors files are + written in `save_directory`; load these with `load_distributed_checkpoint()`. kwargs (`dict[str, Any]`, *optional*): Additional key word arguments passed along to the [`~utils.PushToHubMixin.push_to_hub`] method. """ + if not distributed_checkpoint and not consolidate_distributed_checkpoint: + raise ValueError("Distributed checkpoint options require `distributed_checkpoint=True`.") + if token is not None: kwargs["token"] = token @@ -3406,6 +3414,7 @@ def save_pretrained( self.save_distributed_checkpoint( model_to_save, save_directory, + consolidate=consolidate_distributed_checkpoint, push_to_hub=push_to_hub, save_on_this_rank=save_on_this_rank, token=token, diff --git a/src/transformers/trainer.py b/src/transformers/trainer.py index a4c927f34c12..89e431b65c03 100755 --- a/src/transformers/trainer.py +++ b/src/transformers/trainer.py @@ -57,6 +57,7 @@ from .data.data_collator import DataCollator, DataCollatorWithPadding, default_data_collator from .debug_utils import DebugOption, DebugUnderflowOverflow from .distributed.fsdp import get_fsdp_ckpt_kwargs, update_fsdp_plugin_peft +from .distributed.utils import clip_grad_norm_ from .feature_extraction_sequence_utils import SequenceFeatureExtractor from .feature_extraction_utils import FeatureExtractionMixin from .hyperparameter_search import ALL_HYPERPARAMETER_SEARCH_BACKENDS, default_hp_search_backend @@ -2644,30 +2645,6 @@ def _track_num_input_tokens(self, inputs): input_tokens = torch.as_tensor(input_tokens, device=self.args.device, dtype=torch.int64) self.state.num_input_tokens_seen += self.accelerator.gather(input_tokens).sum().item() - def _mixed_mesh_grad_norm(self, model, max_norm): - """ - Gradient norm (and clip) when the gradients live on different device meshes, which `clip_grad_norm_` cannot - span: one norm per mesh, each already reduced over its own mesh. - """ - from torch.distributed.tensor import DTensor - from torch.nn.utils import clip_grads_with_norm_, get_total_norm - - params_by_mesh = defaultdict(list) - for param in model.parameters(): - if param.grad is not None: - params_by_mesh[param.grad.device_mesh if isinstance(param.grad, DTensor) else None].append(param) - - norms = [] - for params in params_by_mesh.values(): - norm = get_total_norm([p.grad for p in params]) - norms.append(norm.full_tensor() if isinstance(norm, DTensor) else norm) - total_norm = torch.linalg.vector_norm(torch.stack(norms)) - - if max_norm != float("inf"): - for params in params_by_mesh.values(): - clip_grads_with_norm_(params, max_norm, total_norm) - return total_norm - def _has_mixed_mesh_grads(self, model) -> bool: # Static for the life of the run (sharding never changes after setup), so scan the # parameters only on the first call. @@ -2680,7 +2657,7 @@ def _clip_grad_norm(self, model): if is_sagemaker_mp_enabled() and self.args.fp16: return self.optimizer.clip_master_grads(self.args.max_grad_norm) if self._has_mixed_mesh_grads(model): - return self._mixed_mesh_grad_norm(model, self.args.max_grad_norm) + return clip_grad_norm_(model.parameters(), self.args.max_grad_norm) return self.accelerator.clip_grad_norm_(model.parameters(), self.args.max_grad_norm) def _get_grad_norm(self, model, grad_norm=None): @@ -2688,7 +2665,7 @@ def _get_grad_norm(self, model, grad_norm=None): if grad_norm is None: # Compute norm without clipping (inf means no actual clipping happens) if self._has_mixed_mesh_grads(model): - grad_norm = self._mixed_mesh_grad_norm(model, float("inf")) + grad_norm = clip_grad_norm_(model.parameters(), float("inf")) else: grad_norm = self.accelerator.clip_grad_norm_(model.parameters(), float("inf")) diff --git a/tests/test_distributed_utils.py b/tests/test_distributed_utils.py new file mode 100644 index 000000000000..a6897564c597 --- /dev/null +++ b/tests/test_distributed_utils.py @@ -0,0 +1,212 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import os +import tempfile +import unittest +from contextlib import contextmanager +from datetime import timedelta +from unittest.mock import patch + +from transformers.testing_utils import require_torch +from transformers.utils import is_torch_available + + +if is_torch_available(): + import torch + import torch.distributed as dist + import torch.multiprocessing as mp + + from transformers import LlamaConfig, LlamaForCausalLM + from transformers.distributed import DistributedConfig + from transformers.distributed.utils import clip_grad_norm_, load_optimizer_distributed, save_optimizer_distributed + + if dist.is_available(): + from torch.distributed.device_mesh import init_device_mesh + from torch.distributed.tensor import DTensor, Shard, distribute_tensor + + +def _full_tensor(tensor): + return tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + + +def _optimizer(model): + # foreach requires homogeneous groups when TP leaves some parameters unsharded. + groups = [ + [p for p in model.parameters() if isinstance(p, DTensor) == distributed] for distributed in (False, True) + ] + return torch.optim.AdamW([{"params": group} for group in groups if group], lr=0.003, foreach=True) + + +def _step(model, optimizer): + for parameter in model.parameters(): + parameter.grad = parameter.detach().clone() + optimizer.step() + optimizer.zero_grad(set_to_none=True) + + +def _check_optimizer(model, optimizer, reference, reference_optimizer): + for parameter, expected in zip(model.parameters(), reference.parameters()): + actual_state = optimizer.state[parameter] + expected_state = reference_optimizer.state[expected] + assert actual_state.keys() == expected_state.keys() + for key in expected_state: + torch.testing.assert_close(_full_tensor(actual_state[key]), expected_state[key]) + assert all(group["lr"] == reference_optimizer.param_groups[0]["lr"] for group in optimizer.param_groups) + + +@contextmanager +def _distributed_context(rank, directory): + environment = {"RANK": str(rank), "LOCAL_RANK": str(rank), "WORLD_SIZE": "4"} + with patch.dict(os.environ, environment), patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + dist.init_process_group( + "gloo", + init_method=f"file://{directory}/rendezvous", + rank=rank, + world_size=4, + timeout=timedelta(seconds=60), + ) + try: + yield + finally: + dist.destroy_process_group() + + +def _load_model(directory, consolidate, config=None): + if consolidate: + return LlamaForCausalLM.from_pretrained(f"{directory}/saved", distributed_config=config) + + model = LlamaForCausalLM.from_pretrained(f"{directory}/seed", distributed_config=config) + # Clear the destination so unchanged seed weights cannot make the round trip pass. + with torch.no_grad(): + for parameter in model.parameters(): + parameter.zero_() + model.load_distributed_checkpoint(f"{directory}/saved") + return model + + +# Workers stay at module scope so multiprocessing.spawn can pickle them. +def _model_checkpoint_worker(rank, directory, consolidate): + with _distributed_context(rank, directory): + reference = LlamaForCausalLM.from_pretrained(f"{directory}/seed") + model = LlamaForCausalLM.from_pretrained( + f"{directory}/seed", distributed_config=DistributedConfig(tp_size=2, fsdp_size=2) + ) + model.save_pretrained( + f"{directory}/saved", + distributed_checkpoint=True, + consolidate_distributed_checkpoint=consolidate, + ) + for config in (DistributedConfig(tp_size=2, fsdp_size=2), DistributedConfig(tp_size=4)): + restored = _load_model(directory, consolidate, config) + for name, parameter in restored.state_dict().items(): + torch.testing.assert_close(_full_tensor(parameter), reference.state_dict()[name]) + + +def _optimizer_checkpoint_worker(rank, directory, consolidate): + with _distributed_context(rank, directory): + reference = LlamaForCausalLM.from_pretrained(f"{directory}/seed") + model = LlamaForCausalLM.from_pretrained( + f"{directory}/seed", distributed_config=DistributedConfig(tp_size=2, fsdp_size=2) + ) + optimizer, reference_optimizer = _optimizer(model), _optimizer(reference) + _step(model, optimizer) + _step(reference, reference_optimizer) + save_optimizer_distributed(model, optimizer, f"{directory}/saved", consolidate=consolidate) + checkpoint = f"{directory}/saved/optimizer.pt" if consolidate else f"{directory}/saved" + for config in (DistributedConfig(tp_size=2, fsdp_size=2), DistributedConfig(tp_size=4)): + restored = LlamaForCausalLM.from_pretrained(f"{directory}/seed", distributed_config=config) + restored_optimizer = _optimizer(restored) + load_optimizer_distributed(restored, restored_optimizer, checkpoint) + _check_optimizer(restored, restored_optimizer, reference, reference_optimizer) + + +def _gradient_clipping_worker(rank, directory): + with _distributed_context(rank, directory): + mesh = init_device_mesh("cpu", (4,)) + for distributed in ((False, False), (True, True), (False, True)): + for max_norm in (1.0, 10000.0): + parameters, reference = [], [] + for i, is_distributed in enumerate(distributed): + gradient = torch.arange(1, 65, dtype=torch.float32).reshape(8, 8) * (i + 1) + expected = torch.nn.Parameter(torch.zeros_like(gradient)) + expected.grad = gradient.clone() + reference.append(expected) + if is_distributed: + gradient = distribute_tensor(gradient, mesh, [Shard(0)]) + parameter = torch.nn.Parameter(torch.zeros_like(gradient)) + parameter.grad = gradient + parameters.append(parameter) + expected_norm = torch.nn.utils.clip_grad_norm_(reference, max_norm, foreach=True) + actual_norm = clip_grad_norm_(parameters, max_norm, foreach=True) + torch.testing.assert_close(_full_tensor(actual_norm), expected_norm) + for parameter, expected in zip(parameters, reference): + torch.testing.assert_close(_full_tensor(parameter.grad), expected.grad) + + +@require_torch +@unittest.skipUnless( + is_torch_available() and dist.is_available() and dist.is_gloo_available(), "Requires distributed Gloo" +) +class DistributedUtilsTest(unittest.TestCase): + def setUp(self): + self.config = LlamaConfig( + vocab_size=16, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=4, + num_key_value_heads=4, + ) + + def test_model_checkpoint(self): + for consolidate in (False, True): + with self.subTest(consolidate=consolidate): + with tempfile.TemporaryDirectory() as directory: + reference = LlamaForCausalLM(self.config) + reference.save_pretrained(f"{directory}/seed") + mp.spawn( + _model_checkpoint_worker, + args=(directory, consolidate), + nprocs=4, + join=True, + ) + + # Reload without a process group or distributed configuration. + restored = _load_model(directory, consolidate) + for name, parameter in restored.state_dict().items(): + torch.testing.assert_close(parameter, reference.state_dict()[name]) + + def test_optimizer_checkpoint(self): + for consolidate in (False, True): + with self.subTest(consolidate=consolidate), tempfile.TemporaryDirectory() as directory: + reference = LlamaForCausalLM(self.config) + reference.save_pretrained(f"{directory}/seed") + mp.spawn(_optimizer_checkpoint_worker, args=(directory, consolidate), nprocs=4, join=True) + + # Reload without a process group or distributed configuration. + reference_optimizer = _optimizer(reference) + _step(reference, reference_optimizer) + restored = LlamaForCausalLM.from_pretrained(f"{directory}/seed") + restored_optimizer = _optimizer(restored) + checkpoint = f"{directory}/saved/optimizer.pt" if consolidate else f"{directory}/saved" + load_optimizer_distributed(restored, restored_optimizer, checkpoint) + _check_optimizer(restored, restored_optimizer, reference, reference_optimizer) + + def test_gradient_clipping(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_gradient_clipping_worker, args=(directory,), nprocs=4, join=True) + + +if __name__ == "__main__": + unittest.main()