Skip to content

FSDP x TP Saving / Loading, Planner / Writer based solution - #48830

Closed
michaelbenayoun wants to merge 26 commits into
huggingface:mainfrom
michaelbenayoun:fsdp_tp_saving_planner_solution
Closed

michaelbenayoun wants to merge 26 commits into
huggingface:mainfrom
michaelbenayoun:fsdp_tp_saving_planner_solution

Conversation

@michaelbenayoun

@michaelbenayoun michaelbenayoun commented Sep 15, 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

Provides the same functionality as #48802, while avoiding the _StridedShard -> Shard redistribution used for checkpointing.

It extends DCP with:

  • HuggingFaceSavePlanner: describes each _StridedShard as non-overlapping global regions and resolves each write request to the corresponding local tensor view.

  • HuggingFaceStorageWriter: stores chunks under unique physical keys, preserving their original parameter names and global offsets in metadata. This prevents multiple chunks of the same parameter from overwriting each other within one safetensors file.

  • HuggingFaceLoadPlanner and HuggingFaceStorageReader: map saved chunks into the destination tensor layout, supporting the same mesh, a different mesh, or ordinary tensors.

Consolidation produces standard Hugging Face checkpoints.

2. Optimizer checkpoints

Optimizer checkpoints use the same workaround, then restore the original placements after loading. Optimizer state is flattened by parameter name so it can be restored when parameter groups change across FSDP × TP → TP.
DTensors and ordinary tensors use separate optimizer groups, allowing foreach=True for both.

3. clip_grad_norm_ alignment

There were two independent implementation for clip_grad_norm_ to work with a mixture of ordinary and distributed tensors:

Current implementation is more aligned with the general torch.nn.utils.clip_grad_norm_ as it supports multiple norm types and tries to get the best of the 2 implementations.

A PR to align with the current implementation has been opened in Accelerate as well: huggingface/accelerate#4266

4. Distributed checkpointing / Consolidation

Load and save distributed and consolidated checkpoints for both the model and the optimizer.

Script

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

warning: No `requires-python` value found in the workspace. Defaulting to `>=3.13`.
W0916 18:49:36.663000 447125 torch/distributed/run.py:874]
W0916 18:49:36.663000 447125 torch/distributed/run.py:874] *****************************************
W0916 18:49:36.663000 447125 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:49:36.663000 447125 torch/distributed/run.py:874] *****************************************
Loading weights: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 134.68it/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.74s/it]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5614.19it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5543.20it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5481.65it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 5320.41it/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
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1443.70it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1755.45it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1684.86it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1680.29it/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.99s/it]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1773.71it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1670.99it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1587.16it/s]
Loading weights: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 111/111 [00:00<00:00, 1597.87it/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.'

@michaelbenayoun michaelbenayoun changed the title Fsdp tp saving planner solution FSDP x TP Saving / Loading, Planner / Writer based solution Sep 15, 2026
@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.

Comment thread tp_fsdp_overfit.py Outdated
)
# 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())

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.

what is the motivation to have our own writer ? The following works out of the box already iirc

from torch.distributed.checkpoint.hf_storage import HuggingFaceStorageWriter

 dcp.save(
        state_dict,
        storage_writer=HuggingFaceStorageWriter(
            path=checkpoint_dir,
            save_distributed=True,
            enable_consolidation=True,
        ),
    ) 

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 core issue for saving a DTensor with a _StridedShard placement is that the DTensor's default checkpoint hooks create incorrect WriteItem. Check the docstring here to understand what is happening.

So our _CheckpointView creates multiple WriteItems associated with actual contiguous regions in the global tensor.

When the time comes to actual write data the HuggingFaceStorageWriter from PyTorch uses the name of the parameter as key, which in the case of _StridedShard is not unique anymore: we have multiple chunks of a global tensor referring to the same parameter name.

The proposed writer uses unique keys instead: chunk_{idx}, and stores a mapping chunk_{idx} -> parameter name to be able to resolve who is who at loading time.

@github-actions

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 35221920081:1
Result: success | Jobs: 16 | Tests: 190,485 | Failures: 0 | Duration: 17h 21m

@michaelbenayoun

Copy link
Copy Markdown
Member Author

Closing this PR in favor of #48802 which offers the same features.
This PR provides a cleaner but too complex solution for the _StridedShard issue.

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.

3 participants