Loading video...

Video Failed to Load

Go Home

(1/5) FP4 hardware is here, but 4-bit attention still kills model quality, blocking true end-to-end FP4 serving. To fix that, we propose Attn-QAT, the first systematic study of quantization-aware training for attention. The result: FP4 attention quality is comparable to BF16 attention with 1.1x–1.5x higher throughput than SageAttention3 on...

38,252 views • 6 months ago •via X (Twitter)

12 Comments

Hao AI Lab's profile picture
Hao AI Lab6 months ago

(2/5)Naive QAT breaks when applied to FlashAttention kernels. We found two fixes are needed: 1. Store a small high-precision auxiliary output so the gradient computation stays mathematically consistent 2. Recompute attention probabilities in the backward pass using the same low precision as the forward pass These two changes stabilize 4-bit attention training.

Hao AI Lab's profile picture
Hao AI Lab6 months ago

(3/5) Across both video diffusion models and language models, Attn-QAT recovers the quality drop of 4-bit attention without the extra outlier-mitigation heuristics. For continued pretraining, Attn-QAT recovers most of the quality loss caused by FP4 attention on Qwen3-14B and partially recovers it on Llama 3.1-70B. For supervised fine-tuning, Attn-QAT can be used as a drop-in replacement for BF16 attention. On Qwen3-14B, it achieves nearly identical downstream benchmark performance to BF16 attention. On Llama 3.1-70B, it remains close with a small gap. For randomly-selected example videos (generated by Wan-2.1-14B), we see that with Attn-QAT, FP4 attention produces videos comparable to BF16 attention, whereas SageAttention3 produces videos with artifacts.

Hao AI Lab's profile picture
Hao AI Lab6 months ago

(4/5) Because the model learns to account for quantization error during QAT, inference needs no extra heuristics! No Q/K smoothing, no two-level quantization. Simpler kernel → faster inference compared to SageAttention3 on an RTX 5090.

Hao AI Lab's profile picture
Hao AI Lab6 months ago

(5/5) On a B200, Blackwell's tensor cores are so fast that the softmax now becomes a bottleneck in addition to the GEMMs. Quantizing PV adds scale-factor overhead that piles onto the softmax warps. So we run NVFP4 QK + BF16 PV, with a careful TMEM overlap schedule to fit scale factors into an already-full pipeline. Result: 1801 TFLOPS and up to 1.39x over FlashAttention-4. 2x/4x faster exp on B300/Rubin should push this further, and end-to-end FP4 serving, once blocked by attention quality, is now within reach.

Lee Penkman's profile picture
Lee Penkman6 months ago

dang i see this on PRs a lot :D tricky now theres all these agents to control what they are all doing, maybe they should set a githook or smth @grok whats best way to prevent this

Vipul kumar's profile picture
Vipul kumar6 months ago

@gork is this true man

Vipul kumar's profile picture
Vipul kumar6 months ago

@grok hi

vik's profile picture
vik6 months ago

nice work

David Güera's profile picture
David Güera6 months ago

Great work!

Cliff Lattner's profile picture
Cliff Lattner6 months ago

@ye_combinator Did you try FP8? It should be the same performance uplift as FP4

homuraakemifan's profile picture
homuraakemifan6 months ago

WOW!

Lee Penkman's profile picture
Lee Penkman6 months ago

friends from @fal might like this :)

Related Videos

New short course: Attention in Transformers: Concepts and Code in PyTorch. Last week we released a course on how LLM transformers work. This week, go deeper and learn about the technical ideas behind the attention mechanism, and see how to code it in PyTorch. This course is built with Joshua Starmer, Founder and CEO of StatQuest. The attention mechanism was a breakthrough that led to transformers, the architecture powering large language models like ChatGPT. Transformers, introduced in the 2017 paper: "Attention is All You Need" by Viswani and others, took off because of its highly scalable design. In this course, you’ll learn how the attention mechanism, a key element of transformer-based LLMs, works and implement it in PyTorch. You'll develop deep intuition about building reliable, functional, and scalable AI applications. What you will do: - Understand the evolution of the attention mechanism, a key breakthrough that led to transformers. - Learn the relationships between word embeddings, positional embeddings, and attention. - Learn about the Query, Key, and Value matrices, and how to produce and use them in attention. - Walk through the math required to calculate self-attention and masked self-attention to learn why and how they work. - Understand the difference between self-attention and masked self-attention and how one is used in the encoder to build context-aware embeddings and the other is used in the decoder for generative outputs. - Learn the details of the encoder-decoder architecture, cross-attention, and multi-head attention and how they are all incorporated into a transformer. - Use PyTorch to code a class that implements self-attention, masked self-attention, and multi-head attention. There're lots of exciting technical details in this course. Please sign up here:

Andrew Ng

132,544 views • 1 year ago

[Self-Attention] by Hand ✍️ Self-attention is what enables LLMs to understand context. How does it work? This exercise demonstrates how to calculate a 6-3 attention head by hand. Note that if we have two instances of this, we get 6-6 attention (i.e., multi-head attention, n=2). -- 𝗚𝗼𝗮𝗹 -- Transform [6D Features 🟧] to [3D Attention Weighted Features 🟦] -- 𝗪𝗮𝗹𝗸𝘁𝗵𝗿𝗼𝘂𝗴𝗵 -- [1] Given ↳ A set of 4 feature vectors (6-D): x1,x2,x3,x4 [2] Query, Key, Value ↳ Multiply features x's with linear transformation matrices WQ, WK, and WV, to obtain query vectors (q1,q2,q3,q4), key vectors (k1,k2,k3,k4), and value vectors (v1,v2,v3,v4). ↳ "Self" refers to the fact that both queries and keys are derived from the same set of features. [3] 🟪 Prepare for MatMul ↳ Copy query vectors ↳ Copy the transpose of key vectors [4] 🟪 MatMul ↳ Multiply K^T and Q ↳ This is equivalent to taking dot product between every pair of query and key vectors. ↳ The purpose is to use dot product as an estimate of the "matching score" between every key-value pair. ↳ This estimate makes sense because dot product is the numerator of Cosine Similarity between two vectors. [5] 🟨 Scale ↳ Scale each element by the square root of dk, which is the dimension of key vectors (dk=3). ↳ The purpose is to normalize the impact of the dk on matching scores, even if we scale dk to 32, 64, or 128. ↳ To simplify hand calculation, we approximate [ □/sqrt(3) ] with [ floor(□/2) ]. [6] 🟩 Softmax: e^x ↳ Raise e to the power of the number in each cell ↳ To simplify hand calculation, we approximate e^□ with 3^□. [7] 🟩 Softmax: ∑ ↳ Sum across each column [8] 🟩 Softmax: 1 / sum ↳ For each column, divide each element by the column sum ↳ The purpose is normalize each column so that the numbers sum to 1. In other words, each column is a probability distribution of attention, and we have four of them. ↳ The result is the Attention Weight Matrix (A) (yellow) [9] 🟦 MatMul ↳ Multiply the value vectors (Vs) with the Attention Weight Matrix (A) ↳ The results are the attention weighted features Zs. ↳ They are fed to the position-wise feed forward network in the next layer.

Tom Yeh

101,225 views • 2 years ago