§ 参考资料 共 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 结果