How Grouped Query Attention (GQA) Functions with n_kv_heads in Llama
Grouped Query Attention (GQA) reduces memory and compute during inference by sharing key/value heads among multiple query heads when n_kv_heads is set lower than n_heads in the Llama 2 architecture.
In the meta-llama/llama repository, the n_kv_heads parameter controls whether the model uses classic Multi-Head Attention (MHA) or the more efficient Grouped Query Attention (GQA) mechanism. This implementation is particularly critical for the 70B parameter model, where reducing KV cache size significantly improves inference scalability.
The Role of n_kv_heads in ModelArgs
The configuration for GQA begins in the model definition file. In llama/model.py, the ModelArgs dataclass exposes n_kv_heads as an optional integer that overrides the default symmetry between query and key/value heads.
Default Behavior vs. GQA Activation
When n_kv_heads is omitted from the configuration, the model defaults to standard multi-head attention. The Attention class constructor resolves the effective KV head count with the following logic:
self.n_kv_heads = args.n_heads if args.n_kv_heads is None else args.n_kv_heads # lines 200-201
If you supply a value smaller than n_heads (for example, 8 KV heads against 32 query heads), the system activates GQA. Each KV head then serves multiple query heads, reducing the dimensionality of the key and value projections while maintaining the full query capacity.
Internal Mechanics of GQA
Under model parallelism, Llama 2 distributes attention heads across multiple devices. The GQA mechanism requires careful calculation of how many query heads share each KV head on every device.
Head Distribution and the Replication Factor
The constructor computes local head counts and the replication ratio:
self.n_local_heads = args.n_heads // model_parallel_size # line 202
self.n_local_kv_heads = self.n_kv_heads // model_parallel_size # line 203
self.n_rep = self.n_local_heads // self.n_local_kv_heads # line 204
The self.n_rep variable defines the group size—specifically, how many local query heads share a single local KV head. For instance, with 32 total query heads and 8 KV heads running on a single device, n_rep equals 4.
Reduced Linear Projections
The linear transformation layers reflect the head count asymmetry:
- Queries (
wq): Project ton_heads * head_dimdimensions - Keys and Values (
wk,wv): Project ton_kv_heads * head_dimdimensions
This reduction means the KV cache stores fewer vectors than the query tensor, directly decreasing memory bandwidth requirements during autoregressive generation.
The repeat_kv Function
Before the attention score calculation, the system expands the compressed KV tensors to match the query head count. The helper function repeat_kv (defined at lines 64-73 in llama/model.py) tiles each KV head n_rep times along the head dimension:
keys = repeat_kv(keys, self.n_rep) # lines 291-292
values = repeat_kv(values, self.n_rep) # lines 293-294
After replication, both keys and values have shape (batch, seq_len, n_local_heads, head_dim), identical to the query tensor shape. This allows the subsequent torch.matmul operations to proceed without modification to the core attention algorithm.
Attention Computation Flow
Following replication, the attention mechanism executes standard scaled dot-product attention. The softmax and output projections operate identically to MHA, but the underlying KV representations remain shared among query groups. This architectural choice preserves model quality while reducing the KV cache memory footprint by a factor of n_heads / n_kv_heads.
Configuration Example
To instantiate a Llama 2 model with GQA enabled, specify n_kv_heads in the model arguments:
from llama.model import ModelArgs, Transformer
# 32 query heads with 8 KV heads creates 4 query heads per KV group
args = ModelArgs(
dim=4096,
n_layers=32,
n_heads=32, # Total query heads
n_kv_heads=8, # Total KV heads (GQA active)
vocab_size=32000,
max_seq_len=2048,
)
model = Transformer(args)
In this configuration, each of the 8 KV heads handles attention weights for 4 distinct query heads internally, cutting the KV cache size by 75% compared to standard MHA.
Summary
n_kv_headsinllama/model.py(line 24) activates GQA when set lower thann_heads- The
Attentionclass calculatesn_rep(lines 200-204) to determine how many query heads share each KV head - Linear projections for keys and values use the reduced
n_kv_headscount, saving memory - The
repeat_kvfunction (lines 64-73) expands KV tensors during the forward pass to maintain compatibility with existing attention math - GQA reduces KV cache memory by a factor of
n_heads / n_kv_headswithout requiring changes to the core attention computation logic
Frequently Asked Questions
What happens if I set n_kv_heads equal to n_heads?
When n_kv_heads equals n_heads, the model operates as standard Multi-Head Attention. The n_rep variable becomes 1, meaning each query head has its own dedicated KV head. This is the default behavior when n_kv_heads is left as None in the configuration.
How much memory does GQA save compared to standard attention?
GQA reduces the key/value cache memory footprint proportionally to the ratio of total query heads to KV heads. For example, with n_heads=32 and n_kv_heads=8, the KV cache uses 75% less memory than standard MHA, since only 8 KV vectors are stored and reused across 32 query heads.
Does GQA affect model accuracy or training requirements?
According to the Llama 2 implementation, GQA is designed to maintain model quality while improving inference efficiency. The mechanism is primarily beneficial for large models (70B parameters) where KV cache memory constraints are severe. The training dynamics differ from MHA, but the inference-time n_kv_heads parameter specifically optimizes memory bandwidth during generation without altering the training architecture.
Where is the KV head replication logic implemented in the source code?
The replication occurs in llama/model.py within the Attention.forward method at lines 291-294, where repeat_kv is called on both keys and values. The helper function itself resides at lines 64-73, using torch.repeat_interleave to expand the head dimension by the n_rep factor calculated during initialization.
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 →