Skip to content

[core] Shard tensor-parallel checkpoints on load and save - #14544

Merged
sayakpaul merged 38 commits into
huggingface:mainfrom
JingyaHuang:add-shard-ckpt-loading
Oct 1, 2026
Merged

sayakpaul merged 38 commits into
huggingface:mainfrom
JingyaHuang:add-shard-ckpt-loading

Conversation

@JingyaHuang

@JingyaHuang JingyaHuang commented Aug 20, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do?

Fixes #14533

This is a follow-up of the Tensor Parallelism support in #13781, TP previously required loading the whole checkpoint on every rank and resharding it afterwards, so per-rank memory was the full model size. In this PR, we adapt the shard loading (.from_pretrained()) and saving (.save_pretrained()) to be tp-aware:

  • Load: from_pretrained(..., parallel_config=...) shards while reading, each rank slices only its own part of every _tp_plan weight off disk, straight into a DTensor on its device. Unsupported combinations: device_map / quantization / use_flashpack / DDUF / non-safetensors -> raise.
  • Save: save_pretrained() gathers the shards back to a normal checkpoint
    • collective: all on all ranks, rank 0 writes
    • save_pretrained(..., dcp=True) writes per-rank
import torch
import torch.distributed as dist

from diffusers import Flux2Transformer2DModel, TensorParallelConfig

dist.init_process_group()                                  # "nccl" on CUDA, "neuron" on Trainium
tp = TensorParallelConfig(tp_degree=dist.get_world_size())

# LOAD: shard while reading, never materialize the full checkpoint ----
model = Flux2Transformer2DModel.from_pretrained(
    "black-forest-labs/FLUX.2-klein-9B",
    subfolder="transformer",
    torch_dtype=torch.bfloat16,
    parallel_config=tp,
)

# SAVE: gather back
model.save_pretrained("out/full")

# SAVE (b): sharded distributed checkpoint, nothing gathered
# model.save_pretrained("out/sharded", dcp=True)

# LOAD BACK a distributed ckpt
# model = Flux2Transformer2DModel.from_pretrained("out/sharded", parallel_config=tp)

Besides above:

  • _check_tp_model_state, called from apply_tensor_parallel, rejects a model that is quantized, group-offloaded, placed by accelerate (device_map or CPU offload), or has PEFT layers injected.
  • load_lora_adapter refuses a tensor-parallel model.
  • save_pretrained refuses a quantized TP model.

Before submitting

  • Did you use an AI agent (Claude Code, Codex, Cursor, etc.) to help with this PR? If so:
    • Did you read the Coding with AI agents guide?
    • Did you run the self-review skill on the diff?
    • Did you share the final self-review notes in the PR description or a comment?
  • Did you read the contributor guideline?
  • Did you read our philosophy doc? (important for complex PRs)
  • Was this discussed/approved via a GitHub issue or the forum? Please add a link to it if that's the case.
  • Did you make sure to update the documentation with your changes? Here are the
    documentation guidelines, and
    here are tips on formatting docstrings.
  • Did you write any new necessary tests?
  • Are you the author (or part of the team) of the model/pipeline (only applicable for model/pipeline related PRs)?

Who can review?

Anyone in the community is free to review the PR once the tests have passed. Feel free to tag
members/contributors who may be interested in your PR.

@github-actions github-actions Bot added documentation Improvements or additions to documentation fixes-issue size/M PR with diff < 200 LOC labels Aug 20, 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.

@github-actions github-actions Bot added models tests utils hooks size/L PR with diff > 200 LOC and removed size/M PR with diff < 200 LOC labels Aug 20, 2026
Stream each rank's slice of a tensor-parallel checkpoint straight off disk
instead of materializing the full checkpoint on every rank and resharding it
afterwards, and gather the shards back on save.

- `from_pretrained(..., parallel_config=TensorParallelConfig(...))` resolves
  the shard specs on the still-meta model, then slices each safetensors tensor
  before the dtype cast, so host memory peaks at ~1/tp_degree of the checkpoint.
- `save_pretrained` all-gathers the DTensors into an ordinary checkpoint, or
  writes a distributed checkpoint with `dcp=True` so no full tensor is ever
  formed. The writing `tp_degree` is recorded, since a packed weight's stored
  layout is interleaved by it.
- Factor the plan interpretation out of the Neuron pre-shard path into shared
  `TPShardSpec` / `resolve_tp_shard_specs` / `_local_shard` / `_hooks_only_styles`
  helpers, so both backends and both the load and save paths shard identically.
@JingyaHuang
JingyaHuang force-pushed the add-shard-ckpt-loading branch from a5e135c to acdf4bf Compare August 20, 2026 16:09
JingyaHuang and others added 3 commits August 20, 2026 18:11
…ng or LoRA

Addresses the remaining two items of the review on huggingface#13718: tensor parallelism was rejected
alongside quantization and `device_map` only on the `from_pretrained` streaming path, while
`enable_parallelism` — which the quantization error message itself recommended — accepted a
quantized, offloaded or adapter-injected model and sharded it anyway.

- Add `_check_tp_model_state`, called from `apply_tensor_parallel`, the one chokepoint every TP
  entry point funnels through. It rejects a model that is quantized, group-offloaded, placed by
  accelerate (`device_map` or CPU offload), or has PEFT layers injected. Placed before the
  device-type check so the reported reason is the useful one.
- Guard the reverse order too: `enable_group_offload`, the two pipeline CPU-offload methods, and
  `load_lora_adapter` now refuse a tensor-parallel model.
- `save_pretrained` refuses a quantized tensor-parallel model. Previously the `dcp=True` branch
  returned before the quantizer's serialization step, writing shards with no quantization
  metadata and no error.
- The DCP load guard checked the `quantization_config` kwarg only, so a pre-quantized checkpoint
  directory loaded silently; check the config's own entry too, and add the missing `_tp_plan`
  check that otherwise surfaced as a raw `AttributeError`.
- Correct the `from_pretrained` message and the doc sentence that pointed at `enable_parallelism`
  as a way to shard a quantized model.

The new tests are the first tensor-parallel tests that need neither an accelerator nor more than
one rank: every case asserts a raise before any collective, so they run single-process on gloo.
@JingyaHuang
JingyaHuang marked this pull request as ready for review August 26, 2026 16:19
JingyaHuang and others added 4 commits August 26, 2026 18:19
Resolve conflict in `_load_pretrained_model`: keep the tensor-parallel
`load_fn` branch from this PR, and drop `dduf_entries` from the ordinary
branch — DDUF loading was removed upstream. The dangling `dduf_entries`
references this PR added (the DCP unsupported-options list and the
`_check_tp_streaming_supported` guard) go away with it, since the kwarg
no longer exists.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Resolve conflict in `_load_pretrained_model`: keep the tensor-parallel
`load_fn` branch from this PR, and drop `dduf_entries` from the ordinary
branch — DDUF loading was removed upstream. The dangling `dduf_entries`
references this PR added (the DCP unsupported-options list and the
`_check_tp_streaming_supported` guard) go away with it, since the kwarg
no longer exists.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

@sayakpaul sayakpaul left a comment

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.

Left some high-level comments.

My main comment is if we want to ship the advanced features of rank aware save and load yet. I am leaning towards raising when we encouter those situations and simplify the code a bit. This way, we can see if the community wants this feature and ship it when we have enough interest. But I would like to double-check with @DN6 on this too.

Comment thread docs/source/en/training/distributed_inference.md
Comment thread docs/source/en/training/distributed_inference.md Outdated
Comment on lines +488 to +490
### Saving a tensor-parallel model

[`~ModelMixin.save_pretrained`] gathers the shards back into ordinary full tensors, so the result is a normal checkpoint that loads with or without tensor parallelism. Gathering is a collective, so call it on **every** rank; only rank 0 writes.

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.

That is cool! However, do we have to ship this yet? I don't have any strong opinions. @DN6 WDYT?

Comment thread docs/source/en/training/distributed_inference.md Outdated
Comment thread src/diffusers/hooks/tensor_parallel.py Outdated
Comment thread src/diffusers/hooks/tensor_parallel.py Outdated
@sayakpaul
sayakpaul requested a review from DN6 September 2, 2026 11:36
Comment thread src/diffusers/hooks/tensor_parallel.py Outdated
Comment thread src/diffusers/models/model_loading_utils.py
@JingyaHuang

Copy link
Copy Markdown
Contributor Author

Just validated the changes on cuda / neuron / tpu.

@sayakpaul sayakpaul left a comment

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.

Just a few last comments and I think we're good after that!

Comment thread docs/source/en/training/distributed_inference.md Outdated
from torch.distributed.tensor import Replicate, distribute_tensor
from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel

def _make_packed_col(marker: PackedColwiseParallel) -> ColwiseParallel:

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 like the separation here.

Comment thread src/diffusers/pipelines/pipeline_utils.py Outdated
Comment thread src/diffusers/hooks/tensor_parallel.py
Comment thread src/diffusers/hooks/tensor_parallel.py

@sayakpaul sayakpaul left a comment

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.

Thanks, last set of comment, one being major.

Comment on lines +445 to +448
| tp_degree | method | load time | peak CPU/rank |
|---|---|---|---|
| 4 | `from_pretrained(parallel_config=...)` | 12.5s | 6.8GB |
| 4 | `from_pretrained` + `enable_parallelism` | 30.4s | 64.1GB |

@sayakpaul sayakpaul Oct 1, 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.

Am I reading it right that it leads to a decrease of ~10x in the peak CPU memory per rank? 😳

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.

Just reproduced it myself. Wow, this is very cool!

def test_tensor_parallel_batch_inputs(self):
self.test_tensor_parallel_inference(batch_size=2)

def _tp_checkpoint_and_reference(self, tmp_path, world_size):

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.

Hmm this is good but this doesn't ensure that the model checkpoint is saved with sharded state dict files, though. This is something that reflects real-world use cases.

We usually test tiny models so the sharding is not obvious. To do it explicitly, we need to specify the shard_size argument:

def test_sharded_checkpoints(self, base_model_output, tmp_path, atol=1e-5, rtol=0):

@sayakpaul sayakpaul Oct 1, 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.

I pushed some updates to in 95c44cc. Hope, that is okay with you :) I ran the tests and they passed.

Comment on lines +331 to +332
sharded = {k: v for k, v in model.state_dict().items() if isinstance(v, DTensor)}
assert sharded, "No parameter was sharded into a DTensor by the streaming load."

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.

Can be done in a follow-up but can we go even specific to assert DTensor on the keys we expect to be sharded based on the _tp_plan?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Sure, will improve the test in #14039

@sayakpaul
sayakpaul merged commit 863092f into huggingface:main Oct 1, 2026
40 of 41 checks passed
@sayakpaul sayakpaul added the performance Anything related to performance improvements, profiling and benchmarking label Oct 1, 2026
@JingyaHuang
JingyaHuang deleted the add-shard-ckpt-loading branch October 1, 2026 13:08
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation fixes-issue hooks lora models performance Anything related to performance improvements, profiling and benchmarking pipelines size/L PR with diff > 200 LOC tests

Projects

Status: Done

Development

Successfully merging this pull request may close these issues.

better reporting of errors when using TP

4 participants