/
AI Flashcards
Save to my account
Sign up
AI Flashcards
FlashAttention-3 essentials
Study
1
Question
What is incoherent processing in FlashAttention-3?
Page 1
Answer
Multiplying Q and K with a random orthogonal matrix M before FP8 quantization to even out outliers.
2
Question
Why multiply Q and K with matrix M in incoherent processing?
Page 1
Answer
To spread out outliers, reducing quantization error in FP8.
3
Question
What property ensures multiplying Q and K with M doesn't change attention output?
Page 1
Answer
M is orthogonal, so M M^T = I, thus (Q M)(K M)^T = Q K^T.
4
Question
How does incoherent processing reduce quantization error?
Page 1
Answer
Each entry in Q M or K M is a random sum of original entries, spreading outlier impact.
5
Question
What is the choice of M in incoherent processing?
Page 1
Answer
Product of random diagonal matrices of ±1 and a Hadamard matrix.
6
Question
What is the time complexity of multiplying with the chosen M?
Page 1
Answer
O(d log d), faster than O(d^2).
7
Question
Can M multiplication fuse with rotary embedding?
Page 1
Answer
Yes, at no extra computation cost.
8
Question
By how much do the techniques reduce numerical error?
Page 1
Answer
Up to 2.6x, validated in section 4.3.
9
Question
What primitives from CUTLASS are used in FlashAttention-3?
Page 1
Answer
WGMMA and TMA abstractions for implementation.
10
Question
What is FlashAttention-3 compared to in benchmarking?
Page 1
Answer
PyTorch standard, FlashAttention-2, Triton FA-2, cuDNN FA-2 on H100.
11
Question
How much faster is FlashAttention-3 than FlashAttention-2?
Page 1
Answer
Up to 2.0x faster.
12
Question
How does FlashAttention-3 compare to Triton FA-2?
Page 1
Answer
1.5x faster.
13
Question
What peak performance does FlashAttention-3 achieve?
Page 1
Answer
Up to 740 TFLOPs/s, 75% of H100 theoretical max.
14
Question
What contributes to FlashAttention-3 speedup in ablation?
Page 1
Answer
Warp-specialization and GEMM-softmax pipelining.
15
Question
How much error reduction from FP8 techniques?
Page 1
Answer
2.6x reduction in numerical error.
16
Question
What GPU is used for benchmarking?
Page 1
Answer
H100 80GB SXM5.
17
Question
What input precision for main benchmarks?
Page 1
Answer
FP16 inputs, with/without causal mask, head dims 64 or 128.
18
Question
How much faster is FA-3 forward pass than FA-2?
Page 1
Answer
1.5-2.0x faster.
19
Question
Backward pass speedup of FA-3 over FA-2?
Page 1
Answer
1.5-1.75x faster.
20
Question
Compared to standard PyTorch attention, FA-3 speedup?
Page 1
Answer
Up to 3-16x faster.
21
Question
When does FA-3 surpass cuDNN FA-2?
Page 1
Answer
For medium and long sequences (1k and above).
22
Question
Benchmark sequence lengths varied?
Page 1
Answer
512, 1k, ..., 16k.
23
Question
How is batch size set in benchmarks?
Page 1
Answer
So total tokens = 16k.
24
Question
Hidden dimension in benchmarks?
Page 1
Answer
2048.
25
Question
Head dimensions tested?
Page 1
Answer
64, 128, 256 (32, 16, 8 heads).
26
Question
Forward pass FLOPs formula?
Page 1
Answer
4 * seqlen^2 * head_dim * num_heads.
27
Question
Adjustment for causal masking in FLOPs?
Page 1
Answer
Divide by 2, as half the entries are computed.
28
Question
Backward pass FLOPs relative to forward?
Page 1
Answer
2.5x forward FLOPs.
29
Question
FP8 benchmarks focus on?
Page 1
Answer
Forward pass, head dim 256, results in Fig. 7 and Appendix C.2.
30
Question
Ablation parameters fixed?
Page 1
Answer
batch=4, seqlen=8448, nheads=16, hdimg=128.