1. One matrix multiplication computes every query–key comparison
1. 一次矩阵乘法并行计算全部 query–key comparisons
Q ∈ ℝ3×dk,
K ∈ ℝ3×dk
Q ∈ ℝ3×dk
×
KT ∈ ℝdk×3
→
S = QKT/√dk ∈ ℝ3×3
S ∈ ℝ3×3: rows = queries, columns = keys
| k₁ | k₂ | k₃ |
| q₁ | 1.2 | 0.3 | −0.4 |
| q₂ | 0.0 | 1.0 | 1.5 |
| q₃ | 1.0 | 0.0 | 2.0 |
The upper-right scores are computed, but have not been allowed to carry information.
右上角的 future scores 虽已算出,但尚未被允许传递信息。
2. Add the causal mask before row-wise softmax
2. 在 row-wise softmax 之前加入 causal mask
S ∈ ℝ3×3
| 1.2 | 0.3 | −0.4 |
| 0.0 | 1.0 | 1.5 |
| 1.0 | 0.0 | 2.0 |
+
→
S + M ∈ ℝ3×3
| 1.2 | −∞ | −∞ |
| 0.0 | 1.0 | −∞ |
| 1.0 | 0.0 | 2.0 |
3. Masked softmax normalizes each query row independently
3. Masked softmax 分别 normalization 每一条 query row
A = softmaxrow(S + M) ∈ ℝ3×3
A ∈ ℝ3×3: attention weights
| v₁ | v₂ | v₃ |
| q₁ | 1.000 | 0 | 0 |
| q₂ | 0.269 | 0.731 | 0 |
| q₃ | 0.245 | 0.090 | 0.665 |
×
=
O = AV ∈ ℝ3×dv
o₁ = 1.000v₁
o₂ = 0.269v₁ + 0.731v₂
o₃ = 0.245v₁ + 0.090v₂ + 0.665v₃
S₁₃ existed temporarily; masking made A₁₃ = 0 before multiplication by V. Therefore q₁ receives no information from v₃.
S₁₃ 可以暂时存在;但在乘以 V 之前,mask 已使 A₁₃ = 0。因此 q₁ 不会从 v₃ 获得任何信息。