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

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

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:

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, the repeat_kv function expands each KV head to match the query head count:

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:

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

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:

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:

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 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.

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 →