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

> Understand Grouped-Query Attention (GQA) vs Multi-Head Attention (MHA). Learn how GQA optimizes memory and KV-cache for LLMs, offering efficiency with comparable performance.

- Repository: [Sebastian Raschka/LLMs-from-scratch](https://github.com/rasbt/LLMs-from-scratch)
- Tags: deep-dive
- Published: 2026-05-12

---

**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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/ch04/04_gqa/gpt_with_kv_mha.py), the `MultiHeadAttention` class creates distinct linear layers for each head:

```python
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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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:

```python
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`:

```python
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:

```python

# 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:

```python

# 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

```python

# 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

```python

# 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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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`](https://github.com/rasbt/LLMs-from-scratch/blob/main/gpt_with_kv_mha.py) and the `GroupedQueryAttention` class in [`gpt_with_kv_gqa.py`](https://github.com/rasbt/LLMs-from-scratch/blob/main/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.