Grouped-Query Attention vs Multi-Head Attention: Understanding GQA Implementation

Grouped-Query Attention (GQA) reduces memory consumption by sharing key and value projections across groups of attention heads while maintaining independent query projections, achieving comparable expressivity to Multi-Head Attention (MHA) with significantly smaller KV-cache footprints.

When implementing transformer architectures from scratch in the rasbt/LLMs-from-scratch repository, choosing between attention mechanisms directly impacts your model's memory efficiency and inference speed. This guide examines the precise architectural differences between Grouped-Query Attention and Multi-Head Attention as implemented in the source code, demonstrating how GQA optimizes the KV-cache through selective parameter sharing.

Architectural Differences Between MHA and GQA

Both mechanisms compute scaled dot-product attention, but they diverge in how they project and share keys and values across heads.

Parameter Projection Patterns

In standard Multi-Head Attention, every head maintains independent projection matrices for keys and values. According to the implementation in ch04/04_gqa/gpt_with_kv_mha.py, the MultiHeadAttention class creates distinct linear layers for each head:

self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)

After projection, the tensors are reshaped to (B, num_heads, T, head_dim), giving each head exclusive access to its own key-value pairs.

Grouped-Query Attention, implemented in ch04/04_gqa/gpt_with_kv_gqa.py, follows a different pattern in the GroupedQueryAttention class. Keys and values are projected once per KV-group rather than per head, significantly reducing parameters:

self.group_size = num_heads // num_kv_groups  # heads sharing the same KV

# Single projection per KV-group (shared across its heads)

self.W_key   = nn.Linear(d_in, num_kv_groups * self.head_dim,
                         bias=qkv_bias, dtype=dtype)
self.W_value = nn.Linear(d_in, num_kv_groups * self.head_dim,
                         bias=qkv_bias, dtype=dtype)

# Independent projection for queries (per head)

self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias, dtype=dtype)

The assertion assert num_heads % num_kv_groups == 0 ensures even group distribution.

Memory Footprint and KV-Cache Optimization

The dimensional differences create massive memory savings during autoregressive generation.

MHA stores a distinct KV tensor for every head, resulting in a cache shape of (B, num_heads, T, head_dim).

GQA stores only one tensor per KV-group with shape (B, num_kv_groups, T, head_dim). During the forward pass, these shared tensors are expanded to match the full head count using repeat_interleave:

keys   = keys_base.repeat_interleave(self.group_size, dim=1)   # (B, num_heads, T, head_dim)

values = values_base.repeat_interleave(self.group_size, dim=1)

For a model with 12 heads and 2 KV-groups, the GQA approach reduces KV-cache size by approximately 6× while preserving the same query diversity.

Source Code Implementation Details

The rasbt/LLMs-from-scratch repository provides complete, runnable implementations that share identical interfaces despite their internal architectural differences.

MultiHeadAttention Class Structure

The baseline MHA implementation uses simple linear projections followed by reshaping without repetition:


# Inside ch04/04_gqa/gpt_with_kv_mha.py

class MultiHeadAttention(nn.Module):
    def __init__(self, d_in, d_out, dropout, num_heads, qkv_bias=False):
        super().__init__()
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.out_proj = nn.Linear(d_out, d_out)
        # ... additional initialization

Each head receives unique key and value tensors, consuming O(num_heads) memory complexity for the KV cache.

GroupedQueryAttention Class Structure

The GQA version introduces group-based sharing and explicit cache management buffers:


# Inside ch04/04_gqa/gpt_with_kv_gqa.py

class GroupedQueryAttention(nn.Module):
    def __init__(self, d_in, d_out, dropout, num_heads,
                 num_kv_groups, dtype=None, qkv_bias=False):
        super().__init__()
        assert d_out % num_heads == 0
        assert num_heads % num_kv_groups == 0

        self.head_dim = d_out // num_heads
        self.num_heads = num_heads
        self.num_kv_groups = num_kv_groups
        self.group_size = num_heads // num_kv_groups

        # Shared KV projections per group

        self.W_key   = nn.Linear(d_in, num_kv_groups * self.head_dim,
                                 bias=qkv_bias, dtype=dtype)
        self.W_value = nn.Linear(d_in, num_kv_groups * self.head_dim,
                                 bias=qkv_bias, dtype=dtype)
        
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias, dtype=dtype)
        
        # KV-cache buffers for autoregressive generation

        self.register_buffer("cache_k", None, persistent=False)
        self.register_buffer("cache_v", None, persistent=False)

Both classes integrate into the same TransformerBlock interface, allowing you to swap attention mechanisms by simply changing the imported class.

Practical Implementation Examples

You can instantiate either architecture using the same GPTModel wrapper with configuration-specific parameters.

Standard Multi-Head Attention Model


# File: ch04/04_gqa/gpt_with_kv_mha.py

from gpt_with_kv_mha import GPTModel

config = {
    "vocab_size": 50257,
    "context_length": 256,
    "emb_dim": 768,
    "n_heads": 12,
    "n_layers": 6,
    "drop_rate": 0.1,
    "qkv_bias": False,
}
model = GPTModel(config).to("cpu")

# Generate without KV-cache

output_ids = model.generate_text_simple_cached(
    model, 
    torch.tensor([[1, 2, 3]], dtype=torch.long), 
    max_new_tokens=20
)

Grouped-Query Attention Configuration


# File: ch04/04_gqa/gpt_with_kv_gqa.py

from gpt_with_kv_gqa import GPTModel

config = {
    "vocab_size": 50257,
    "context_length": 256,
    "emb_dim": 768,
    "n_heads": 12,
    "n_kv_groups": 2,      # Fewer groups → memory savings

    "n_layers": 6,
    "drop_rate": 0.1,
    "qkv_bias": False,
}
model = GPTModel(config).to("cpu")

# Generate with KV-cache enabled (optimized for autoregressive generation)

output_ids = model.generate_text_simple_cached(
    model, 
    torch.tensor([[1, 2, 3]], dtype=torch.long),
    max_new_tokens=20, 
    use_cache=True
)

Both examples utilize the generate_text_simple_cached function, which handles cache reset, prompt processing, and token-by-token generation. When use_cache=True, the GQA implementation stores only the compressed group tensors, reducing memory bandwidth during inference.

Summary

  • Parameter Efficiency: GQA uses O(num_kv_groups) KV matrices versus O(num_heads) in standard MHA, directly reducing model size.
  • Memory Optimization: The KV-cache stores (B, num_kv_groups, T, head_dim) tensors instead of full head-sized tensors, cutting memory usage proportionally to num_heads / num_kv_groups.
  • Implementation Simplicity: Both attention types in rasbt/LLMs-from-scratch share identical interfaces, allowing seamless swapping via the GroupedQueryAttention and MultiHeadAttention classes.
  • Computational Pattern: GQA employs repeat_interleave to broadcast shared keys and values to multiple heads without duplicating parameters.
  • Practical Impact: For large-scale deployments, GQA achieves comparable expressivity to MHA while significantly reducing memory pressure during autoregressive generation.

Frequently Asked Questions

What is the main advantage of GQA over standard Multi-Head Attention?

Grouped-Query Attention primarily reduces KV-cache memory consumption by sharing key and value projections across groups of heads. According to the source code in ch04/04_gqa/gpt_with_kv_gqa.py, this reduces cache storage from (B, num_heads, T, head_dim) to (B, num_kv_groups, T, head_dim), which is particularly beneficial for long-context inference in large language models.

How does repeat_interleave enable parameter sharing in GQA?

The repeat_interleave function expands the compressed KV-group tensors to match the full head count during the forward pass. As implemented in the GroupedQueryAttention class, keys and values are projected once per group, then repeated across group_size heads using keys_base.repeat_interleave(self.group_size, dim=1), allowing multiple heads to attend to identical KV data without storing duplicate parameters.

Can I easily switch between MHA and GQA in the codebase?

Yes. Both the MultiHeadAttention class in gpt_with_kv_mha.py and the GroupedQueryAttention class in gpt_with_kv_gqa.py implement the same interface, accepting att(x, use_cache=...) calls within the TransformerBlock. Swapping between them requires only changing the imported attention class and adding the n_kv_groups parameter to your configuration dictionary.

When should I choose GQA over MHA for production models?

Choose GQA when memory efficiency and inference speed are critical, particularly for deployments with long-context windows or limited GPU memory, as it maintains model quality while reducing KV-cache overhead. Use MHA when maximum head independence is required and memory constraints are less restrictive, though modern implementations like LLaMA demonstrate that GQA achieves comparable performance with significantly fewer resources.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →