softmax进化
Naive Softmax
- d₀ ← 0
- for j = 1 to V do
- dⱼ ← dⱼ₋₁ + e^(xⱼ)
- end for
- for i = 1 to V do
- yᵢ ← e^(xᵢ) / d_V
- end for
Safe spftmax
- m₀ ← −∞
- for k = 1 to V do
- mₖ ← max(mₖ₋₁, xₖ)
- end for
- d₀ ← 0
- for j = 1 to V do
- dⱼ ← dⱼ₋₁ + e^(xⱼ − m_V)
- end for
- for i = 1 to V do
- yᵢ ← e^(xᵢ − m_V) / d_V
- end for
Safe softmax with online normalizer calculation
- m₀ ← −∞
- d₀ ← 0
- for j = 1 to V do
- mⱼ ← max(mⱼ₋₁, xⱼ)
- dⱼ ← dⱼ₋₁ × e^(mⱼ₋₁ − mⱼ) + e^(xⱼ − mⱼ)
- end for
- for i = 1 to V do
- yᵢ ← e^(xᵢ − m_V) / d_V
- end for
Online softmax and top-k
- m₀ ← −∞
- d₀ ← 0
- u ← {−∞, −∞, …, −∞}ᵀ , u ∈ ℝ^(K+1) ⊲ 前 K 个元素保存当前 TopK 值
- p ← {−1, −1, …, −1}ᵀ , p ∈ ℤ^(K+1) ⊲ 保存对应索引
- for j = 1 to V do
- mⱼ ← max(mⱼ₋₁, xⱼ)
- dⱼ ← dⱼ₋₁ × e^(mⱼ₋₁ − mⱼ) + e^(xⱼ − mⱼ)
- uₖ₊₁ ← xⱼ ⊲ 将当前值放入第 K+1 个位置
- pₖ₊₁ ← j ⊲ 保存其索引
- k ← K ⊲ 对 u 做降序插入排序(前 K 个已有序)
- while k ≥ 1 and uₖ < uₖ₊₁ do
- swap(uₖ, uₖ₊₁)
- swap(pₖ, pₖ₊₁)
- k ← k − 1
- end while
- end for
- for i = 1 to K do
- vᵢ ← e^(uᵢ − m_V) / d_V
- zᵢ ← pᵢ
- end for
补充
作者在论文中还证明了online safe softmax可以在gpu中并行运行