[core] Shard tensor-parallel checkpoints on load and save - #14544
Conversation
|
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. |
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.
a5e135c to
acdf4bf
Compare
…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.
…sers into add-shard-ckpt-loading
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>
…sers into add-shard-ckpt-loading
sayakpaul
left a comment
There was a problem hiding this comment.
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.
| ### 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. |
There was a problem hiding this comment.
That is cool! However, do we have to ship this yet? I don't have any strong opinions. @DN6 WDYT?
|
Just validated the changes on cuda / neuron / tpu. |
sayakpaul
left a comment
There was a problem hiding this comment.
Just a few last comments and I think we're good after that!
| from torch.distributed.tensor import Replicate, distribute_tensor | ||
| from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel | ||
|
|
||
| def _make_packed_col(marker: PackedColwiseParallel) -> ColwiseParallel: |
sayakpaul
left a comment
There was a problem hiding this comment.
Thanks, last set of comment, one being major.
| | 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 | |
There was a problem hiding this comment.
Am I reading it right that it leads to a decrease of ~10x in the peak CPU memory per rank? 😳
There was a problem hiding this comment.
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): |
There was a problem hiding this comment.
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:
There was a problem hiding this comment.
I pushed some updates to in 95c44cc. Hope, that is okay with you :) I ran the tests and they passed.
| 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." |
There was a problem hiding this comment.
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?
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:from_pretrained(..., parallel_config=...)shards while reading, each rank slices only its own part of every_tp_planweight off disk, straight into a DTensor on its device. Unsupported combinations: device_map / quantization / use_flashpack / DDUF / non-safetensors -> raise.save_pretrained()gathers the shards back to a normal checkpointBesides above:
_check_tp_model_state, called fromapply_tensor_parallel, rejects a model that is quantized, group-offloaded, placed by accelerate (device_mapor CPU offload), or has PEFT layers injected.load_lora_adapterrefuses a tensor-parallel model.save_pretrainedrefuses a quantized TP model.Before submitting
self-reviewskill on the diff?documentation guidelines, and
here are tips on formatting docstrings.
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.