Online Softmax
Softmax
-
给定输入向量 \(\boldsymbol{x} = [x_1, x_2, ..., x_N]\), Softmax 函数将其映射为概率分布向量 \(\boldsymbol{y}\), 第 \(i\) 个输出 \(y_i\) 的计算公式为:
\[y_i = Softmax(x_i) = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}} \]- 其满足: \(0 \lt y_i \lt 1, \sum_{i=1}^{N} y_i = 1\).
Safe Softmax
-
平移不变性:
-
Softmax 函数对输入向量 \(\boldsymbol{x}\) 同时加上或减去一个常数 \(C\), 输出结果不变, 即:
\[Softmax(\boldsymbol{x}) = Softmax(\boldsymbol{x} - C) \] -
证明:
\[\frac{e^{x_i - C}}{\sum_{j=1}^{N} e^{x_j - C}} = \frac{e^{x_i}\cdot e^{-C}}{\sum_{j=1}^{N} (e^{x_j} \cdot e^{-C})} = \frac{e^{x_i}\cdot e^{-C}}{ e^{-C} \cdot \sum_{j=1}^{N} e^{x_j}} = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}} \]
-
-
直接计算 \(e^{x_i}\) 易导致上溢 (当 \(x_i\) 很大时), 基于平移不变性, 工程上通常取常数 \(C = max(\boldsymbol{x})\), 此时计算公式变形为
\[y_i = \frac{e^{x_i - max(\boldsymbol{x})}}{\sum_{j=1}^{N} e^{x_j - max(\boldsymbol{x})}} \]
Online Softmax
- 标准的 Safe Softmax 计算步骤需要三次读取输入数据 \(\boldsymbol{x}\): 找最大值 \(m\), 计算 \(d = \sum {e^{x_i - m}}\), 计算 \(y_i = \frac{e^{x_i - m}}{d}\).
- Online Softmax 的目标是尽可能少次读取输入数据 \(\boldsymbol{x}\), 同时计算出最大值、指数和、以及最终结果.
- 串行递推公式:
- 定义在 \(k\) 时刻(即处理完输入序列的前 \(k\) 个元素 \(x_1, ..., x_k\))的状态为 \((m_k, d_k)\):
- 局部最大值: \(m_k = max(x_1, ..., x_k)\)
- 局部指数和: \(d_k = \sum_{j=1}^{k} e^{x_j - m_k}\)
- 公式推导:
- 已有 \(k\) 时刻的状态 \((m_k, d_k)\), 现在输入新元素 \(x_{k+1}\), 需计算 \(k+1\) 时的状态 \((m_{k+1}, d_{k+1})\).
- 第一步: 更新最大值 \(m_{k+1}\)
- 显然 \(m_{k+1} = max(m_k, x_{k+1})\).
- 第二步: 计算 \(d_{k+1}\)
- 根据定义展开 \(d_{k+1}\):\[ \begin{aligned} d_{k+1} &= \sum_{j=1}^{k+1} e^{x_j - m_{k+1}} \\&= (\sum_{j=1}^{k} e^{x_j - m_{k+1}}) + e^{x_{k+1} - m_{k+1}} \\&= [\sum_{j=1}^{k} (e^{x_j - m_{k}} \cdot e^{m_{k} - m_{k+1}})] + e^{x_{k+1} - m_{k+1}} \\&= e^{m_{k} - m_{k+1}} \cdot \sum_{j=1}^{k} e^{x_j - m_{k}} + e^{x_{k+1} - m_{k+1}} \\&= e^{m_{k} - m_{k+1}} \cdot d_{k} + e^{x_{k+1} - m_{k+1}} \end{aligned} \]即: \(d_{k+1} = e^{m_{k} - m_{k+1}} \cdot d_{k} + e^{x_{k+1} - m_{k+1}}\)
- 根据定义展开 \(d_{k+1}\):
- 定义在 \(k\) 时刻(即处理完输入序列的前 \(k\) 个元素 \(x_1, ..., x_k\))的状态为 \((m_k, d_k)\):
- 并行规约计算:
- 将输入数据 \(\boldsymbol{x}\) 切分成两个不相交的子集 \(A\) 和 \(B\).
- 对任意集合 \(S\), 维护一个二元组状态 \((m_S, d_S)\):
- 局部最大值: \(m_S = max_{x \in S}x\)
- 局部指数和: \(d_S = \sum_{x \in S} e^{x - m_S}\)
- 已知 \((m_A, d_A)\) 和 \((m_B, d_B)\), 求集合 \(C = A \cup B\) 的状态 \((m_C, d_C)\).
- 全局最大值 \(m_{C}\):
- \(m_C = max(m_A, m_B)\)
- 全局指数和 \(d_{C}\):\[ \begin{aligned} d_C &= \sum_{x \in A \cup B} e^{x - m_C} \\&= \sum_{x \in A} e^{x - m_C} + \sum_{x \in B} e^{x - m_C} \\&= \sum_{x \in A} e^{x - m_A + m_A - m_C} + \sum_{x \in B} e^{x - m_B + m_B - m_C} \\&= e^{m_A - m_C} \cdot \sum_{x \in A} e^{x - m_A} + e^{m_B - m_C} \cdot \sum_{x \in B} e^{x - m_B} \\&= e^{m_A - m_C} \cdot d_A + e^{m_B - m_C} \cdot d_B \end{aligned} \]
- 即 \(d_C = e^{m_A - m_C} \cdot d_A + e^{m_B - m_C} \cdot d_B\)
- 当 \(m_A > m_B\), 此时 \(m_C = m_A\), \(d_C = d_A + d_B \cdot e^{m_B-m_A}\);
- 当 \(m_A < m_B\), 此时 \(m_C = m_B\), \(d_C = d_A \cdot e^{m_A-m_B} + d_B\);
- 当 \(m_A = m_B\), 此时 \(m_C = m_B = m_A\), \(d_C = d_A + d_B\);
- 串行递推与并行规约的关系
- 当前 \(k\) 时刻的状态 \((m_k, d_k)\) 对应集合 A 的状态 \((m_A, d_A)\)
- 将要读取的下一元素 \(x_{k+1}\), 可将其视为集合只有一个元素的集合 B, 即 \(B = \{x_{k+1}\}\); 由于集合 B 只有一个元素 \(x_{k+1}\), 其最大值仍是 \(x_{k+1}\), 则指数和 \(d_B = e^{x_{k+1}-m_B} = e^0 = 1\), 故集合 B 的状态为 \((m_B, d_B) = (x_{k+1}, 1)\).
- 将其代入并行规约的全局指数和公式, 可得 \(d_{k+1} = e^{m_{k} - m_{k+1}} \cdot d_{k} + e^{x_{k+1} - m_{k+1}}\). \((m_{k+1}, d_{k+1})\) 对应集合 C 的状态.