# How MiniMind Implements Grouped-Query Attention (GQA): A Deep Dive into the Source Code

> Explore how MiniMind implements Grouped-Query Attention GQA by adjusting key-value heads and replicating them. Dive into the source code for a technical deep dive.

- Repository: [jingyaogong/minimind](https://github.com/jingyaogong/minimind)
- Tags: deep-dive
- Published: 2026-03-24

---

**MiniMind implements Grouped-Query Attention by allowing `num_key_value_heads` to be smaller than `num_attention_heads`, then replicating KV heads via the `repeat_kv` function to match query head count during the attention computation.**

Grouped-Query Attention (GQA) reduces memory bandwidth and computational cost during inference by sharing key and value representations across multiple query heads. In the MiniMind repository, this efficiency technique is implemented through a flexible configuration system that decouples the number of query heads from key/value heads. This article examines the specific implementation details in [`model/model_minimind.py`](https://github.com/jingyaogong/minimind/blob/main/model/model_minimind.py), showing how the architecture handles head grouping and replication.

## Configuring GQA in MiniMind

MiniMind exposes GQA controls through the `MiniMindConfig` dataclass, allowing users to specify different head counts for queries versus keys and values.

### The MiniMindConfig Parameters

The configuration relies on two critical fields in [`model/model_minimind.py`](https://github.com/jingyaogong/minimind/blob/main/model/model_minimind.py):

```python
num_attention_heads: int = 8,
num_key_value_heads: int = 8,   # Set smaller than num_attention_heads to enable GQA

```

When `num_key_value_heads` is set to a value smaller than `num_attention_heads`, the model automatically activates Grouped-Query Attention. For example, setting `num_attention_heads=8` and `num_key_value_heads=2` creates four query heads per key/value head, significantly reducing the KV cache memory footprint.

## Core Implementation Details

The GQA mechanism centers on calculating a replication factor and expanding KV tensors to match query head dimensions during the forward pass.

### Head Replication Factor Calculation

In the `Attention.__init__` method (lines 150-158), MiniMind computes how many query heads share each KV head:

```python
self.num_key_value_heads = args.num_attention_heads if args.num_key_value_heads is None else args.num_key_value_heads
assert args.num_attention_heads % self.num_key_value_heads == 0
self.n_local_heads = args.num_attention_heads
self.n_local_kv_heads = self.num_key_value_heads
self.n_rep = self.n_local_heads // self.n_local_kv_heads      # Groups per KV head

```

The `n_rep` variable stores the replication factor. The assertion ensures that query heads divide evenly into KV head groups, which is required for the grouping logic to function correctly.

### The repeat_kv Helper Function

Located at lines 140-147 in [`model/model_minimind.py`](https://github.com/jingyaogong/minimind/blob/main/model/model_minimind.py), the `repeat_kv` function expands each KV head to match the query head count:

```python
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
    bs, slen, num_key_value_heads, head_dim = x.shape
    if n_rep == 1:
        return x
    return (
        x[:, :, :, None, :].expand(bs, slen, num_key_value_heads,
                                 n_rep, head_dim)
         .reshape(bs, slen, num_key_value_heads * n_rep, head_dim)
    )

```

This function uses PyTorch's `expand` and `reshape` operations to broadcast each KV head `n_rep` times without allocating additional memory for the expanded tensor until necessary.

### Forward Pass Integration

During the `Attention.forward` method (lines 190-194), the query, key, and value projections are transformed and the KV tensors are repeated:

```python
xq, xk, xv = (
    xq.transpose(1, 2),
    repeat_kv(xk, self.n_rep).transpose(1, 2),
    repeat_kv(xv, self.n_rep).transpose(1, 2)
)

```

After replication, `xk` and `xv` have the same number of heads as `xq`, allowing standard scaled dot-product attention to proceed. The rest of the attention algorithm—including optional FlashAttention, masking, and scaling—remains unchanged from standard multi-head attention.

## Practical Code Examples

### Instantiating a MiniMind Model with GQA

To create a model with 8 query heads and 2 key/value heads (4:1 grouping ratio):

```python
from model.model_minimind import MiniMindConfig, MiniMindForCausalLM

cfg = MiniMindConfig(
    num_attention_heads=8,
    num_key_value_heads=2,   # Activates Grouped-Query Attention

    hidden_size=512,
    vocab_size=6400,
)

model = MiniMindForCausalLM(cfg)

```

### Running Inference with GQA Enabled

The model automatically applies GQA during forward passes:

```python
import torch

input_ids = torch.randint(0, cfg.vocab_size, (1, 16))   # Batch-size 1, sequence length 16

outputs = model(input_ids, use_cache=False)

logits = outputs.logits               # Shape: (1, 16, vocab_size)

```

### Verifying the Replication Factor

You can inspect the grouping ratio at runtime:

```python
print(model.model.layers[0].self_attn.n_rep)   # Output: 4 (8 heads / 2 KV heads)

```

## Summary

- **Configuration-driven**: GQA is activated by setting `num_key_value_heads` lower than `num_attention_heads` in `MiniMindConfig`.
- **Replication mechanism**: The `n_rep` factor calculated in `Attention.__init__` determines how many query heads share each KV head.
- **Efficient broadcasting**: The `repeat_kv` function in [`model/model_minimind.py`](https://github.com/jingyaogong/minimind/blob/main/model/model_minimind.py) expands KV tensors to match query dimensions without unnecessary memory allocation.
- **Transparent integration**: After KV expansion, the standard attention computation proceeds unchanged, supporting optional optimizations like FlashAttention.

## Frequently Asked Questions

### What is Grouped-Query Attention and why does MiniMind use it?

Grouped-Query Attention is an optimization where multiple query heads share the same key and value head projections, reducing the memory bandwidth required for KV cache during inference. MiniMind implements this to allow flexible trade-offs between model capacity (controlled by query heads) and computational efficiency (controlled by KV heads), particularly beneficial for deployment on resource-constrained devices.

### How do I enable GQA when training a MiniMind model?

Enable GQA by modifying the `MiniMindConfig` before model initialization. Set `num_attention_heads` to your desired total head count, then set `num_key_value_heads` to a divisor of that number (e.g., 8 query heads and 2 KV heads). The training scripts in the `trainer/` directory accept these configuration parameters, and the `Attention` layer automatically handles the grouping logic during both training and inference.

### What is the performance impact of using GQA in MiniMind?

Using GQA reduces the memory footprint and computational cost of the KV projections and cache by a factor of `n_rep` (the ratio of query heads to KV heads). For example, with 8 query heads and 2 KV heads, MiniMind uses 75% less memory for key and value storage compared to standard multi-head attention, while maintaining most of the model's representational capacity through the higher query head count.

### Does MiniMind's GQA implementation support FlashAttention?

Yes, MiniMind's GQA implementation is compatible with FlashAttention optimizations. The `repeat_kv` operation occurs before the attention computation, meaning the standard scaled dot-product attention (including FlashAttention variants) operates on tensors where KV heads have already been expanded to match query head counts. This allows the efficiency benefits of GQA to combine with the speed improvements of optimized attention kernels.