What Is QK Normalization in NanoChat's Attention Mechanism?
NanoChat applies RMS normalization to queries and keys immediately after rotary position embeddings and before attention computation, replacing the traditional $1/\sqrt{d}$ scaling with unit-magnitude vectors and an empirical 1.15 sharpening factor.
Karpathy's NanoChat repository implements a modern transformer variant that stabilizes attention computation through query-key (QK) normalization. This technique, implemented in the CausalSelfAttention module of nanochat/gpt.py, normalizes Q and K vectors to unit root-mean-square (RMS) magnitude before the attention dot product. The approach proves particularly effective for training stability across mixed-precision environments including bfloat16 and FP8 hardware.
The RMS-Norm Implementation
NanoChat defines a lightweight normalization helper in nanochat/gpt.py that leverages PyTorch's native RMS normalization function. Unlike LayerNorm, which centers and scales data, RMSNorm maintains only the scaling component, forcing each vector to unit RMS magnitude.
# nanochat/gpt.py – RMS-Norm helper (lines 42-44)
def norm(x):
return F.rms_norm(x, (x.size(-1),)) # unit RMS per token
This norm function operates on the last dimension of the input tensor, ensuring each token's embedding vector possesses uniform magnitude regardless of its position in the sequence.
Application Inside CausalSelfAttention
The normalization occurs within the CausalSelfAttention.forward method immediately after rotary position embeddings are applied to the projected queries and keys. According to the source code, the implementation applies norm to both tensors, followed by a multiplicative sharpening factor.
# nanochat/gpt.py – Q-K normalization (lines 99-101)
q, k = norm(q), norm(k) # QK norm
q = q * 1.15 # optional sharpening factor
k = k * 1.15
The 1.15 sharpening factor empirically sharpens the attention distribution without destabilizing training, effectively increasing the contrast between high and low attention weights while the underlying RMS normalization eliminates variance-related numerical instability.
Why Normalize Queries and Keys?
Traditional transformer implementations scale the dot-product attention by $1/\sqrt{d}$ to prevent softmax saturation when dimensionality increases. However, this approach presents several drawbacks that QK normalization addresses:
- Precision sensitivity: The $1/\sqrt{d}$ scaling becomes unstable in mixed-precision training (bfloat16/FP8) where dynamic ranges are limited.
- Variance dependence: Raw Q/K vectors may exhibit varying magnitudes across layers or heads, making static scaling suboptimal.
- Simplified implementation: By forcing unit RMS magnitude before the attention computation, the model removes the need for hand-crafted scaling factors and enables consistent behavior across different model sizes and head dimensions.
The normalization ensures that q @ k.T operates on standardized vectors, allowing the model to share a single scaling strategy across all attention heads and layers regardless of the underlying hardware precision.
Integration with Flash Attention
After QK normalization, tensors flow into NanoChat's unified flash-attention wrapper without additional processing. The nanochat/flash_attention.py module consumes the pre-normalized Q and K tensors, dispatching to either Flash Attention 3 kernels or PyTorch's scaled dot-product attention (SDPA) fallback.
# nanochat/flash_attention.py (lines 19-27)
y = flash_attn.flash_attn_func(q, k, v, causal=True, window_size=window_size)
This architecture decouples the normalization logic from the attention kernel implementation, allowing the normalized tensors to benefit from optimized memory-efficient attention without modification to the flash attention internals.
Inspecting QK Normalization in Practice
You can verify the normalization behavior by inspecting the RMS magnitude of Q and K vectors during a forward pass:
import torch
from nanochat.gpt import GPT, GPTConfig, norm, COMPUTE_DTYPE
# Build a tiny model (depth 2 for demonstration)
config = GPTConfig(n_layer=2, n_head=4, n_kv_head=4, n_embd=128, sequence_len=16)
model = GPT(config)
# Dummy input (batch=1, seq_len=4)
x = torch.randn(1, 4, config.n_embd, dtype=COMPUTE_DTYPE)
# Grab the first block's attention module
attn = model.transformer["h"][0].attn
# Project to Q and K (pre-rotary)
q = attn.c_q(x).view(1, 4, attn.n_head, attn.head_dim)
k = attn.c_k(x).view(1, 4, attn.n_kv_head, attn.head_dim)
# Apply rotary (skipped here) then Q-K norm
q_norm = norm(q)
k_norm = norm(k)
# RMS of each vector should be ≈1
print("Mean RMS of Q:", q_norm.norm(dim=-1).mean().item())
print("Mean RMS of K:", k_norm.norm(dim=-1).mean().item())
Running this snippet outputs values very close to 1.0, confirming that norm successfully constrains the vectors to unit RMS magnitude as implemented in nanochat/gpt.py.
Summary
- Location: QK normalization is implemented in
nanochat/gpt.pywithin theCausalSelfAttentionclass, specifically after rotary embeddings and before attention computation. - Mechanism: Uses
F.rms_normvia thenormhelper function to force unit RMS magnitude on Q and K vectors. - Sharpening: Applies a 1.15 multiplicative factor to normalized Q and K tensors to sharpen attention distributions.
- Benefits: Eliminates the need for $1/\sqrt{d}$ scaling, improves stability in bfloat16/FP8 precision, and standardizes attention behavior across model configurations.
- Compatibility: Works seamlessly with both Flash Attention kernels and PyTorch SDPA fallbacks through
nanochat/flash_attention.py.
Frequently Asked Questions
What is the difference between QK Normalization and traditional attention scaling?
Traditional transformers scale the dot product $QK^T$ by $1/\sqrt{d}$ to prevent softmax saturation, while QK normalization (as in NanoChat) applies RMS normalization to Q and K vectors individually before the dot product, removing the variance dependence entirely. This approach eliminates the need for dimension-based scaling and proves more stable in low-precision training environments.
Why does NanoChat use a 1.15 sharpening factor?
The 1.15 factor empirically increases the sharpness of attention weights after RMS normalization, effectively amplifying the contrast between attended and ignored tokens. Without this factor, the unit-normalized vectors might produce overly diffuse attention distributions; the 1.15 multiplier restores selectivity while maintaining the numerical stability benefits of normalization.
How does QK normalization affect Flash Attention compatibility?
QK normalization occurs before the flash attention call in CausalSelfAttention.forward, meaning Flash Attention receives pre-normalized tensors exactly as it would receive scaled tensors in other implementations. Since nanochat/flash_attention.py passes Q, K, and V directly to the underlying kernel without rescaling, the normalization integrates transparently with both Flash Attention 3 and PyTorch SDPA backends.
Where is the norm function defined in the NanoChat codebase?
The norm function is defined in nanochat/gpt.py at lines 42-44 as a standalone helper using torch.nn.functional.rms_norm. This function is imported and utilized within the CausalSelfAttention class to normalize queries and keys, and can be imported directly for testing or custom attention implementations as shown in the COMPUTE_DTYPE example above.
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 →