Tensor & Pipeline Parallelism¶
Overview¶
When data parallelism (DDP/FSDP) stops scaling— because the model is too big for one GPU, or the interconnect is too slow— you split the model itself:
- Tensor Parallelism (TP): split each layer's weight matrices across ranks; each rank computes a slice of the output. High comm per layer (all-to-all style) → needs fast NVLink.
- Pipeline Parallelism (PP): split the model by layers into stages; ranks run different stages with micro-batch pipelining. Comm is small (stage boundaries) but bubbles appear.
- PyTorch implements TP via
torch.distributed.tensor.parallel(DTensor), PP viatorch.distributed.pipelining(new) ortorch.distributed.pipeline.sync(legacy, Pipe).
Golden rule: TP trades bandwidth for memory; PP trades latency/bubble for memory. Combine with DDP/FSDP (hybrid) for big models.
Tensor Parallelism— How a Linear Splits¶
For y = x @ W^T with 2 ranks:
- Split
W: rank 0 getsW[:,:K], rank 1 getsW[:, K:](column split). - Each rank computes
x @ W_part^T→ partialy. - All-reduce the partials → full
y.
Row-split variant: for the preceding layer, split x vertically so the matmul of the next layer starts from reduced data— PyTorch's DTensor handles this transparently.
import torch, torch.nn as nn, torch.distributed as dist
from torch.distributed.tensor.parallel import parallelize_module
from torch.distributed.tensor import DeviceMesh
def main(rank, world_size):
dist.init_process_group("nccl", rank=rank, world_size=world_size)
torch.cuda.set_device(rank)
mesh = DeviceMesh("cuda", list(range(world_size))) # 1 mesh dim
model = nn.Linear(8, 8).cuda()
# parallelize the Linear's weight by "colwise" (output dim)
tp_model = parallelize_module(model, mesh, {"": {"linear": "ColwiseParallel"}})
x = torch.randn(4, 8).cuda()
out = tp_model(x)
# out now requires all-reduce internally; result same as serial model
torch.cuda.synchronize()
if rank == 0:
print("TP output works, shape:", out.shape)
dist.destroy_process_group()
Attention-specific TP (classic LLM pattern)¶
| OP | Split scheme |
|---|---|
| QKV projection | colwise → each rank has full Q/K/V rows |
| attention softmax + matmul | all-reduce partial outputs |
| out projection | rowwise → gather+reduce |
torch.distributed.tensor.parallelprovidesColwiseParallel,RowwiseParallel,SequenceParallel; always verify equivalence to the un-parallelized model on a tiny seed.
-
Pipeline Parallelism— The Bubble¶
Micro-batches pipelined across stages fill the pipeline; throughput gap = bubble time. Standard schedule (GPipe in PyTorch):
stream: 1F1B (one-forward-one-backward) — classic:
f0 f1 f2 f3
stage0: [1][2][3][4]...
Schedule example (4 stages, 4 micro-batches):
step: 1 2 3 4 5 6 7 8
s0: f1 f2 f3 f4 b4 b3 b2 b1
s1: f1 f2 f3 f4 b4 b3 b2 b1
s2: f1 f2 f3 f4 b4 b3 b2 b1
s3: f1 f2 f3 f4 b4 b3 b2 b1
The bubble = idle slots at ramp-up/down ≈ (P-1)/P of pipeline depth.
More micro-batches = smaller bubble ratio, but if you don't have enough memory to hold them, you can't win.
-
Manual PP (the honest version)¶
import torch, torch.nn as nn, torch.distributed as dist
def make_stage(rank):
return nn.Sequential(nn.Linear(8, 8), nn.ReLU())
def main(rank, world_size):
dist.init_process_group("nccl", rank=rank, world_size=world_size)
torch.cuda.set_device(rank)
stage = make_stage(rank).cuda()
x = torch.randn(4, 8).cuda()
if rank == 0:
# forward stage 0 -> send to stage 1
out = stage(x)
dist.send(out.contiguous(), dst=1)
elif rank == world_size - 1:
# receive from previous, compute last stage
recv = torch.empty(4, 8).cuda()
dist.recv(recv, src=rank - 1)
out = stage(recv)
if rank == world_size - 1:
print("PP forward complete, shape:", out.shape)
else:
recv = torch.empty(4, 8).cuda()
dist.recv(recv, src=rank - 1)
out = stage(recv)
dist.send(out.contiguous(), dst=rank + 1)
dist.destroy_process_group()
if __name__ == "__main__":
import torch.multiprocessing as mp
mp.spawn(main, args=(3,), nprocs=3, join=True)
Real libraries (
torch.distributed.pipelining.Pipeline,torchpipe, DeepSpeed PP) implement 1F1B + scheduling + backward for you— manual send/recv above is for understanding, not production.
-
Hybrid Parallelism (TP × PP × DP)— the LLM recipe¶
mesh: [DP=2, PP=2, TP=2]
rank layout: dp-rank * (pp * tp)... organized as a 3D mesh
- FSDP/DDP over DP dim, pipeline over PP, TensorParallel over TP dim
- Communication: TP on NVLink; PP on node-links; DP across nodes
from torch.distributed.tensor import DeviceMesh
mesh = DeviceMesh("cuda", list(range(8)), mesh_dim_names=["dp", "pp", "tp"])
# PP and TP use DeviceMesh from torch.distributed.tensor.parallel / pipelining
Practical sizing: TP ≤ 8 (NVLink width), PP = #stages/GPU count, DP fills the rest. Keep TP fastest interconnect, PP to the next, DP slowest.
Which to Use When (Decision Table)¶
| Situation | Strategy | Why |
|---|---|---|
| Model fits one GPU | DDP/FSDP | simplest |
| Model > 1 GPU memory (params) | FSDP or TP | shard params |
| Model huge, nodes slow | TP + PP + FSDP | hybrid |
| Latency-sensitive inference | TP (not PP) | no bubble |
| Training on 2 GPUs, small model | DDP | TP overhead not worth it |
Verification & Debugging¶
- Numerical check: TP must equal the serial model (same seed)— any mismatch = split bug.
- Comm %:
torch.profiler— if all-reduce >30% of step, interconnect-bound: reduce TP size or use FSDP instead. - Bubble measuring: PP idle kernel slots in the timeline— increase micro-batches.
- Use
TORCH_DISTRIBUTED_DEBUG=DETAILto spot deadlocks and shape mismatches.
-
Key Takeaways¶
- TP splits weight matrices; comm per layer; needs fast interconnect; verify numerics.
- PP splits by layers; small comm, bubble cost; micro-batches fill the pipe.
- Hybrid (DP × PP × TP) is the production LLM pattern— arrange mesh dims by speed.
- Never reimplement send/recv in prod— use
torch.distributed.pipelining/tensor.parallelAPIs. - Measure comm% vs compute% before and after every parallelizaton change.
-