Skip to main content
Create your own

Tensor-Parallel Execution for Multi-GPU Deployment

Introduction

In the previous lesson, we tackled the challenge of model size on a single GPU by integrating quantization into our custom inference engine. You implemented a QuantizedLinear layer, allowing the engine to load and execute models with a significantly smaller memory footprint.

However, even with 4-bit quantization, the largest and most powerful models—those with 70 billion parameters or more—still exceed the VRAM capacity of a single consumer or even enterprise GPU. To run these models, we must move from a single-GPU to a multi-GPU architecture. This lesson addresses the learning outcome: Add tensor-parallel execution capabilities to the engine for multi-GPU deployment.

We will explore tensor parallelism (TP), a powerful technique for splitting the weights of a single model layer across multiple GPUs. You will learn the core mechanics of how TP works for transformer blocks and then implement the necessary components to integrate these capabilities directly into your capstone inference engine.

1. The Challenge of "Too-Big-to-Fit" Models

First, let's establish why we need a new parallelism strategy. Quantization reduces memory usage, but there's a hard limit to what it can achieve before model quality degrades unacceptably. For truly massive models, we have no choice but to combine the memory of multiple GPUs.

LLM Inference Optimization #2: Tensor, Data & Expert Parallelism (TP, DP, EP, MoE)

To start, let's watch a brief segment that introduces the problem of growing model sizes and the fundamental need for parallelism in LLM inference.

Watch the first two minutes of this video by Faradawn Yang, from 00:00 to 01:54. It clearly frames the memory and compute challenges that motivate the use of multi-GPU strategies.

As the video explains, when a model's weights cannot fit into one GPU's memory, we must "shard" the model. Tensor parallelism is a specific, highly efficient way of doing this for inference. It works by splitting the computation within a single layer, as opposed to pipeline parallelism which puts different layers on different GPUs.

Tensor Parallelism Illustration for Multi-GPU Deployment
This diagram illustrates the core concept of tensor parallelism. The weight matrix `A` is split into four column-wise shards (A1-A4). The input `X` is sent to all four GPUs, each of which computes a partial result (Y1-Y4) using its local shard. A communication step is then required to combine these partial results into the final output `Y`.

2. The Tensor Parallelism Recipe: Column vs. Row Splitting

The most common and effective method for tensor parallelism in transformers was popularized by the Megatron-LM paper. It follows a specific, clever pattern of splitting weight matrices to minimize communication overhead. The key is to alternate between splitting matrices by columns and by rows.

To build a concrete, step-by-step understanding of this process, we will work through a detailed manual example.

Tensor parallelism by hand

This article by Lewis Won provides an exceptionally clear, by-the-numbers walkthrough of tensor parallelism for a two-layer MLP block. It will give you a solid foundation before we write any code.

Please read the following two sections: 'Forward pass': Follow the detailed calculations from 'Step 1: Initial Setup' through 'Step 3: Forward Pass for Z = Dropout(Y * B)'. Pay close attention to how matrix A is split by columns and B is split by rows, and note where the All-Reduce communication step becomes necessary. 'Why split Matrix A by columns, Matrix B by rows': This section explains the mathematical and efficiency rationale behind the 'Column -> Row' pattern. Understanding this is crucial for implementing it correctly. There is no need to go through the backward pass section for our inference-focused engine.

As you've just read, the Column -> Row pattern is a deliberate design choice.

  • Column Parallelism: The operation Y = f(X @ A) where A is split into [A1, A2] results in a naturally sharded output [Y1, Y2] = [f(X @ A1), f(X @ A2)]. This requires zero communication.
  • Row Parallelism: The operation Z = [Y1, Y2] @ [B1; B2] (where ; denotes row-stacking) is mathematically equivalent to (Y1 @ B1) + (Y2 @ B2). This requires an All-Reduce operation to sum the partial results from each GPU.

By sequencing a column-parallel layer followed by a row-parallel layer, the expensive All-Reduce operation is needed only once for the entire block. This pattern is applied to both the MLP and Attention blocks in a transformer.

Tensor Parallelism in MLP and Self-Attention Layers
This diagram from the original Megatron-LM paper illustrates tensor parallelism for (a) an MLP block and (b) a Self-Attention block. Notice the pattern: the first operation involves splitting weights by columns (e.g., `A1`/`A2`), and the second involves splitting by rows (`B1`/`B2`), requiring communication (an All-Reduce, denoted by the 'g' function) to combine results. The attention block follows a similar pattern, splitting the Q, K, V projections per head and using an All-Reduce for the output projection.

3. Implementing Parallel Layers in PyTorch

Now, let's translate this concept into code for our engine. We'll create two custom nn.Module classes: ColumnParallelLinear and RowParallelLinear. These will act as drop-in replacements for nn.Linear.

The following blog post provides excellent context by showing how production systems like Hugging Face's TGI and vLLM implement this very idea.

Tensor Parallelism and Sequence Parallelism: Detailed Analysis

This post by Insu Jang analyzes tensor parallelism implementations. It connects the concepts we've discussed directly to code.

Please read the section titled 'Tensor Model Parallelism'. Focus on how q_proj, k_proj, and v_proj are mapped to column parallelism, while out_proj is mapped to row parallelism. Pay special attention to the forward method of TensorParallelRowLinear which explicitly shows the torch.distributed.all_reduce call.

Inspired by these examples, we will implement our own versions. Create a new file in your project named parallel.py.

First, add some helper functions to interact with PyTorch's distributed process group:




# parallel.py

import torch
import torch.distributed as dist

_TENSOR_MODEL_PARALLEL_GROUP = None
_TENSOR_MODEL_PARALLEL_WORLD_SIZE = None
_TENSOR_MODEL_PARALLEL_RANK = None

def initialize_model_parallel(tensor_model_parallel_size: int = 1):
    """
    Initializes the tensor model parallel group.
    """
    global _TENSOR_MODEL_PARALLEL_GROUP
    global _TENSOR_MODEL_PARALLEL_WORLD_SIZE
    global _TENSOR_MODEL_PARALLEL_RANK

    if not dist.is_available() or not dist.is_initialized():
        raise RuntimeError("PyTorch distributed is not available or not initialized.")

    _TENSOR_MODEL_PARALLEL_WORLD_SIZE = tensor_model_parallel_size
    



    # Create a process group for tensor parallelism
    # In a more complex setup (with pipeline/data parallelism), you would create subgroups.
    # For now, our TP group is the entire world.
    num_tensor_model_parallel_groups = dist.get_world_size() // tensor_model_parallel_size
    
    for i in range(num_tensor_model_parallel_groups):
        ranks = list(range(i * tensor_model_parallel_size, (i + 1) * tensor_model_parallel_size))
        group = dist.new_group(ranks)
        if dist.get_rank() in ranks:
            _TENSOR_MODEL_PARALLEL_GROUP = group
            _TENSOR_MODEL_PARALLEL_RANK = dist.get_rank() % tensor_model_parallel_size


def get_tensor_model_parallel_world_size():
    return _TENSOR_MODEL_PARALLEL_WORLD_SIZE

def get_tensor_model_parallel_rank():
    return _TENSOR_MODEL_PARALLEL_RANK

def get_tensor_model_parallel_group():
    return _TENSOR_MODEL_PARALLEL_GROUP

Now, add the implementation for the parallel linear layers:




# parallel.py (continued)

import torch.nn as nn
import torch.nn.functional as F
from torch.nn.parameter import Parameter

class _CopyToParallelRegion(torch.autograd.Function):
    """Pass the input to the parallel region."""
    @staticmethod
    def forward(ctx, input_):
        return input_

    @staticmethod
    def backward(ctx, grad_output):
        return dist.all_reduce(grad_output, group=get_tensor_model_parallel_group())

class _ReduceFromParallelRegion(torch.autograd.Function):
    """All-reduce the input from the parallel region."""
    @staticmethod
    def forward(ctx, input_):
        return dist.all_reduce(input_, group=get_tensor_model_parallel_group())

    @staticmethod
    def backward(ctx, grad_output):
        return grad_output


class ColumnParallelLinear(nn.Module):
    def __init__(self, in_features: int, out_features: int, bias: bool = True):
        super().__init__()
        
        self.in_features = in_features
        self.out_features = out_features
        world_size = get_tensor_model_parallel_world_size()

        if out_features % world_size != 0:
            raise ValueError(f"out_features ({out_features}) must be divisible by "
                             f"tensor parallel world_size ({world_size})")

        self.out_features_per_partition = out_features // world_size
        
        self.weight = Parameter(torch.empty(
            self.out_features_per_partition, self.in_features
        ))
        if bias:
            self.bias = Parameter(torch.empty(self.out_features_per_partition))
        else:
            self.register_parameter('bias', None)

    def forward(self, input_):



        # The input is expected to be a complete tensor (not sharded).
        # We perform a local matmul with the sharded weight.
        output_parallel = F.linear(input_, self.weight, self.bias)
        



        # The output is now sharded. For the backward pass, we need to gather gradients.
        # This is handled by applying the autograd function.
        return _CopyToParallelRegion.apply(output_parallel)


class RowParallelLinear(nn.Module):
    def __init__(self, in_features: int, out_features: int, bias: bool = True):
        super().__init__()

        self.in_features = in_features
        self.out_features = out_features
        world_size = get_tensor_model_parallel_world_size()

        if in_features % world_size != 0:
            raise ValueError(f"in_features ({in_features}) must be divisible by "
                             f"tensor parallel world_size ({world_size})")

        self.in_features_per_partition = in_features // world_size

        self.weight = Parameter(torch.empty(
            self.out_features, self.in_features_per_partition
        ))
        if bias:
            self.bias = Parameter(torch.empty(self.out_features))
        else:
            self.register_parameter('bias', None)

    def forward(self, input_):



        # The input is expected to be sharded.
        # Apply the autograd function to handle gradient propagation correctly.
        input_parallel = _ReduceFromParallelRegion.apply(input_)




        # Perform a local matmul. This gives a partial result.
        output_parallel = F.linear(input_parallel, self.weight)




        # All-reduce the partial results to get the final output.
        # This is the main communication step in the forward pass.
        output = dist.all_reduce(output_parallel, group=get_tensor_model_parallel_group(), async_op=False)

        if self.bias is not None:
            output = output + self.bias
            
        return output
  • The _CopyToParallelRegion and _ReduceFromParallelRegion autograd functions are crucial for ensuring gradients are correctly handled during backpropagation (even if we're only doing inference, using Parameter requires correct grad handling). The key takeaway is how the backward pass communication is the reverse of the forward pass.
  • ColumnParallelLinear splits its weight and bias along the output dimension. Its forward pass produces a sharded output.
  • RowParallelLinear splits its weight along the input dimension. Its forward pass takes a sharded input and performs an all_reduce to produce a complete output.

4. Integrating Tensor Parallelism into the Engine

The final step is to modify your engine's model loading procedure to use these new layers. This is analogous to how you integrated the QuantizedLinear layer. You'll traverse the model and replace nn.Linear layers with either ColumnParallelLinear or RowParallelLinear based on their name and role in the transformer block.

1. Create the replacement function in parallel.py:




# parallel.py (continued)
import logging

def replace_with_tensor_parallel_layers(module, prefix=""):
    for name, child in module.named_children():
        new_prefix = f"{prefix}.{name}" if prefix else name
        if isinstance(child, nn.Linear):



            # For decoder-only models, these are the typical layer names
            # that get column-wise parallelism.
            if any(n in new_prefix for n in ["q_proj", "k_proj", "v_proj", "gate_proj", "up_proj"]):
                new_layer = ColumnParallelLinear(child.in_features, child.out_features, child.bias is not None)



            # These layers get row-wise parallelism.
            elif any(n in new_prefix for n in ["o_proj", "down_proj"]):
                new_layer = RowParallelLinear(child.in_features, child.out_features, child.bias is not None)
            else:
                logging.warning(f"Layer {new_prefix} not recognized for TP. Keeping as is.")
                continue

            setattr(module, name, new_layer)
            logging.info(f"Replaced {new_prefix} with {type(new_layer).__name__}")
        else:



            # Recurse into submodules
            replace_with_tensor_parallel_layers(child, prefix=new_prefix)

2. Modify your engine launcher script:

Your main script that launches the engine must now initialize the distributed environment.




# main_engine.py (example)
import torch
import torch.distributed as dist
from parallel import initialize_model_parallel, replace_with_tensor_parallel_layers, get_tensor_model_parallel_rank

def main():



    # 1. Initialize Distributed Environment
    dist.init_process_group(backend="nccl")
    world_size = dist.get_world_size()
    rank = dist.get_rank()
    torch.cuda.set_device(rank)




    # For now, let's assume TP size is the world size
    tensor_parallel_size = world_size
    initialize_model_parallel(tensor_parallel_size)




    # 2. Load Model Architecture (on CPU first to avoid OOM)
    # model = AutoModelForCausalLM.from_pretrained(..., torch_dtype=torch.float16, use_cache=False)
    # # model is on meta device or CPU




    # 3. Replace layers with TP versions
    # replace_with_tensor_parallel_layers(model)




    # 4. Load Sharded Weights
    # Now, you need a mechanism to load the weights. Each rank loads only its shard.
    # For a layer replaced with ColumnParallelLinear, rank 'i' would load columns
    # from (i * shard_size) to ((i+1) * shard_size) of the original weight matrix.
    # This is a non-trivial step. Frameworks like Hugging Face's Accelerate can help.
    # For this exercise, you can manually shard a saved checkpoint.
    



    # A simplified loading example:
    # with torch.no_grad():
    #     for name, param in model.named_parameters():
    #         # logic to load the correct shard of 'param' from a file
    #         # into the parameter on the current rank.
    #         pass
    



    # model.to(torch.cuda.current_device())
    



    # ... rest of your engine setup (Scheduler, PagedKVCacheManager, Executor)

if __name__ == "__main__":
    main()

To run this, you would use torchrun:
torchrun --nproc_per_node=2 main_engine.py (for 2 GPUs on one machine).

The sharded weight loading is a complex topic. For this capstone, you can simplify it by saving a model's state dict, then writing a separate utility script that splits the weights for each relevant layer into tp_rank_00_*.bin, tp_rank_01_*.bin, etc., which your main_engine.py script can then load based on its rank.

Conclusion

This lesson has been a deep dive into one of the most fundamental techniques for scaling LLM inference. You've gone from the high-level concept of model sharding to the low-level mechanics of column and row parallelism and their implementation in PyTorch.

Key Takeaways:

  • Tensor Parallelism (TP) splits individual weight matrices across GPUs, enabling models that are too large for a single device.
  • The Column -> Row parallelism pattern is a highly efficient strategy that minimizes communication by requiring only one all_reduce operation per transformer block in the forward pass.
  • Integrating TP into an engine involves creating custom parallel layers (ColumnParallelLinear, RowParallelLinear) and a model surgery process to replace the standard nn.Linear layers.
  • Running a TP model requires a distributed environment (initialized via torch.distributed) and a sharded checkpoint loading mechanism where each GPU rank loads only its specific slice of the weights.

Preview of the next lesson:
Your custom engine is now architecturally complete. It has a scheduler, an advanced memory manager, and support for both single-GPU optimization (quantization) and multi-GPU scaling (tensor parallelism). In the next lesson, we will complete the capstone project by implementing the final user-facing components: request ingestion and response streaming layers for an end-to-end test. This will turn your powerful backend into a functional service.

Can't find a good explanation? Sign up and we'll make it for you

Sign up