Online Softmax

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}}\)
  • 并行规约计算:
    • 将输入数据 \(\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 的状态.