Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
26 commits
Select commit Hold shift + click to select a range
4454e4b
wip: add script to track potential changes
michaelbenayoun Sep 14, 2026
ec979c4
feat: workaround for _StridedShard
michaelbenayoun Sep 14, 2026
6d9a239
feat: handle optimizer loading / saving
michaelbenayoun Sep 14, 2026
6868815
wip: add clip_grad_norm_
michaelbenayoun Sep 14, 2026
f20be8d
doc: add docstring to clip_grad_norm_
michaelbenayoun Sep 14, 2026
fec373e
Merge branch 'main' into fsdp_tp_saving
michaelbenayoun Sep 14, 2026
f9341a8
feat: uniformize the clip_grad_norm_ implem
michaelbenayoun Sep 15, 2026
b09c7ee
refactor: use the proper function for clip_grad_norm_
michaelbenayoun Sep 15, 2026
9f677ef
feat: save shards / consolidate utility functions
michaelbenayoun Sep 15, 2026
9c742be
feat: save / load sharded / consolidated optimizer states
michaelbenayoun Sep 15, 2026
dc192d9
feat: add checkpoint state check
michaelbenayoun Sep 15, 2026
24558d8
Merge branch 'main' into fsdp_tp_saving
michaelbenayoun Sep 15, 2026
b271952
feat: add distributed checkpoint loading
michaelbenayoun Sep 15, 2026
6a5075c
test: add saving, loading and clipping tests
michaelbenayoun Sep 15, 2026
4288507
feat: writer / reader / planer based solution
michaelbenayoun Sep 15, 2026
78ece67
wip: remove .bin support
michaelbenayoun Sep 16, 2026
67f41e5
doc: update docstring
michaelbenayoun Sep 16, 2026
931d76f
chore: remove testing script
michaelbenayoun Sep 16, 2026
a4dc979
Merge branch 'main' into fsdp_tp_saving_planner_solution
michaelbenayoun Sep 16, 2026
ab8df86
doc: improve docstring for _CheckpointView
michaelbenayoun Sep 16, 2026
0309cb1
doc: add comment about regions
michaelbenayoun Sep 16, 2026
5840a2e
Merge branch 'main' into fsdp_tp_saving_planner_solution
michaelbenayoun Sep 16, 2026
60862c9
fix: remove remaining torch saving implem
michaelbenayoun Sep 17, 2026
ebab4c4
doc: improve docstring
michaelbenayoun Sep 17, 2026
5b0aa3b
Merge branch 'main' into fsdp_tp_saving_planner_solution
michaelbenayoun Sep 18, 2026
b4f8baa
fix: drop redundant optimizer state dict alias
michaelbenayoun Sep 18, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
199 changes: 199 additions & 0 deletions src/transformers/distributed/checkpoint.py
Original file line number Diff line number Diff line change
@@ -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)
172 changes: 172 additions & 0 deletions src/transformers/distributed/hf_storage.py
Original file line number Diff line number Diff line change
@@ -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)
Loading
Loading