简化版FlashAttention

难度: 高 | 预计时间: 1-2周

🎯 项目目标

📊 内存对比

📖 核心原理

标准Attention流程: ┌─────────────────────────────────────────────────────────┐ │ 1. S = Q × K^T (N×N矩阵,O(N²)内存) │ │ 2. P = softmax(S) (N×N矩阵) │ │ 3. O = P × V (N×d矩阵,输出) │ └─────────────────────────────────────────────────────────┘ FlashAttention流程: ┌─────────────────────────────────────────────────────────┐ │ for each tile of Q, K, V: │ │ 1. 加载Q_tile, K_tile, V_tile到SRAM │ │ 2. 计算 S_tile = Q_tile × K_tile^T │ │ 3. Online Softmax更新 (维护m, l统计量) │ │ 4. 累加输出 O_tile │ │ │ │ 结果: O(N)内存,减少HBM访问 │ └─────────────────────────────────────────────────────────┘

💻 完整代码 (简化版)

/**
 * 简化版FlashAttention实现
 * 
 * 注意: 这是教学版本,用于理解核心原理
 * 生产环境请使用Dao-AILab/flash-attention
 * 
 * 编译: nvcc -o flash_attn flash_attn.cu -arch=sm_80
 */

#include <stdio.h>
#include <math.h>
#include <cuda_runtime.h>

#define BLOCK_SIZE 32  // 每个Block处理的序列长度
#define HEAD_DIM 64   // 注意力头维度

// ==================== 标准Attention ====================
// 用于对比:O(N²)内存版本
__global__ void standard_attention(
    const float* Q, const float* K, const float* V,
    float* O, int seq_len, int head_dim
) {
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    int col = blockIdx.x * blockDim.x + threadIdx.x;
    
    if (row < seq_len && col < seq_len) {
        // 计算注意力分数 S = Q × K^T
        float sum = 0.0f;
        for (int i = 0; i < head_dim; i++) {
            sum += Q[row * head_dim + i] * K[col * head_dim + i];
        }
        sum /= sqrt(head_dim);
        
        // Softmax需要先找最大值
        // 这里简化处理,实际需要所有行参与
        float score = expf(sum);
        
        // O = P × V
        float out = 0.0f;
        for (int i = 0; i < head_dim; i++) {
            out += score * V[row * head_dim + i];
        }
        O[row * head_dim + col] = out;
    }
}

// ==================== FlashAttention Kernel ====================
// O(N)内存版本
__global__ void flash_attention_kernel(
    const float* Q, const float* K, const float* V,
    float* O, int seq_len, int head_dim
) {
    // Shared Memory用于存储Q_tile, K_tile, V_tile
    extern __shared__ float smem[];
    float* q_tile = smem;                          // BLOCK_SIZE × head_dim
    float* k_tile = q_tile + BLOCK_SIZE * head_dim; // BLOCK_SIZE × head_dim
    float* v_tile = k_tile + BLOCK_SIZE * head_dim; // BLOCK_SIZE × head_dim
    
    // 当前处理的行
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    
    // Online Softmax统计量
    float m_prev = -INFINITY;  // 之前的最大值
    float l_prev = 0.0f;       // 之前的exp和
    
    // 输出累加器
    float o_acc[HEAD_DIM] = {0};
    
    // 遍历所有K,V tiles
    for (int t = 0; t < seq_len; t += BLOCK_SIZE) {
        // 1. 协作加载Q_tile, K_tile, V_tile到Shared Memory
        // ... (省略加载代码) ...
        
        // 2. 计算当前tile的注意力分数
        float s_local = 0.0f;
        for (int i = 0; i < head_dim; i++) {
            s_local += q_tile[threadIdx.x * head_dim + i] * 
                       k_tile[threadIdx.y * head_dim + i];
        }
        s_local /= sqrt(head_dim);
        
        // 3. Online Softmax更新
        float m_new = fmaxf(m_prev, s_local);
        float l_new = l_prev * expf(m_prev - m_new) + expf(s_local - m_new);
        
        // 4. 更新输出累加器
        for (int i = 0; i < head_dim; i++) {
            o_acc[i] = o_acc[i] * (l_prev * expf(m_prev - m_new) / l_new) +
                       expf(s_local - m_new) * v_tile[threadIdx.y * head_dim + i] / l_new;
        }
        
        m_prev = m_new;
        l_prev = l_new;
    }
    
    // 写入最终输出
    for (int i = 0; i < head_dim; i++) {
        O[row * head_dim + i] = o_acc[i];
    }
}

// ==================== 主函数 ====================
int main() {
    const int seq_len = 4096;
    const int head_dim = 64;
    
    printf("FlashAttention演示: seq_len=%d, head_dim=%d\n", seq_len, head_dim);
    printf("标准Attention内存: %.2f MB\n", 
           4.0 * seq_len * seq_len * head_dim / 1024 / 1024);
    printf("FlashAttention内存: %.2f MB\n",
           4.0 * seq_len * head_dim * 4 / 1024 / 1024);
    
    // ... (完整代码请参考GitHub仓库)
    
    return 0;
}

📝 代码详解

1. Online Softmax

2. Tiling策略

3. IO感知

💡 扩展任务: