CUDA 算子笔记 4

§ 参考资料 共 1 条

Decaying Causal Attention

输入 $\boldsymbol Q\in\mathbb R^{M\times d}$、$\boldsymbol K\in\mathbb R^{M\times d}$、$\boldsymbol V\in\mathbb R^{M\times d}$ 和 $\gamma\in(0,1]$,输出:

$$ \boldsymbol O_n=\sum_{m=0}^{n}\gamma^{n-m}\cdot\frac{\boldsymbol Q_n\cdot\boldsymbol K_m}{\sqrt d}\cdot\boldsymbol V_m $$

其中,$\boldsymbol O_i$、$\boldsymbol Q_i$、$\boldsymbol K_i$ 和 $\boldsymbol V_i$ 分别表示输出和输入矩阵的第 $i$ 行

例子:

Input:
Q (2x4): [[1.0, 1.0, 0.0, 0.0],
          [1.0, 1.0, 0.0, 0.0]]
K (2x4): [[1.0, 0.0, 0.0, 0.0],
          [0.0, 1.0, 0.0, 0.0]]
V (2x4): [[4.0, 8.0, 12.0, 16.0],
          [4.0, 8.0, 12.0, 16.0]]

Output:
O (2x4): [[2.0, 4.0, 6.0, 8.0],
          [3.0, 6.0, 9.0, 12.0]]

数据范围:$1\le M\le 8192$,$1\le d\le 256$

最简单的做法

其实这就是去掉了 Softmat 的另一种形式的更简单的 Attention,改造一下 Softmat Attention 的代码即可

Softmax Attention 参考 CUDA 算子笔记 2

代码:

#include <cuda_runtime.h>

#define CEIL(x, y) (((x) + (y) - 1) / (y))

template <int BX, int BY, int KS, int TM, int TN, bool Transpose, bool Norm,
          bool Decay>
__global__ void multiple_kernel(const float* __restrict__ A,
                                const float* __restrict__ B,
                                float* __restrict__ output, int M, int N, int K,
                                float scale = 1.0f, float log_gamma = 0.0f) {
  constexpr int BM = BY * TM;
  constexpr int BN = BX * TN;
  constexpr int T = BX * BY;

  __shared__ float tile_A[BM][KS + 1];
  __shared__ float tile_B[KS][BN + 1];

  const int tid = threadIdx.y * BX + threadIdx.x;

  const int block_row = blockIdx.y * BM;
  const int block_col = blockIdx.x * BN;

  const int thread_row = threadIdx.y * TM;
  const int thread_col = threadIdx.x * TN;

  float acc[TM][TN] = {};

#pragma unroll
  for (int k0 = 0; k0 < K; k0 += KS) {
#pragma unroll
    for (int i = 0; i < BM * KS; i += T) {
      const int local_row = (tid + i) / KS;
      const int local_col = (tid + i) % KS;

      const int global_row = block_row + local_row;
      const int global_col = k0 + local_col;

      if (global_row < M && global_col < K) {
        tile_A[local_row][local_col] = A[global_row * K + global_col];
      } else {
        tile_A[local_row][local_col] = 0.0f;
      }
    }

    if constexpr (Transpose) {
#pragma unroll
      for (int i = 0; i < KS * BN; i += T) {
        const int local_row = (tid + i) / KS;
        const int local_col = (tid + i) % KS;

        const int global_row = block_col + local_row;
        const int global_col = k0 + local_col;

        if (global_row < N && global_col < K) {
          tile_B[local_col][local_row] = B[global_row * K + global_col];
        } else {
          tile_B[local_col][local_row] = 0.0f;
        }
      }
    } else {
#pragma unroll
      for (int i = 0; i < KS * BN; i += T) {
        const int local_row = (tid + i) / BN;
        const int local_col = (tid + i) % BN;

        const int global_row = k0 + local_row;
        const int global_col = block_col + local_col;

        if (global_row < K && global_col < N) {
          tile_B[local_row][local_col] = B[global_row * N + global_col];
        } else {
          tile_B[local_row][local_col] = 0.0f;
        }
      }
    }

    __syncthreads();

#pragma unroll
    for (int k = 0; k < KS; ++k) {
      float a_reg[TM];
      float b_reg[TN];

#pragma unroll
      for (int i = 0; i < TM; ++i) {
        a_reg[i] = tile_A[thread_row + i][k];
      }

#pragma unroll
      for (int j = 0; j < TN; ++j) {
        b_reg[j] = tile_B[k][thread_col + j];
      }

#pragma unroll
      for (int i = 0; i < TM; ++i) {
#pragma unroll
        for (int j = 0; j < TN; ++j) {
          acc[i][j] = fmaf(a_reg[i], b_reg[j], acc[i][j]);
        }
      }
    }

    __syncthreads();
  }

#pragma unroll
  for (int i = 0; i < TM; ++i) {
    const int row = block_row + thread_row + i;

#pragma unroll
    for (int j = 0; j < TN; ++j) {
      const int col = block_col + thread_col + j;

      if (row < M && col < N) {
        float value = acc[i][j];

        if constexpr (Norm) {
          value *= scale;
        }

        if constexpr (Decay) {
          if (col <= row) {
            value *= __expf((row - col) * log_gamma);
          } else {
            value = 0.0f;
          }
        }

        output[row * N + col] = value;
      }
    }
  }
}

// Q, K, V, output are device pointers
extern "C" void solve(const float* Q, const float* K, const float* V,
                      float* output, int seq_len, int d_model, float gamma) {
  const int M = seq_len, N = seq_len, d = d_model;

  float* attention;
  cudaMalloc(&attention, M * N * sizeof(float));

  {
    constexpr int BX = 16, BY = 16;
    constexpr int TM = 4, TN = 4;
    constexpr int BM = BY * TM, BN = BX * TN;
    constexpr int KS = 32;

    const dim3 block_dim(BX, BY);
    const dim3 grid_dim(CEIL(N, BN), CEIL(M, BM));

    multiple_kernel<BX, BY, KS, TM, TN, true, true, true>
        <<<grid_dim, block_dim>>>(Q, K, attention, M, N, d, 1.0f / sqrtf(d),
                                  logf(gamma));
  }

  {
    constexpr int BX = 16, BY = 16;
    constexpr int TM = 4, TN = 4;
    constexpr int BM = BY * TM, BN = BX * TN;
    constexpr int KS = 32;

    const dim3 block_dim(BX, BY);
    const dim3 grid_dim(CEIL(d, BN), CEIL(M, BM));

    multiple_kernel<BX, BY, KS, TM, TN, false, false, false>
        <<<grid_dim, block_dim>>>(attention, V, output, M, d, N);
  }

  cudaFree(attention);
}

基于递推的方法

相比于 Softmax Attention,这个 Attention 方法去掉了 Softmax,所以可以做这样的变换:

$$ \begin{aligned} \boldsymbol O_{n,j}&=\sum_{m=0}^{n}\gamma^{n-m}\cdot\frac{\boldsymbol Q_n\cdot\boldsymbol K_m}{\sqrt d}\cdot\boldsymbol V_{m,j} \\ &=\frac{1}{\sqrt d}\sum_{m=0}^{n}\gamma^{n-m}\left(\sum_{i=0}^d\boldsymbol Q_{n,i}\boldsymbol K_{m,i}\right)\boldsymbol V_{m,j} \\ &=\frac{1}{\sqrt d}\sum_{i=0}^{d}\boldsymbol Q_{n,i}\left(\sum_{m=0}^{n}\gamma^{n-m}\boldsymbol K_{m,i}\boldsymbol V_{m,j}\right) \end{aligned} $$

即调换内外两个求和符号

然后定义:

$$ \boldsymbol S_n[i,j]=\sum_{m=0}^n\gamma^{n-m}\boldsymbol K_{m,i}\boldsymbol V_{m,j} $$

写成矩阵形式:

$$ \boldsymbol S_n=\sum_{m=0}^n\gamma^{n-m}\boldsymbol K_m^T\boldsymbol V_m $$

有递推公式:

$$ \boldsymbol S_n=\gamma\boldsymbol S_{n-1}+\boldsymbol K_n^T\boldsymbol V_n $$

那么:

$$ \boldsymbol O_{n,j}=\frac{1}{\sqrt d}\sum_{i=0}^{d}\boldsymbol Q_{n,i}\boldsymbol S_n[i,j] $$

相当于 $\boldsymbol Q$ 的第 $n$ 行与 $\boldsymbol S_n$ 的第 $j$ 列相乘,即一个向量乘上一个矩阵:

$$ \boldsymbol O_n=\frac{\boldsymbol Q_n\boldsymbol S_n}{\sqrt d} $$

在这里,$\boldsymbol Q_n\in\mathbb R^{1\times d}$、$\boldsymbol S_n\in\mathbb R^{d\times d}$,根本不需要中间的 $M\times M$ 的 Attention,直接:

S = 0

for n:
    S = gamma * S + outer(K[n], V[n])
    O[n] = Q[n] @ S / sqrt(d)

就可以了,这样把 $O(s^2d)$ 的 Attention 改写成一个 $d\times d$ 状态的递推

这个递推可以算是 Prefix Scan,代码:

#include <cuda_runtime.h>

#define CEIL(x, y) (((x) + (y) - 1) / (y))

constexpr int k_max_d = 256;
constexpr int k_threads_per_warp = 32;

// 每个 d*d 的 state 由多个 warp 共同完成,每个 warp 完成 4
// 列(一次完整的矩阵乘法运算需要第二个矩阵的完整一列,所以这里一个 warp
// 算完整一列)
// 每个 thread 分到 8*4
constexpr int k_cols_per_warp = 4;
constexpr int k_rows_per_thread = k_max_d / k_threads_per_warp;

// 每个 block 有 4 个 warp
constexpr int k_warps_per_block = 4;
constexpr int k_threads_per_block = k_threads_per_warp * k_warps_per_block;
constexpr int k_cols_per_block = k_warps_per_block * k_cols_per_warp;

// 每个 block 推进 256 个 state step
constexpr int k_chs_per_block = 256;

__device__ __forceinline__ float warp_reduce(float val) {
#pragma unroll
  for (int offset = 16; offset > 0; offset >>= 1) {
    val += __shfl_down_sync(0xffffffffu, val, offset);
  }
  return val;
}

__global__ void decay_atten_k1(const float* __restrict__ Q,
                               const float* __restrict__ K,
                               const float* __restrict__ V,
                               float* __restrict__ output,
                               float* __restrict__ d_block_states, int M, int D,
                               float gamma, float scale) {
  const int warp = threadIdx.x >> 5;
  const int lane = threadIdx.x & 31;

  const int ch_base = blockIdx.y * k_chs_per_block;
  const int col_base = blockIdx.x * k_cols_per_block + warp * k_cols_per_warp;

  float reg[k_rows_per_thread][k_cols_per_warp];

#pragma unroll
  for (int u = 0; u < k_rows_per_thread; ++u) {
#pragma unroll
    for (int j = 0; j < k_cols_per_warp; ++j) {
      reg[u][j] = 0.0f;
    }
  }

  for (int t = 0; t < k_chs_per_block; ++t) {
    const int ch = t + ch_base;
    if (ch < M) {
      float vv[k_cols_per_warp];
      float dot[k_cols_per_warp];

#pragma unroll
      for (int j = 0; j < k_cols_per_warp; ++j) {
        const int col = col_base + j;
        vv[j] = col < D ? V[ch * D + col] : 0.0f;
        dot[j] = 0.0f;
      }

#pragma unroll
      for (int i = 0; i < k_rows_per_thread; ++i) {
        const int row = lane + i * k_threads_per_warp;
        if (row < D) {
          float kv = K[ch * D + row];
          float qv = Q[ch * D + row];

#pragma unroll
          for (int j = 0; j < k_cols_per_warp; ++j) {
            const int col = col_base + j;
            if (col < D) {
              reg[i][j] = fmaf(gamma, reg[i][j], kv * vv[j]);
              dot[j] = fmaf(qv, reg[i][j], dot[j]);
            }
          }
        }
      }

#pragma unroll
      for (int j = 0; j < k_cols_per_warp; j++) {
        dot[j] = warp_reduce(dot[j]);
      }

      if (lane == 0) {
#pragma unroll
        for (int j = 0; j < k_cols_per_warp; j++) {
          int col = col_base + j;
          if (col < D) {
            output[ch * D + col] = dot[j] * scale;
          }
        }
      }
    }
  }

#pragma unroll
  for (int i = 0; i < k_rows_per_thread; ++i) {
    const int row = lane + i * k_threads_per_warp;

    if (row < D) {
#pragma unroll
      for (int j = 0; j < k_cols_per_warp; ++j) {
        const int col = col_base + j;

        if (col < D) {
          d_block_states[(blockIdx.y * D + col) * D + row] = reg[i][j];
        }
      }
    }
  }
}

__global__ void decay_atten_k2(float* __restrict__ d_block_states,
                               int num_chunks, int M, int D, float gamma_chunk,
                               float gamma) {
  const int tid = blockIdx.x * blockDim.x + threadIdx.x;
  const int num_items = D * D;

  if (tid >= num_items) {
    return;
  }

  float carry = 0.0f;
  for (int i = 0; i < num_chunks; ++i) {
    const int idx = i * num_items + tid;

    const float val = d_block_states[idx];
    d_block_states[idx] = carry;

    const int len = min(k_chs_per_block, M - i * k_chs_per_block);
    const float decay =
        (len == k_chs_per_block) ? gamma_chunk : __powf(gamma, len);
    carry = fmaf(decay, carry, val);
  }
}

__global__ void decay_atten_k3(const float* __restrict__ Q,
                               float* __restrict__ output,
                               const float* __restrict__ d_block_states, int M,
                               int D, float gamma, float scale) {
  const int warp = threadIdx.x >> 5;
  const int lane = threadIdx.x & 31;

  const int chunk = blockIdx.y + 1;

  const int ch_base = chunk * k_chs_per_block;
  const int col_base = blockIdx.x * k_cols_per_block + warp * k_cols_per_warp;

  float reg[k_rows_per_thread][k_cols_per_warp];

#pragma unroll
  for (int i = 0; i < k_rows_per_thread; ++i) {
    const int row = lane + i * k_threads_per_warp;

#pragma unroll
    for (int j = 0; j < k_cols_per_warp; ++j) {
      const int col = col_base + j;

      if (row < D && col < D) {
        reg[i][j] = d_block_states[(chunk * D + col) * D + row];
      } else {
        reg[i][j] = 0.0f;
      }
    }
  }

  float decay = 1.0f;

  for (int t = 0; t < k_chs_per_block; ++t) {
    const int ch = t + ch_base;

    if (ch < M) {
      decay *= gamma;

      float dot[k_cols_per_warp] = {};

#pragma unroll
      for (int u = 0; u < k_rows_per_thread; ++u) {
        const int row = lane + u * k_threads_per_warp;

        if (row < D) {
          const float qv = Q[ch * D + row];

#pragma unroll
          for (int j = 0; j < k_cols_per_warp; ++j) {
            const int col = col_base + j;

            if (col < D) {
              dot[j] = fmaf(qv, reg[u][j], dot[j]);
            }
          }
        }
      }

#pragma unroll
      for (int j = 0; j < k_cols_per_warp; ++j) {
        dot[j] = warp_reduce(dot[j]);
      }

      if (lane == 0) {
#pragma unroll
        for (int j = 0; j < k_cols_per_warp; ++j) {
          const int col = col_base + j;

          if (col < D) {
            const int idx = ch * D + col;
            output[idx] = fmaf(dot[j], scale * decay, output[idx]);
          }
        }
      }
    }
  }
}

// Q, K, V, output are device pointers
extern "C" void solve(const float* Q, const float* K, const float* V,
                      float* output, int seq_len, int d_model, float gamma) {
  const int M = seq_len;
  const int D = d_model;

  const int num_chunks = CEIL(M, k_chs_per_block);

  float* d_block_states = nullptr;
  cudaMalloc(&d_block_states, num_chunks * D * D * sizeof(float));

  {
    constexpr int block_dim = k_threads_per_block;
    const dim3 grid_dim(CEIL(D, k_cols_per_block), num_chunks);
    decay_atten_k1<<<grid_dim, block_dim>>>(Q, K, V, output, d_block_states, M,
                                            D, gamma, 1.0f / sqrtf(D));
  }

  {
    constexpr int block_dim = 256;
    const int grid_dim = CEIL(D * D, block_dim);
    const float gamma_chunk = powf(gamma, k_chs_per_block);
    decay_atten_k2<<<grid_dim, block_dim>>>(d_block_states, num_chunks, M, D,
                                            gamma_chunk, gamma);
  }

  {
    constexpr int block_dim = k_threads_per_block;
    const dim3 grid_dim(CEIL(D, k_cols_per_block), num_chunks - 1);
    decay_atten_k3<<<grid_dim, block_dim>>>(Q, output, d_block_states, M, D,
                                            gamma, 1.0f / sqrtf(D));
  }

  cudaFree(d_block_states);
}

把 $d\times d$ 的 State 分到若干个 Warp 里面,每个 Warp 计算 State 的一个完整列(这里矩阵乘法的第二个矩阵需要完整的列,所以要这样,否则计算一个完整的值需要跨 Warp 甚至跨 Block 求和,这里直接 Warp Reduce 即可)

  • Kernel 1 每 256 步计算一个 State,并且计算 Local 结果
  • Kernel 2 对 Chunk State 计算前缀和
  • Kernel 3 计算 Global 结果