LogoCookLLM Docs
LogoCookLLM Docs
HomeCookLLM

Principles

Tokenization
Tokenization BasicsBPE AlgorithmGPT TokenizersBPE Training Engineering
Model Architecture
Transformer LM
From token ids to logitsEmbedding and LM Head
Attention Mechanisms
From Self-Attention to GQAAttention Sink
Position Encoding
Position Encoding BasicsRoPE Math DerivationRoPE ImplementationLength Extrapolation
GPU Programming Basics
GPU Architecture BasicsTensor LayoutTriton Basics: Vector Add
FlashAttention
Flash Attention PrinciplesFrom Naive Implementation to Auto-TuningBlock Pointers and Multi-Dim SupportCausal Masking OptimizationGrouped Query AttentionBackward Pass Implementation
Distributed Training
Data ParallelismZeRO OptimizerFully Sharded Data ParallelTensor ParallelismPipeline ParallelismMulti-Dimensional Hybrid Parallelism

Hands-on Training

Overview
Pretraining
Pretraining DataTokenizer TrainingModel ArchitectureData PipelineTraining LoopMonitoring and Validation
X (Twitter)
SystemsFlashAttention

Grouped Query Attention

Premium

Add GQA/MQA support so multiple query heads share KV, reducing KV cache memory.

Get code access

In previous chapters, we built a full Flash Attention kernel with arbitrary sequence length, multi-dim parallelism, and causal masking. Now we add the final key feature: Grouped Query Attention (GQA).

This is standard in Llama 2/3, Mistral, Falcon, etc. It can reduce KV cache memory by 4–8× with minimal quality loss.

Quick GQA Recap

GQA shares Key/Value heads across multiple Query heads:

Log in to continue reading

This is premium content. Please log in to access the full article.

# MHA: each Q head has its own K/V head
Q: (B, 8, N, D)
K: (B, 8, N, D)
V: (B, 8, N, D)

# GQA: multiple Q heads share KV heads (groups=4)
Q: (B, 8, N, D)
K: (B, 2, N, D)  # 2 KV heads → 4x memory reduction



Head index mapping:

kv_head_idx(hq)=⌊hqHQ/HKV⌋\text{kv\_head\_idx}(h_q) = \left\lfloor \frac{h_q}{H_Q / H_{KV}} \right\rfloorkv_head_idx(hq​)=⌊HQ​/H

For concept, history, and tradeoffs, see Attention Mechanisms. This chapter focuses on zero-copy, zero-extra-memory GQA in Flash Attention.

Causal Masking Optimization

Implement causal attention for autoregressive models, achieving ~2x speedup by skipping the upper-triangular computation.

Backward Pass Implementation

Implement Flash Attention gradient computation, achieving memory-efficient training through recomputation.

Table of Contents

Quick GQA Recap
Problem with Standard PyTorch GQA
Efficient Flash Attention Implementation
Core Idea: Pointer Indexing, Not Data Copy
Concrete Example
Unified Support: MHA/GQA/MQA
Full Implementation
Performance Validation
Correctness
Performance
Autotune Best Config
Tradeoffs and Recommendations
Quality vs Memory
Key Implementation Summary
Summary
V: (B, 2, N, D)
# Q heads 0-3 → KV head 0
# Q heads 4-7 → KV head 1
KV
​
hq​
​
⌋