Skip to content

FSDP x TP Saving / Loading - #48802

Open
michaelbenayoun wants to merge 47 commits into
huggingface:mainfrom
michaelbenayoun:fsdp_tp_saving
Open

michaelbenayoun wants to merge 47 commits into
huggingface:mainfrom
michaelbenayoun:fsdp_tp_saving

Conversation

@michaelbenayoun

@michaelbenayoun michaelbenayoun commented Sep 14, 2026 •

Copy link
Copy Markdown
Member

CPU CI GPU run-slow

What does this PR do?

This PR handles a few thing related to distributed training.

1. _StridedShard placement issue

The _StridedShard placement 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

"""
torchrun --nproc_per_node=4 tp_fsdp_overfit.py

The script overfit one sentence following the steps:
    - Train first half using FSDP=2+TP=2
    - Save the model and optimizer in distributed checkpoint
    - Reload the model and optimizer from the distributed checkpoint
    - Train the rest in TP=4 (change distributed config)
    - Save the model and optimizer in distributed checkpoint
    - Reload the model in a single safetensors file.
    - Do inference in TP=4  and assert greedy generation reproduces the sentence verbatim
"""

import os

import torch
from torch.distributed.tensor import DTensor

from transformers import AutoModelForCausalLM, AutoTokenizer
from transformers.distributed import DistributedConfig
from transformers.distributed.utils import (
    clip_grad_norm_,
    load_optimizer_distributed,
    save_optimizer_distributed,
)


def create_optimizer(model):
    dtensor_params, tensor_params = [], []
    for parameter in model.parameters():
        (dtensor_params if isinstance(parameter, DTensor) else tensor_params).append(parameter)
    groups = []
    if dtensor_params:
        groups.append({"params": dtensor_params})
    if tensor_params:
        groups.append({"params": tensor_params})
    return torch.optim.AdamW(groups, lr=1e-3, foreach=True)


NAME = "Isotonic/TinyMixtral-4x248M-MoE"
TEXT = "In a quiet village nestled between rolling hills and a slow river, the autumn mornings arrived with mist that hung low over the fields and a sky that turned from grey to pale gold as the sun climbed."
CKPT = "./checkpoints"
PADDED_MODEL = "./padded_model"
OPT = os.path.join(CKPT, "optimizer")
STEPS = 10
HALF = STEPS // 2

rank, local_rank = int(os.environ["RANK"]), int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
torch.distributed.init_process_group("nccl", device_id=torch.device(f"cuda:{local_rank}"))

tokenizer = AutoTokenizer.from_pretrained(NAME)
ids = tokenizer(TEXT, return_tensors="pt").input_ids.to(f"cuda:{local_rank}")

# Pad before applying TP: both TP=2 and TP=4 need a divisible lm_head size.
if rank == 0:
    model = AutoModelForCausalLM.from_pretrained(NAME, dtype=torch.bfloat16, device_map="cpu")
    model.resize_token_embeddings(pad_to_multiple_of=4)
    model.save_pretrained(PADDED_MODEL)
    del model
torch.distributed.barrier()

# Train first half, distributed-save model + optimizer.
model = AutoModelForCausalLM.from_pretrained(
    PADDED_MODEL,
    distributed_config=DistributedConfig(tp_size=2, fsdp_size=2, enable_sequence_parallel=True),
    dtype=torch.bfloat16,
)
optimizer = create_optimizer(model)
model.train()
for step in range(0, HALF):
    loss = model(ids, labels=ids).loss
    loss.backward()
    total_norm = clip_grad_norm_(model.parameters(), max_norm=1.0)
    optimizer.step()
    optimizer.zero_grad()
    if rank == 0:
        print(f"step {step:>2} | loss {loss.item():.5f} grad norm {total_norm.item():.5f}")

model.save_pretrained(CKPT, distributed_checkpoint=True)
save_optimizer_distributed(model, optimizer, OPT)
del model, optimizer
torch.cuda.empty_cache()

# Reload model + optimizer from the distributed checkpoint, train the rest.

model = AutoModelForCausalLM.from_pretrained(
    CKPT,
    distributed_config=DistributedConfig(tp_size=4, enable_sequence_parallel=True),
    dtype=torch.bfloat16,
)
optimizer = create_optimizer(model)
load_optimizer_distributed(model, optimizer, OPT)

model.train()
for step in range(HALF, STEPS):
    loss = model(ids, labels=ids).loss
    loss.backward()
    total_norm = clip_grad_norm_(model.parameters(), max_norm=1.0)
    optimizer.step()
    optimizer.zero_grad()
    if rank == 0:
        print(f"step {step:>2} | loss {loss.item():.5f} grad norm {total_norm.item():.5f}")

# INFERENCE in TP=4
model.save_pretrained(CKPT)
save_optimizer_distributed(model, optimizer, OPT + "_tp4")
del model, optimizer
torch.cuda.empty_cache()

model = AutoModelForCausalLM.from_pretrained(
    CKPT,
    distributed_config=DistributedConfig(tp_size=4),
    dtype=torch.bfloat16,
)
model.eval()
prompt = tokenizer("In a quiet village", return_tensors="pt").to(f"cuda:{local_rank}")
out = model.generate(**prompt, max_new_tokens=ids.shape[-1] - prompt.input_ids.shape[-1], do_sample=False)

got, want = out[0].tolist(), ids[0].tolist()
if rank == 0:
    print(f"generated: {tokenizer.decode(got, skip_special_tokens=True)!r}")
    print(f"expected: {tokenizer.decode(want, skip_special_tokens=True)!r}")
assert got == want, (
    f"generation mismatch at index {next((i for i, (g, e) in enumerate(zip(got, want)) if g != e), -1)}"
)

torch.distributed.destroy_process_group()

And the output is:

warning: No `requires-python` value found in the workspace. Defaulting to `>=3.13`.
W0916 18:52:21.262000 454691 torch/distributed/run.py:874]
W0916 18:52:21.262000 454691 torch/distributed/run.py:874] *****************************************
W0916 18:52:21.262000 454691 torch/distributed/run.py:874] Setting OMP_NUM_THREADS environment variable for each process to be 1 in default, to avoid your system being overloaded, please further tune the variable for optimal performance in your application as needed.
W0916 18:52:21.262000 454691 torch/distributed/run.py:874] *****************************************
Loading weights: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 135.09it/s]
[transformers] The new embeddings will be initialized from a multivariate normal distribution that has old embeddings' mean and covariance. As described in this article: https://nlp.stanford.edu/~johnhew/vocab-expansion.html. To disable this, use `mean_resizing=False`
[transformers] The new lm_head weights will be initialized from a multivariate normal distribution that has old embeddings' mean and covariance. As described in this article: https://nlp.stanford.edu/~johnhew/vocab-expansion.html. To disable this, use `mean_resizing=False`
Writing model shards: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 1/1 [00:02<00:00,  2.63s/it]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5646.87it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5430.94it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5556.63it/s]

/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralRMSNorm, FSDPLinear) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralRMSNorm, FSDPLinear) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralRMSNorm, FSDPLinear) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralRMSNorm, FSDPLinear) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralForCausalLM) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralForCausalLM) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralForCausalLM) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
/home/michael/transformers/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_state.py:331: UserWarning: FSDP2-wrapped module (FSDPMixtralForCausalLM) returned a view tensor. An in-place op on this view (e.g., `x += y`) will silently drop the pre-backward hook and skip the all-gather, which can cause backward to fail or produce wrong gradients. Use out-of-place ops (`out = out + y`, not `out += y`) or `.clone()` the output before any in-place op.
  output = self._register_pre_backward_hook(output)
step  0 | loss 4.97490 grad norm 10.75000
step  1 | loss 1.45883 grad norm 5.71875
step  2 | loss 0.23285 grad norm 2.67188
step  3 | loss 0.21926 grad norm 3.20312
step  4 | loss 0.38996 grad norm 10.81250
[rank3]:W0916 18:52:40.667000 454786 torch/distributed/tensor/_redistribute.py:390] While redistributing from (_StridedShard(dim=0, sf=2), Shard(dim=0)) to (Shard(dim=0), Shard(dim=0)), 2 sequential all_gather operations will be performed. This is suboptimal: multiple collective operations have higher latency (separate kernel launches and synchronization points) and may give inconsistent results between ranks due to different reduction orders. it is not possible to merge non-ascending order all_gather operations.
[rank2]:W0916 18:52:40.668000 454785 torch/distributed/tensor/_redistribute.py:390] While redistributing from (_StridedShard(dim=0, sf=2), Shard(dim=0)) to (Shard(dim=0), Shard(dim=0)), 2 sequential all_gather operations will be performed. This is suboptimal: multiple collective operations have higher latency (separate kernel launches and synchronization points) and may give inconsistent results between ranks due to different reduction orders. it is not possible to merge non-ascending order all_gather operations.
[rank1]:W0916 18:52:40.669000 454784 torch/distributed/tensor/_redistribute.py:390] While redistributing from (_StridedShard(dim=0, sf=2), Shard(dim=0)) to (Shard(dim=0), Shard(dim=0)), 2 sequential all_gather operations will be performed. This is suboptimal: multiple collective operations have higher latency (separate kernel launches and synchronization points) and may give inconsistent results between ranks due to different reduction orders. it is not possible to merge non-ascending order all_gather operations.
[rank0]:W0916 18:52:40.674000 454783 torch/distributed/tensor/_redistribute.py:390] While redistributing from (_StridedShard(dim=0, sf=2), Shard(dim=0)) to (Shard(dim=0), Shard(dim=0)), 2 sequential all_gather operations will be performed. This is suboptimal: multiple collective operations have higher latency (separate kernel launches and synchronization points) and may give inconsistent results between ranks due to different reduction orders. it is not possible to merge non-ascending order all_gather operations.
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1784.55it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1709.28it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1762.23it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1561.45it/s]
step  5 | loss 0.00217 grad norm 0.06689
step  6 | loss 0.00039 grad norm 0.01117
step  7 | loss 0.00025 grad norm 0.00674
step  8 | loss 0.00023 grad norm 0.00922
step  9 | loss 0.00008 grad norm 0.00201
Writing model shards: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 1/1 [00:02<00:00,  2.96s/it]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1942.67it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1496.11it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1503.06it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1467.24it/s]
generated: 'In a quiet village nestled between rolling hills and a slow river, the autumn mornings arrived with mist that hung low over the fields and a sky that turned from grey to pale gold as the sun climbed.'
expected: 'In a quiet village nestled between rolling hills and a slow river, the autumn mornings arrived with mist that hung low over the fields and a sky that turned from grey to pale gold as the sun climbed.'

Comment thread tp_fsdp_overfit.py Outdated
@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

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.

@michaelbenayoun michaelbenayoun changed the title Fsdp tp saving FSDP x TP Saving / Loading Sep 15, 2026
Comment thread src/transformers/distributed/mixin.py Outdated
Comment thread src/transformers/distributed/utils.py Outdated
Comment thread src/transformers/distributed/utils.py Outdated
Comment thread src/transformers/distributed/utils.py
Comment thread src/transformers/distributed/utils.py Outdated
Comment thread src/transformers/distributed/utils.py Outdated
Comment thread src/transformers/distributed/utils.py Outdated
# 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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same import top + guarding distributed as well. Make sure as well to check torch version availabilities for the func

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

But the comment at the top mentions that it is imported here on purpose to avoid warnings.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

okay then add guarding. We have to make sure the user downloaded a torch version that has DCP

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The guard is already there, line 514.

Comment thread src/transformers/distributed/utils.py Outdated
Comment thread src/transformers/distributed/utils.py Outdated
Comment thread src/transformers/distributed/utils.py Outdated

@michaelbenayoun michaelbenayoun left a comment

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I've addressed your comments. WDYT?

Comment thread src/transformers/distributed/checkpoint.py Outdated
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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

moving that in mixin ?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

moving that in mixin ?

@3outeille 3outeille Oct 6, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Comment thread src/transformers/distributed/checkpoint.py Outdated
_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):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why is there 2 paths when checkpoint is consolidate ?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@ArthurZucker

Copy link
Copy Markdown
Collaborator

Sorry a bit late, reviewing now!

import torch


def _check_distributed_checkpointing_available(raise_if_not: bool = True) -> bool:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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():

@3outeille 3outeille Oct 6, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

use is_torch_distributed_available() as it uses across the codebase

@github-actions

github-actions Bot commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 37397912686:1
Result: success | Jobs: 16 | Tests: 193,976 | Failures: 0 | Duration: 13h 48m

token=token,
**hub_kwargs,
# Native DCP checkpoints remain sharded, use `distributed_checkpoint=False` for an interoperable checkpoint.
consolidate=False,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we should have an utils function called consolidate_distributed_checkpoint() in transformers.distributed.checkpoint



def _prepare_state_dict_for_dcp(state_dict):
"""

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

@3outeille 3outeille Oct 6, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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")):

@3outeille 3outeille Oct 6, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

@3outeille

Copy link
Copy Markdown
Member

We are getting there !

Main points:

  • consolidate_distributed_checkpoint()in transformers.distributed.checkpoint
  • up-to-date script for PR description
      1. train + save dcp
      1. resume from the DCP checkpoint (same process group), train, save dcp.
      1. consolidate with consolidate_distributed_checkpoint()
      1. generate in different distributed layout than training
  • better explanation of doctstring
  • move things to mixin
  • remove potential code paths that are not triggered anymore

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants