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_headslower thannum_attention_headsinMiniMindConfig. - Replication mechanism: The
n_repfactor calculated inAttention.__init__determines how many query heads share each KV head. - Efficient broadcasting: The
repeat_kvfunction inmodel/model_minimind.pyexpands 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →