[core] Shard tensor-parallel checkpoints on load and save - #14544
JingyaHuang wants to merge 18 commits into
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.
| Pass a [`TensorParallelConfig`] to [`~ModelMixin.enable_parallelism`]. `tp_degree` is the number of devices to shard across and must divide the model's number of attention heads. The model must define a `_tp_plan` (a flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style). | ||
| Pass a [`TensorParallelConfig`] to the `parallel_config` argument of the model's [`~ModelMixin.from_pretrained`]. `tp_degree` is the number of devices to shard across and must divide the model's number of attention heads. The model must define a `_tp_plan` (a flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style). | ||
|
|
||
| Loading this way shards the checkpoint *while reading it*: each rank reads only its own slice of each sharded weight and places it straight onto its own device. Nothing full-size is ever materialized, so per-rank memory falls as `tp_degree` rises. |
There was a problem hiding this comment.
Oh really! This is very cool. Could we also present a small comparison between the loading time with and without this way of loading?
There was a problem hiding this comment.
Here's a quick benchmark:
Synthetic checkpoint (FLUX.1-shaped, 1.33B params, 2.65GB bf16):
| mode | tp_degree | load time | peak GPU/rank | peak CPU/rank |
|---|---|---|---|---|
streamed (from_pretrained(parallel_config=...)) |
2 | 1.92s ± 0.11 | 1.71GB | 2.70GB |
materialize+reshard (from_pretrained + enable_parallelism) |
2 | 3.06s ± 0.02 | 1.69GB | 4.08GB |
streamed (from_pretrained(parallel_config=...)) |
4 | 1.45s ± 0.03 | 1.23GB | 2.20GB |
materialize+reshard (from_pretrained + enable_parallelism) |
4 | 3.29s ± 0.02 | 1.23GB | 4.08GB |
- Speed: streaming is ~37% faster at tp_degree=2, ~56% faster at tp_degree=4 on the synthetic checkpoint (1.92s vs 3.06s, 1.45s vs 3.29s).
- GPU memory: basically the same between the two.
- CPU/host memory: 34% less at tp_degree=2, 46% less at tp_degree=4 (2.70GB vs 4.08GB, 2.20GB vs 4.08GB).
There was a problem hiding this comment.
Okay nice. So, it's the speed that matters for now. Let's include these numbers in the docs, then?
|
|
||
| `tp_degree` is taken from `world_size` above, so `--nproc-per-node 4` shards the transformer across 4 devices. | ||
|
|
||
| A tensor-parallel `parallel_config` cannot be combined with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. Tensor parallelism also cannot be combined with quantization, offloading, or LoRA adapters at all — the parameters it shards have to be plain parameters owned by the model — so those raise however the model is sharded. To shard a model that is already in memory, call [`~ModelMixin.enable_parallelism`] with the same config instead — that loads everything first and reshards it, so it costs full checkpoint memory on every rank. |
There was a problem hiding this comment.
Interesting that we cannot load TP with quantization. Do we know why?
There was a problem hiding this comment.
It's mostly about scoping... I didn't want to handle "quantization + TP" sharding in the scale of this PR. IMO, it could be possible if we shard and then quantize, and then according to the quantization configs and the backend, some might work, some doesn't... I would rather raise for now, we could consider supporting the combo properly if the community shows interest.
There was a problem hiding this comment.
Yeah raising is totally fine. Maybe we should put this blob of text into a "> [!CAUTION]" block?
| ### 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?
sayakpaul
left a comment
There was a problem hiding this comment.
Thanks for keeping the PR very scoped and easy to review. Left some feedback which should be straightforward to resolve.
Ran the tests on HF Jobs, too, and they pass: https://huggingface.co/jobs/sayakpaul/6aba42356b030d633f69c66d
| Pass a [`TensorParallelConfig`] to [`~ModelMixin.enable_parallelism`]. `tp_degree` is the number of devices to shard across and must divide the model's number of attention heads. The model must define a `_tp_plan` (a flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style). | ||
| Pass a [`TensorParallelConfig`] to the `parallel_config` argument of the model's [`~ModelMixin.from_pretrained`]. `tp_degree` is the number of devices to shard across and must divide the model's number of attention heads. The model must define a `_tp_plan` (a flat mapping of module-name globs to a `"colwise"`/`"rowwise"` style). | ||
|
|
||
| Loading this way shards the checkpoint *while reading it*: each rank reads only its own slice of each sharded weight and places it straight onto its own device. Nothing full-size is ever materialized, so per-rank memory falls as `tp_degree` rises. |
There was a problem hiding this comment.
Okay nice. So, it's the speed that matters for now. Let's include these numbers in the docs, then?
|
|
||
| `tp_degree` is taken from `world_size` above, so `--nproc-per-node 4` shards the transformer across 4 devices. | ||
|
|
||
| A tensor-parallel `parallel_config` cannot be combined with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. Tensor parallelism also cannot be combined with quantization, offloading, or LoRA adapters at all — the parameters it shards have to be plain parameters owned by the model — so those raise however the model is sharded. To shard a model that is already in memory, call [`~ModelMixin.enable_parallelism`] with the same config instead — that loads everything first and reshards it, so it costs full checkpoint memory on every rank. |
There was a problem hiding this comment.
Yeah raising is totally fine. Maybe we should put this blob of text into a "> [!CAUTION]" block?
|
|
||
| A tensor-parallel `parallel_config` cannot be combined with `device_map`, `quantization_config`, `low_cpu_mem_usage=False`, `use_flashpack=True`, or non-safetensors weights; each raises rather than quietly falling back to loading the full checkpoint. Tensor parallelism also cannot be combined with quantization, offloading, or LoRA adapters at all — the parameters it shards have to be plain parameters owned by the model — so those raise however the model is sharded. To shard a model that is already in memory, call [`~ModelMixin.enable_parallelism`] with the same config instead — that loads everything first and reshards it, so it costs full checkpoint memory on every rank. | ||
|
|
||
| Saving a tensor-parallel model isn't supported yet, and [`~ModelMixin.save_pretrained`] raises on one. Save the model before sharding it. |
| for block_size in block_sizes: | ||
| # An uneven split is rejected rather than silently handed to `Shard`, which pads the tail | ||
| # and would break the paired colwise/rowwise matmul. | ||
| if block_size % tp_size != 0: | ||
| raise ValueError( | ||
| f"Cannot shard a block of size {block_size} across {tp_size} tensor-parallel ranks: " | ||
| f"{block_size} is not divisible by {tp_size}." | ||
| ) |
There was a problem hiding this comment.
Let's raise this error earlier.
| blocks = _blocks if _blocks is not None else getattr(module, "_tp_packed_col_blocks") | ||
| rank = device_mesh.get_local_rank() | ||
| tp_size = device_mesh.size() | ||
| blocks = _blocks if _blocks is not None else module._tp_packed_col_blocks |
There was a problem hiding this comment.
Should we not raise if module._tp_packed_col_blocks doesn't have anything?
| non_safetensors = [f for f in resolved_model_file if not str(f).endswith(".safetensors")] | ||
| if non_safetensors: | ||
| raise ValueError( | ||
| f"A tensor-parallel `parallel_config` requires safetensors weights, so that each rank can " | ||
| f"read only its own slice of each tensor. Got {non_safetensors}." | ||
| ) |
There was a problem hiding this comment.
Is this raise sufficiently early in the stack?
| model.eval() | ||
|
|
||
| if parallel_config is not None: | ||
| if tp_shard_specs is not None: |
There was a problem hiding this comment.
I think it's better to keep the same conditionals depend on the same variable, readability wise. So, we could keep parallel_config and then condition on tp_shards_specs?
| return resolved | ||
|
|
||
|
|
||
| def _check_tp_supported(model_name: str, tp_plan: "dict | None", num_heads: "int | None", tp_config) -> None: |
There was a problem hiding this comment.
I feel like the validation methods are a bit scattered along different files. It makes sense to centralize them.
| ) | ||
|
|
||
| @classmethod | ||
| def _check_tp_streaming_supported( |
There was a problem hiding this comment.
This could probably go to check_tp_supported and check_tp_supported_model_state could be consolidated into it?
| return self.linear_2(self.linear_1(hidden_states)) | ||
|
|
||
|
|
||
| @pytest.fixture(scope="module") |
There was a problem hiding this comment.
We don't test these methods. Let's remove this test.
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.