Buckets:

hf-doc-build/doc-dev / transformers /pr_48314 /en /tensor_parallelism.md
|
download
raw
4.24 kB
# Tensor parallelism for training
Tensor parallelism (TP) splits weight matrices column-wise or row-wise across GPUs. Each GPU holds a shard, computes a partial result, and synchronizes with an all-reduce to produce the full output.
TP relies on frequent cross-GPU communication. It works best on hardware with fast intra-node links such as NVLink.
```text
┌─────────────────────────────┐
│ X (replicated) │
└────┬──────────┬─────────┬───┘
│ │ │
┌────▼───┐ ┌────▼───┐ ┌───▼────┐
│ ▓▓▓ W₀ │ │ ░░░ W₁ │ │ ███ W₂ │
│ X@W₀ │ │ X@W₁ │ │ X@W₂ │
└────┬───┘ └────┬───┘ └───┬────┘
└──────────┼─────────┘
Y₀+Y₁+Y₂
┌────────────────────────────┐
│ Y (full) │
└────────────────────────────┘
```
Transformers supports TP for architectures whose config defines `base_model_tp_plan`. Check that field first to see whether a model supports native TP.
```py
from transformers import AutoConfig
config = AutoConfig.from_pretrained("Qwen/Qwen3-0.6B")
print(config.base_model_tp_plan is not None)
print(config.base_model_tp_plan)
```
If a model supports TP, create a [DistributedConfig](/docs/transformers/pr_48314/en/expert_parallelism#transformers.DistributedConfig) with the number of devices in `tp_size` and pass it to [from_pretrained()](/docs/transformers/pr_48314/en/main_classes/model#transformers.PreTrainedModel.from_pretrained). Transformers uses the model's predefined plan, initializes the device mesh, and shards the supported layers for you.
You can also set `tp_plan="auto"` in [DistributedConfig](/docs/transformers/pr_48314/en/expert_parallelism#transformers.DistributedConfig). When `tp_size` is omitted, it is inferred from `WORLD_SIZE`. Passing `tp_plan` directly to [from_pretrained()](/docs/transformers/pr_48314/en/main_classes/model#transformers.PreTrainedModel.from_pretrained) is deprecated and will be removed in v5.18.
> [!WARNING]
> Don't use `device_map` with `distributed_config`. The two conflict at the weight-loading level. `device_map` places whole modules on specific GPUs, while tensor parallelism shards those same parameters across all GPUs.
```py
import torch
from transformers import AutoModelForCausalLM, DistributedConfig
distributed_config = DistributedConfig(tp_size=4)
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3-0.6B",
dtype=torch.bfloat16,
distributed_config=distributed_config,
)
```
[Trainer](/docs/transformers/pr_48314/en/main_classes/trainer#transformers.Trainer) detects the tensor parallel plan, reads `tp_size` from the model, and creates a `ParallelismConfig` automatically.
Launch training on one node with 4 GPUs.
```shell
torchrun --nproc-per-node 4 train_tp.py
```
## ParallelismConfig
Pass `ParallelismConfig` explicitly when combining TP with other parallelism techniques like [FSDP](./fsdp).
```py
import torch
from accelerate import ParallelismConfig
from transformers import AutoModelForCausalLM, DistributedConfig, TrainingArguments
distributed_config = DistributedConfig(tp_size=4)
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3-0.6B",
dtype=torch.bfloat16,
distributed_config=distributed_config,
)
parallelism_config = ParallelismConfig(tp_size=4)
args = TrainingArguments(
...,
parallelism_config=parallelism_config,
)
```
## Next steps
- Read the [Tensor Parallelism](https://huggingface.co/spaces/nanotron/ultrascale-playbook?section=tensor_parallelism) chapter from The Ultra-Scale Playbook for more details about how it works.
- Read the [tensor parallelism inference guide](./perf_infer_gpu_multi) to learn more about partitioning strategies, manual TP plans, and implementation details.

Xet Storage Details

Size:
4.24 kB
·
Xet hash:
358abf30a3bdc885342174e71822f76e0757e7af91ae7060436bdac53b0587b9

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.