Data parallelism copies the whole model to every GPU and splits the batch. That stops working the moment one layer's weights, gradients and activations no longer fit on one card, or when a single request must be served faster than one GPU can compute it. Tensor parallelism (TP) attacks the problem from the other side: it splits the matrices inside each layer across several GPUs, so every GPU holds a slice of every layer and they all work on the same tokens at the same time.

The price is communication inside every layer, on the critical path: two all-reduces per transformer layer in each direction, one for attention and one for the MLP. This article derives the split from matrix algebra, implements it from scratch and with PyTorch's DTensor API, computes the communication budget for a realistic layer, and ends with failure modes and a checklist, so you can choose a TP degree and defend it.

Advertisement

Where tensor parallelism fits

There are three ways to spread one model over many GPUs. Data parallelism splits the batch and synchronises gradients once per step. Pipeline parallelism splits the model by layers and passes activations between stages. Tensor parallelism splits each layer's weight matrices, so a single matrix multiplication is computed jointly by a group of GPUs. Sharded data parallelism (ZeRO and FSDP) also splits weights, but gathers each layer's full weights before computing with them; TP never materialises the full matrix anywhere.

Because TP never gathers weights, it reduces per-GPU weight memory and per-GPU compute for every token, which also cuts latency for inference. Because the GPUs must combine partial results inside each layer, its collectives are small, frequent and blocking. They must run over the fastest link available, which in practice means NVLink inside one server. The usual rule is TP inside a node, pipeline and data parallelism across nodes.

The algebra: column splits and row splits

Take a linear layer Y = X A, where X has shape [tokens, k] and A has shape [k, n]. There are two ways to split A across t GPUs.

Column split. Cut A into t blocks of columns, A = [A1 ... At]. GPU i computes Yi = X Ai, a block of output columns. No communication is needed in the forward pass as long as every GPU has the full X. The output is sharded along its last dimension.

Row split. Cut A into t blocks of rows and cut X into matching blocks of columns. GPU i computes a partial product Xi Ai of the full output shape, and the true result is the sum over i. The partial sums must be combined with an all-reduce, which adds the tensors from all GPUs and leaves the sum on each of them.

The trick in Megatron-LM (Shoeybi et al., 2019) is to chain the two. A column-split layer produces exactly the sharded input a row-split layer wants, so the pair needs no communication in the middle and one all-reduce at the end. Any element-wise function between them, such as GeLU, can be applied locally to each shard because it acts on each column independently. That is why a transformer MLP, a linear up-projection, an activation and a linear down-projection, maps perfectly onto TP.

Advertisement

The MLP block and the f and g operators

Megatron expresses the communication as two conjugate autograd operators. f sits at the entry of the parallel region: in the forward pass it is the identity (every GPU already has X), and in the backward pass it all-reduces the gradient with respect to X, because each GPU only computed the part of that gradient flowing through its own column slice. g sits at the exit: in the forward pass it all-reduces the partial outputs, and in the backward pass it is the identity, because the incoming gradient is already replicated.

Megatron MLP block split over two GPUs: one all-reduce forward, one backwardX [b, s, h]replicated on bothf: identity forwardall-reduce backwardf: identity forwardall-reduce backwardGPU 0: X @ A0A0 = columns 0..4h/2 of AGPU 1: X @ A1A1 = columns 4h/2..4h of AGeLU (local, no comm)GeLU (local, no comm)GPU 0: Y0 @ B0B0 = rows 0..4h/2 of BGPU 1: Y1 @ B1B1 = rows 4h/2..4h of Bg: all-reduce (sum) forwardidentity backwardpartial Z0partial Z1Z = Z0 + Z1
Column-parallel first GEMM, local GeLU, row-parallel second GEMM. The only communication is g in the forward pass and f in the backward pass.

Per GPU, weight memory for the MLP drops by a factor of t and compute drops by t. Communication per MLP block is one all-reduce of a [b, s, h] activation forward and one backward. Dropout after the block and the residual add operate on the replicated output, so they run redundantly on every GPU, and they must use the same random seed on all ranks or the replicas will diverge.

Attention, embeddings and the loss

Self-attention splits even more naturally, because heads are independent. The Q, K and V projections are column-split so that each GPU owns a contiguous set of whole heads, computes attention for those heads locally, and feeds its slice of the context vectors into a row-split output projection followed by g. So attention also costs one all-reduce per direction. This imposes the first hard constraint: the number of attention heads must be divisible by t. With grouped-query attention there are fewer key-value heads than query heads; if t exceeds the number of KV heads, the KV projections are replicated across ranks, which wastes memory and is a reason to cap t at the KV-head count.

The embedding and the output layer carry the largest single matrices in many models, so they are split along the vocabulary. In a vocabulary-parallel embedding each rank holds a contiguous range of token ids, looks up the ids that fall in its range, writes zeros for the rest, and an all-reduce produces the full embedding. For the loss, gathering logits of shape [b, s, vocab] would be enormous, so the parallel cross-entropy computes the maximum and the sum of exponentials locally, all-reduces those two small tensors, and each rank subtracts the target logit if it owns that id. Megatron pads the vocabulary so it divides evenly by t; your tokenizer's size will usually not.

A from-scratch implementation

The core fits in a few dozen lines of PyTorch. The two autograd functions below are f and g; the two modules are the column- and row-parallel linears. Initialisation must produce slices of one full matrix, so every rank initialises the full weight from a shared seed and keeps its slice (real frameworks initialise the slice directly with a partitioned RNG).

import torch, torch.distributed as dist
import torch.nn as nn, torch.nn.functional as F

class CopyToTP(torch.autograd.Function):        # f
    @staticmethod
    def forward(ctx, x, group):
        ctx.group = group
        return x
    @staticmethod
    def backward(ctx, grad):
        grad = grad.contiguous()
        dist.all_reduce(grad, group=ctx.group)
        return grad, None

class ReduceFromTP(torch.autograd.Function):    # g
    @staticmethod
    def forward(ctx, x, group):
        x = x.contiguous()
        dist.all_reduce(x, group=group)
        return x
    @staticmethod
    def backward(ctx, grad):
        return grad, None

class ColumnParallelLinear(nn.Module):
    def __init__(self, k, n, group, seed=0):
        super().__init__()
        t, r = dist.get_world_size(group), dist.get_rank(group)
        assert n % t == 0
        full = torch.empty(n, k)
        torch.manual_seed(seed); nn.init.xavier_normal_(full)
        self.weight = nn.Parameter(full.chunk(t, dim=0)[r].clone())
        self.group = group
    def forward(self, x):
        return F.linear(CopyToTP.apply(x, self.group), self.weight)

class RowParallelLinear(nn.Module):
    def __init__(self, k, n, group, seed=1):
        super().__init__()
        t, r = dist.get_world_size(group), dist.get_rank(group)
        assert k % t == 0
        full = torch.empty(n, k)
        torch.manual_seed(seed); nn.init.xavier_normal_(full)
        self.weight = nn.Parameter(full.chunk(t, dim=1)[r].clone())
        self.group = group
    def forward(self, x_shard):
        return ReduceFromTP.apply(F.linear(x_shard, self.weight), self.group)

class ParallelMLP(nn.Module):
    def __init__(self, h, group):
        super().__init__()
        self.up = ColumnParallelLinear(h, 4 * h, group)
        self.down = RowParallelLinear(4 * h, h, group)
    def forward(self, x):
        return self.down(F.gelu(self.up(x)))

Test it against the unsplit MLP built from the same seeds: outputs and input gradients must match within bf16 tolerance. Most TP bugs are silent numerical errors, not crashes.

In production you rarely write this yourself. PyTorch's DTensor API expresses the same plan declaratively: parallelize_module takes a module, a one-dimensional device mesh and a plan that maps submodule names to ColwiseParallel, RowwiseParallel or SequenceParallel.

from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import (
    parallelize_module, ColwiseParallel, RowwiseParallel)

mesh = init_device_mesh("cuda", (2, 8), mesh_dim_names=("dp", "tp"))
tp_mesh = mesh["tp"]                      # parallelize_module wants a 1-D mesh

for block in model.layers:
    parallelize_module(block, tp_mesh, {
        "attn.q_proj": ColwiseParallel(),
        "attn.k_proj": ColwiseParallel(),
        "attn.v_proj": ColwiseParallel(),
        "attn.o_proj": RowwiseParallel(),
        "mlp.up_proj": ColwiseParallel(),
        "mlp.down_proj": RowwiseParallel(),
    })
    # attention code must use the local head count: n_heads // tp_mesh.size()

Worked example: the communication budget

Take one layer with hidden size h = 8,192, sequence length s = 4,096, micro-batch b = 1, bf16 activations and t = 8. One activation tensor is b × s × h × 2 bytes = 64 MiB. A ring all-reduce sends 2(t − 1)/t of the tensor from each rank, so each all-reduce moves 112 MiB per GPU.

Forward compute per layer is about 24bsh² for the four projections and the MLP plus 4bs²h for the attention scores and weighted sum, which here is 7.15 TFLOP, or 0.89 TFLOP per GPU. Assume the GPUs sustain 600 TFLOPS of dense bf16 (well below the H100's datasheet peak) and that the all-reduce achieves 350 GB/s of bus bandwidth on NVLink. Those two numbers are illustrative assumptions; measure yours with a profiler and nccl-tests.

Quantity (per layer, forward)Value
Compute time per GPU1.49 ms
Two all-reduces over NVLink (assumed 350 GB/s)0.67 ms (45% of compute)
Two all-reduces over one 400 Gb/s NIC per GPU (assumed 50 GB/s)4.70 ms (316% of compute)
Layer weights, full / per GPU (bf16)1.61 GB / 201 MB

Inside one NVLink server the all-reduces are a significant but tolerable tax; over the network they take longer than the compute they serve, which is the quantitative reason TP stays inside a node. The ratio also shrinks as h grows, because compute scales with h² while communication scales with h: large models tolerate TP better than small ones, and a small model split eight ways mostly waits on NCCL collectives.

Sequence parallelism and activation memory

Plain TP splits the GEMMs but leaves LayerNorm, dropout and the residual stream replicated: every GPU stores full [b, s, h] activations for them. Korthikanti et al. (2022) observed that those regions are independent along the sequence dimension, so they can be split by sequence instead. The all-reduce at g is replaced by a reduce-scatter that leaves each rank with s/t tokens of the summed output, and the identity at f is replaced by an all-gather that rebuilds the full sequence before the next column-parallel GEMM.

A reduce-scatter plus an all-gather moves the same bytes as one all-reduce, because a ring all-reduce is implemented as exactly that pair. So sequence parallelism removes the replicated activation memory for roughly zero extra communication. The cost is complexity: LayerNorm gradients now come from sequence shards and must be all-reduced across the TP group, or the replicas drift slowly.

Combining TP with data and pipeline parallelism

Large runs use all three. A typical layout is TP = 8 inside each server, pipeline stages across servers, and data parallelism (often sharded) across groups of pipelines. Rank ordering is what makes that work: the device mesh must place the TP dimension innermost so that the eight members of each TP group are the eight GPUs of one server. If the mesh is built the other way round, the job runs, gives correct results and is several times slower, because every TP all-reduce crosses the network.

Larger t cuts memory and latency but shrinks each GEMM and raises the communication share. Choose the smallest t at which the model state and activations fit with the micro-batch you want, cap it at the NVLink domain size and at the KV-head count, and use pipeline or data parallelism for the rest of the scale. For picking hardware with enough memory in the first place, see GPU selection.

Tensor parallelism for inference

Serving uses the same layout for a different reason. Decoding reads every weight once per token, so it is bound by memory bandwidth; splitting the weights over t GPUs lets t memory systems read in parallel. Servers such as vLLM expose a tensor-parallel size with the same divisibility and NVLink constraints. Each decode step all-reduces a tiny tensor, so the fixed latency of each collective eventually dominates; for throughput rather than latency, several smaller replicas usually beat one wider one.

Failure modes and how to debug them

SymptomLikely causeWhat to check
Job hangs at the first layerRanks issue collectives in different orders or on different groupsNCCL_DEBUG=INFO logs; that every rank takes the same code path
Loss slightly worse than single-GPU baselineDropout seeds identical inside TP regions or different outside themRNG tracker: distinct seeds for split regions, same seed for replicated ones
Replicas drift over thousands of stepsGradients of replicated parameters (LayerNorm with SP) not all-reducedCompare parameter checksums across the TP group
Much slower than expected, no errorsTP group spans nodes or PCIe instead of NVLinknvidia-smi topo -m; mesh ordering; nccl-tests on the TP group
Shape error at load timeCheckpoint saved with another TP degreeReshard offline, or use a distributed checkpoint format that reshards on load
Assertion on head countn_heads or KV heads not divisible by tLower t or replicate KV heads deliberately

Communication overlap is the main performance lever once correctness is established; the techniques for hiding collectives behind compute are covered in collective overlap.

What to do next

  1. Compute your model state per GPU at t = 1, 2, 4, 8 and pick the smallest t that fits with your target micro-batch.
  2. Check that attention heads, KV heads and the padded vocabulary divide by t.
  3. Run nccl-tests all_reduce_perf on exactly the GPUs of one TP group and record the achieved bus bandwidth.
  4. Redo the worked example with your h, s, measured bandwidth and measured TFLOPS to estimate the communication share.
  5. Build the device mesh with TP innermost and verify with a trace that TP collectives never leave the node.
  6. Validate numerics against an unsplit model on a small config before scaling up; enable sequence parallelism and repeat.
  7. Profile a step and look for exposed all-reduce time on the critical path before tuning anything else.
Key takeaway: Tensor parallelism splits each layer's matrices across GPUs: column-split then row-split, with local element-wise work in between, so each attention or MLP block costs one all-reduce forward and one backward, two per layer in each direction. It cuts per-GPU weight memory, compute and inference latency, but its collectives are small, frequent and blocking, so keep it inside the NVLink domain, cap it at the head and KV-head counts, add sequence parallelism to remove replicated activations, and measure the communication share with your own numbers before choosing t.