What Is the Difference Between ColumnParallelLinear and RowParallelLinear in Llama 2?
ColumnParallelLinear shards the output dimension (columns) of weight matrices for input-side projections like Q/K/V and requires an all-gather operation, whereas RowParallelLinear shards the input dimension (rows) for output-side projections like the attention output and requires an all-reduce sum.
The Meta Llama 2 repository implements tensor parallelism through FairScale to distribute massive transformer computations across multiple GPUs. Understanding the difference between ColumnParallelLinear and RowParallelLinear in Llama 2 clarifies how the architecture minimizes communication overhead while balancing memory consumption across devices.
How Tensor Parallelism Splits Linear Layers in Llama 2
Llama 2 employs two complementary sharding strategies from FairScale to parallelize linear transformations. Both approaches partition the weight matrix across tensor-parallel ranks, but they differ in which dimension is split and the subsequent communication required to reconstruct the full tensor.
ColumnParallelLinear Shards the Output Dimension
ColumnParallelLinear divides the weight matrix along its column axis, meaning each GPU stores a distinct slice of the output features. This strategy is employed for input-side projections where the activation size is relatively small. In llama/model.py, the query, key, and value projection layers (wq, wk, wv) within the Attention class are instantiated as ColumnParallelLinear layers (lines 207-221).
After each GPU performs its local matrix multiplication on the input activations, the results must be consolidated. Because each device computed a portion of the full output features, the system executes an all-gather operation to concatenate the partial results along the output dimension, ensuring every GPU receives the complete projected tensor for Q, K, and V.
RowParallelLinear Shards the Input Dimension
RowParallelLinear partitions the weight matrix along its row axis, with each GPU holding a slice of the input features. This approach appears in output-side projections where the weight matrix connects to larger activation tensors. In llama/model.py, the attention output projection (wo) is implemented as a RowParallelLinear layer (lines 228-234).
Following the local matrix multiplication, each GPU holds partial results that sum to the final output. Rather than gathering full tensors, the system performs an all-reduce operation (specifically a sum) across devices to combine these partial outputs. This is more efficient than an all-gather for these specific layers because the reduction produces the final tensor directly without requiring additional memory overhead.
Communication Patterns and Performance Trade-offs
The dual approach exists to optimize both memory balance and communication efficiency. Column-parallel sharding reduces memory pressure for the large Q/K/V weight matrices, while row-parallel sharding minimizes the memory footprint for activation tensors in the output projection.
The communication costs differ significantly between the two. The all-gather operation used after column-parallel layers is inexpensive for Q/K/V projections because the activation sizes remain relatively small. Conversely, the all-reduce operation following row-parallel layers is efficient for output projections because the weight matrix is already partitioned row-wise, allowing the summation to occur without broadcasting full tensors.
Implementation in llama/model.py
The specific wiring of these parallel linear types appears throughout llama/model.py. The FairScale implementations are imported at the top of the file (lines 8-15).
In the Attention class constructor:
self.wq,self.wk, andself.wvare instantiated usingColumnParallelLinear(lines 191-194 and 207-221), corresponding to the input projections for queries, keys, and values.self.wousesRowParallelLinear(lines 228-234) for the output projection that combines attention heads.
The same pattern extends to the feed-forward network, where the first and third linear transformations (w1 and w3) use column-parallel sharding, while the middle layer (w2) uses row-parallel sharding (lines 24-28).
Practical Code Example
The following example demonstrates how these parallel linear layers function within the Attention module:
import torch
from llama.model import Attention, ModelArgs
# Configuration matching a small model
args = ModelArgs(dim=256, n_heads=8, max_batch_size=2, max_seq_len=128)
# Initialize the attention module with parallel linear layers
attn = Attention(args)
# Simulated input: batch_size=2, seq_len=4, hidden_dim=256
x = torch.randn(2, 4, args.dim)
# Simplified frequency tensor for rotary embeddings
freqs_cis = torch.randn(1, args.dim // args.n_heads, dtype=torch.cfloat)
# Forward pass through the attention mechanism
output = attn(x, start_pos=0, freqs_cis=freqs_cis, mask=None)
print(output.shape) # Output: torch.Size([2, 4, 256])
Under the hood, the forward pass executes the column-parallel multiplications for Q/K/V (each followed by all-gather), computes attention, then applies the row-parallel output projection (followed by all-reduce) to produce the final result.
Summary
- ColumnParallelLinear partitions weight matrices by columns, stores slices of output dimensions, and requires all-gather communication. Used in Llama 2 for Q/K/V projections (
wq,wk,wv) in the attention mechanism and the first/third layers of feed-forward blocks. - RowParallelLinear partitions weight matrices by rows, stores slices of input dimensions, and requires all-reduce (sum) communication. Used in Llama 2 for the attention output projection (
wo) and the middle layer (w2) of feed-forward blocks. - File locations: Both classes are imported from FairScale in
llama/model.py(lines 8-15), with instantiations occurring at lines 191-194, 207-221, and 228-234 for attention layers, and lines 24-28 for feed-forward layers.
Frequently Asked Questions
Why does Llama 2 use both column and row parallelism instead of just one?
Using both strategies optimizes the memory and communication trade-offs inherent in transformer architectures. Column parallelism reduces the memory footprint for large weight matrices in input projections, while row parallelism minimizes activation memory in output projections. This hybrid approach ensures that no single layer becomes a bottleneck for GPU memory or inter-device communication bandwidth.
Which operation is more computationally expensive: all-gather or all-reduce?
For Llama 2's specific configuration, the all-gather following column-parallel layers is relatively inexpensive because Q/K/V activations are smaller dimensional tensors. The all-reduce following row-parallel layers is also efficient because it performs a sum reduction rather than concatenating full tensors. The choice of which to use depends on the specific layer dimensions and whether the bottleneck lies in weight matrix storage or activation memory.
Can these parallel linear layers function with pipeline parallelism?
Yes, tensor parallelism (implemented via ColumnParallelLinear and RowParallelLinear) is orthogonal to pipeline parallelism. Llama 2 can combine both: tensor parallelism splits individual layers across GPUs within a single stage, while pipeline parallelism distributes different layers across different stages. The FairScale implementations handle the necessary communication collectives regardless of the broader parallelism strategy.
Where do these classes originate if they are not defined directly in the Llama repository?
Both ColumnParallelLinear and RowParallelLinear are imported from the FairScale library, which Meta developed for efficient large-scale training. In llama/model.py (lines 8-15), these classes are imported from fairscale.nn.model_parallel.layers and wrapped within the Llama model architecture to provide the specific sharding behavior described in the tensor-parallel implementation.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →