CUDA算子优化(3):Softmax

Softmax是Transformer中Attention计算的核心组件,本文从朴素实现暴露的数值问题出发,逐步引入Safe Softmax、Online Softmax等手段,一步一步优化Softmax Kernel的性能。

Softmax的定义

Softmax函数也称为归一化指数函数,它能将一个含任意实数的N维向量z压缩到另一个N维实向量$\sigma(z)$中,使得每个元素的范围都在(0, 1)之间,并且所有元素的和为1。该函数定义按下面式子给出:

$$\large \text{Softmax}(x_i) = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}}$$

其中,$x_i$为第i个节点输出的值,N为输出节点的个数。

Softmax的公式涉及到指数运算$e^{x_{i}}$,存在数值下溢问题,当$x_i$很小时,分母趋近于0,导致NAN。那么怎么解决呢?

解决方案:减最大值技巧

数学上可以证明,对输入向量的每个元素减去同一个常数c,不改变Softmax的输出:

$$\large \text{Softmax}(x_i) = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}} = \frac{e^{x_i - c}}{\sum_{j=1}^{N} e^{x_j - c}}$$

当c=max(x)时,即c为$x_i$的最大值,所有指数的参数都<=0,结果会落在(0, 1]之间,这就是Safe Softmax。

Softmax的cpu实现

Safe Softmax的实现可以分为三个阶段:

  1. 求向量的最大值c
  2. 求和$\sum_{j=1}^{N} e^{x_j - c}$
  3. 归一化$\frac{e^{x_i - c}}{\sum_{j=1}^{N} e^{x_j - c}}$

具体代码实现如下:(B为Batch size,T为序列长度,C为通道数)

void softmax_cpu(float *h_out, float *h_data, int B, int T, int C) {
for (int i = 0; i < B * T; ++i) {
float *inp = h_data + i * C;
float *out = h_out + i * C;

float maxval = -INFINITY;
for (int j = 0; j < C; ++j) {
if (inp[j] > maxval) {
maxval = inp[j];
}
}

float sum = 0.0f;
for (int j = 0; j < C; ++j) {
sum += expf(inp[j] - maxval);
}
for (int j = 0; j < C; ++j) {
out[j] = expf(inp[j] - maxval) / sum;
}
}
}

版本V0一个线程处理一行

假设B表示Batch size,T表示序列长度,C表示通道数,我们要处理的维度是在C这个维度,如果是一个线程处理一行,那么我们总共需要B*T个线程,也就是说block_size为BLOCK_SIZE时,grid_size为(B * T + BLOCK_SIZE - 1) / BLOCK_SIZE。

__global__ void softmax_gpu(float *d_out, float *d_data, int N, int C) {
int tid = threadIdx.x + blockDim.x * blockIdx.x;
int row = tid;
if (row >= N) {
return;
}

float *inp = d_data + row * C;
float *out = d_out + row * C;

float maxval = -INFINITY;
for (int i = 0; i < C; ++i) {
if (inp[i] > maxval) {
maxval = inp[i];
}
}

float sum = 0.0f;
for (int i = 0; i < C; ++i) {
sum += expf(inp[i] - maxval);
}
for (int i = 0; i < C; ++i) {
out[i] = expf(inp[i] - maxval) / sum;
}
}

这种方法有一个问题在于,总共要遍历三遍向量,有没有办法只遍历两遍向量呢?即在求和的同时更新最大值。对于Memory-bound的操作来说,减少一段扫描,相当于提升性能33%。这就引出了下面的方法Online Softmax。

版本V1 Online Softmax

假设我们已经处理了前j个数据,得到了

  • 当前最大值$m_j$(就是前j个数里最大的那个)

  • 当前分母$d_j = \sum_{i=1}^{j} e^{x_i - m_j}$

现在,第j+1个数据$x_{j+1}$来了,我们需要更新$m_j$和$d_j$

步骤1:更新最大值

新来的$x_{j+1}$可能比之前的最大值$m_j$大,也可能小。所以,新的最大值为:

$$m_{j+1} = max(m_j, x_{j+1})$$

这很好理解,新的最大值要么是老的最大值,要么是新来的值,挑大的那个就行。

步骤2:更新分母

我们需要重新计算新的分母$d_{j+1}$,它应该是:

$$d_{j+1} = \sum_{i = 1}^{j+1}{e^{x_i - m_{j+1}}} = \sum_{i = 1}^{j}{(e^{x_i - m_{j}}*e^{m_j - m_{j+1}})} + e^{x_{j+1}-m_{j+1}} = d_j * e^{m_j - m_{j+1}} + e^{x_{j+1}-m_{j+1}}$$

代码如下:

__global__ void online_softmax(float *d_out, float *d_data, int N, int C) {
int tid = threadIdx.x + blockDim.x * blockIdx.x;
int row = tid;
if (row >= N) {
return;
}
float *inp = d_data + row * C;
float *out = d_out + row * C;

float sum = 0.0f;
float maxval = -INFINITY;
for (int i = 0; i < C; ++i) {
float val = inp[i];
if (inp[i] > maxval) {
sum *= expf(maxval - val);
maxval = val;
}
sum += expf(val - maxval);
}

float r = 1 / sum;
for (int i = 0; i < C; ++i) {
out[i] = expf(inp[i] - maxval) * r;
}
}

版本V2 一个Warp处理一行

#define BLOCK_SIZE 32

__global__ void softmax_gpu(float *d_out, float *d_data, int N, int C) {
int tid = threadIdx.x + blockDim.x * blockIdx.x;
int row_id = tid / blockDim.x;
int lane_id = tid % blockDim.x;
if (row_id >= N) {
return;
}

float *inp = d_data + row_id * C;
float *out = d_out + row_id * C;

float maxval = -INFINITY;
float sum = 0.0f;
for (int i = lane_id; i < C; i += BLOCK_SIZE) {
float val = inp[i];
if (val > maxval) {
sum *= expf(maxval - val);
maxval = val;
}
sum += expf(val - maxval);
}

unsigned mask = 0xFFFFFFFF;
for (int offset = BLOCK_SIZE / 2; offset > 0; offset /= 2) {
float sum_other = __shfl_down_sync(mask, sum, offset);
float maxval_other = __shfl_down_sync(mask, maxval, offset);
float maxval_new = fmaxf(maxval, maxval_other);
sum = sum * expf(maxval - maxval_new) + sum_other * expf(maxval_other - maxval_new);
maxval = maxval_new;
}
sum = __shfl_sync(mask, sum, 0);
maxval = __shfl_sync(mask, maxval, 0);

float r = 1 / sum;
for (int i = lane_id; i < C; i+=BLOCK_SIZE) {
out[i] = expf(inp[i] - maxval) * r;
}
}