online softmax论文解读


softmax进化

Naive Softmax

  1. d₀ ← 0
  2. for j = 1 to V do
  3.     dⱼ ← dⱼ₋₁ + e^(xⱼ)
  4. end for
  5. for i = 1 to V do
  6.     yᵢ ← e^(xᵢ) / d_V
  7. end for

Safe spftmax

  1. m₀ ← −∞
  2. for k = 1 to V do
  3.     mₖ ← max(mₖ₋₁, xₖ)
  4. end for
  5. d₀ ← 0
  6. for j = 1 to V do
  7.     dⱼ ← dⱼ₋₁ + e^(xⱼ − m_V)
  8. end for
  9. for i = 1 to V do
  10.     yᵢ ← e^(xᵢ − m_V) / d_V
  11. end for

Safe softmax with online normalizer calculation

  1. m₀ ← −∞
  2. d₀ ← 0
  3. for j = 1 to V do
  4.     mⱼ ← max(mⱼ₋₁, xⱼ)
  5.     dⱼ ← dⱼ₋₁ × e^(mⱼ₋₁ − mⱼ) + e^(xⱼ − mⱼ)
  6. end for
  7. for i = 1 to V do
  8.     yᵢ ← e^(xᵢ − m_V) / d_V
  9. end for

Online softmax and top-k

  1. m₀ ← −∞
  2. d₀ ← 0
  3. u ← {−∞, −∞, …, −∞}ᵀ , u ∈ ℝ^(K+1) ⊲ 前 K 个元素保存当前 TopK 值
  4. p ← {−1, −1, …, −1}ᵀ , p ∈ ℤ^(K+1) ⊲ 保存对应索引
  5. for j = 1 to V do
  6.     mⱼ ← max(mⱼ₋₁, xⱼ)
  7.     dⱼ ← dⱼ₋₁ × e^(mⱼ₋₁ − mⱼ) + e^(xⱼ − mⱼ)
  8.     uₖ₊₁ ← xⱼ ⊲ 将当前值放入第 K+1 个位置
  9.     pₖ₊₁ ← j ⊲ 保存其索引
  10.     k ← K ⊲ 对 u 做降序插入排序(前 K 个已有序)
  11.     while k ≥ 1 and uₖ < uₖ₊₁ do
  12.         swap(uₖ, uₖ₊₁)
  13.         swap(pₖ, pₖ₊₁)
  14.         k ← k − 1
  15.     end while
  16. end for
  17. for i = 1 to K do
  18.     vᵢ ← e^(uᵢ − m_V) / d_V
  19.     zᵢ ← pᵢ
  20. end for

补充

作者在论文中还证明了online safe softmax可以在gpu中并行运行


文章作者: Austin
版权声明: 本博客所有文章除特別声明外,均采用 CC BY 4.0 许可协议。转载请注明来源 Austin !
评论
  目录