PyTorch Distributed Training Strategies: Data, Pipeline, Tensor, and Model Parallelism

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 W column-wise → output dimension is sharded. Requires all-gather to reconstruct full output.
  • Row Parallelism: Splits W row-wise → input dimension is sharded. Requires all-reduce to 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.

Tags: pytorch distributed-training data-parallelism pipeline-parallelism tensor-parallelism

Posted on Mon, 31 Aug 2026 16:09:06 +0000 by progman