Distributed training in PyTorch enables efficient scaling of deep learning models across multiple GPUs or nodes. This article explains four core parallelism paradigms—data, pipeline, tensor, and model parallelism—with concise conceptual breakdowns and rewritten, production-ready code examples that avoid redundancy while preserving correctness and clarity.
Data Parallelism (DDP)
Data parallelism replicates the entire model across devices; each device processes a distinct shard of the batch. Gradients are synchronized via all-reduce, ensuring consistent parameter updates. DistributedDataParallel (DDP) is the standard implementation—unlike legacy DataParallel, it avoids master-worker bottlenecks by letting each process maintain its own optimizer and gradient reduction.
- No model broadcasting before forward pass
- Full GPU utilization—no single-device coordination overhead
- Automatic gradient synchronization and buffer synchronization
Minimal DDP Implementation
import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
def setup_ddp():
dist.init_process_group(
backend="nccl",
init_method="env://",
timeout=torch.timedelta(minutes=5)
)
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
return local_rank
class SimpleNet(nn.Module):
def __init__(self, in_dim=784, hidden=256, out_dim=10):
super().__init__()
self.layers = nn.Sequential(
nn.Linear(in_dim, hidden),
nn.ReLU(),
nn.Linear(hidden, out_dim)
)
def forward(self, x):
return self.layers(x)
def main():
local_rank = setup_ddp()
world_size = dist.get_world_size()
# Dataset & loader with sharding
train_dataset = torch.utils.data.TensorDataset(
torch.randn(10000, 784),
torch.randint(0, 10, (10000,))
)
sampler = DistributedSampler(train_dataset, num_replicas=world_size, rank=local_rank)
train_loader = DataLoader(train_dataset, batch_size=64, sampler=sampler, pin_memory=True)
# Model & DDP wrapper
model = SimpleNet().to(local_rank)
ddp_model = nn.parallel.DistributedDataParallel(
model, device_ids=[local_rank], output_device=local_rank
)
optimizer = torch.optim.Adam(ddp_model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
for epoch in range(2):
sampler.set_epoch(epoch)
for x, y in train_loader:
x, y = x.to(local_rank), y.to(local_rank)
optimizer.zero_grad()
loss = criterion(ddp_model(x), y)
loss.backward()
optimizer.step()
if local_rank == 0:
print("Training completed.")
dist.destroy_process_group()
if __name__ == "__main__":
main()
Pipeline Parallelism
Pipeline parallelism partitions the model into ordered stages (e.g., layers), assigning each to a separate device. Input batches are split into micro-batches, enabling overlapping computation: while stage i processes micro-batch k+1, stage i−1 computes micro-batch k. This reduces idle time ("bubbles") and improves hardware utilization.
This example splits an 8-layer transformer across two GPUs using torch.distributed.pipelining:
from torch.distributed.pipelining import PipelineStage, ScheduleGPipe
from torch.nn import TransformerDecoderLayer, LayerNorm, Linear, Embedding
class PipelinedTransformer(nn.Module):
def __init__(self, dim=512, n_layers=8, n_heads=8, vocab=10000):
super().__init__()
self.tok_emb = Embedding(vocab, dim)
self.layers = nn.ModuleList([
TransformerDecoderLayer(dim, n_heads) for _ in range(n_layers)
])
self.norm = LayerNorm(dim)
self.proj = Linear(dim, vocab)
def forward(self, x):
x = self.tok_emb(x)
for layer in self.layers:
x = layer(x, x)
return self.proj(self.norm(x))
def build_stage(model, stage_idx, num_stages, device):
if stage_idx == 0:
# First stage: embedding + first 4 layers
model.layers = nn.ModuleList(model.layers[:4])
model.norm = None
model.proj = None
else:
# Second stage: last 4 layers + projection
model.tok_emb = None
model.layers = nn.ModuleList(model.layers[4:])
return PipelineStage(model, stage_idx, num_stages, device)
# Usage (run with `torchrun --nproc_per_node=2 script.py`)
if __name__ == "__main__":
dist.init_process_group("nccl")
rank = int(os.environ["LOCAL_RANK"])
device = torch.device(f"cuda:{rank}")
model = PipelinedTransformer().to(device)
stage = build_stage(model, rank, world_size=2, device=device)
dummy_input = torch.randint(0, 10000, (32, 128)).to(device)
schedule = ScheduleGPipe(stage, n_microbatches=4)
if rank == 0:
schedule.step(dummy_input)
else:
losses = []
schedule.step(losses=losses)
print(f"Rank {rank} final loss: {losses[-1]:.4f}")
dist.destroy_process_group()
Tensor Parallelism
Tensor parallelism decomposes large matrix operations (e.g., y = x @ W) across devices. Two primary strategies exist:
- Column Parallelism: Splits
Wcolumn-wise → output dimension is sharded. Requiresall-gatherto reconstruct full output. - Row Parallelism: Splits
Wrow-wise → input dimension is sharded. Requiresall-reduceto sum partial outputs.
class ColumnParallelLinear(nn.Module):
def __init__(self, in_features, out_features, bias=True):
super().__init__()
self.world_size = dist.get_world_size()
self.out_features_per_rank = out_features // self.world_size
self.linear = nn.Linear(in_features, self.out_features_per_rank, bias=bias)
def forward(self, x):
local_out = self.linear(x)
# Gather outputs across all ranks
gathered = [torch.empty_like(local_out) for _ in range(self.world_size)]
dist.all_gather(gathered, local_out)
return torch.cat(gathered, dim=-1)
def backward_hook(self, grad_output):
# Sync gradients for weight & bias
if hasattr(self.linear.weight, 'grad') and self.linear.weight.grad is not None:
dist.all_reduce(self.linear.weight.grad, op=dist.ReduceOp.AVG)
if self.linear.bias is not None and self.linear.bias.grad is not None:
dist.all_reduce(self.linear.bias.grad, op=dist.ReduceOp.AVG)
Model Parallelism
Model parallelism assigns disjoint submodules (e.g., encoder/decoder blocks) to different devices. Unlike pipeline parallelism, it does not assume sequential data flow or micro-batching—it’s purely about memory partitioning. Forward and backward passes require explicit inter-device tensor transfers.