从零实现 vLLM(1.3):如何加速 Attention 计算

1. 从零实现 vLLM(1.3):如何加速 Attention 计算

vLLM 是目前最主流的 LLM 推理框架,它主要解决了 LLM 推理时的内存瓶颈问题。
本文主要介绍 vLLM 如何加速 Attention 计算。

1.1. Attention 计算过程

1.1.1. Attention 的基本原理

Attention 机制是 Transformer 架构的核心,它允许模型在处理序列时动态关注不同位置的信息
基本的 Attention 计算过程如下:

  1. 计算查询(Query)、键(Key)和值(Value):通过线性变换从输入中生成 Q、K、V 矩阵
  2. 计算注意力分数:通过 Q 和 K 的点积计算注意力分数
  3. 应用 Softmax:对注意力分数进行归一化
  4. 加权求和:使用归一化后的注意力分数对 V 进行加权求和

1.1.2. 标准实现的问题

标准的 Attention 实现存在以下问题:

  • 内存占用大:需要存储完整的注意力矩阵,空间复杂度为 O(N²)
  • 计算效率低:需要多次访问 GPU 全局内存
  • 显存碎片化:不同长度的序列导致显存利用率低

1.2. FlashAttention 原理

1.2.1. 核心思想

FlashAttention 的核心思想是避免构建完整的注意力矩阵,通过分块计算和在线更新策略,显著减少内存访问次数。

1.2.2. I/O 感知计算

FlashAttention 通过以下方式实现 I/O 感知计算:

  1. 分块处理:将 Q、K、V 矩阵分成小块,在 SRAM 中完成计算
  2. 在线更新:使用递推公式在线更新中间结果
  3. 减少全局内存访问:最大程度减少对慢速 GPU 全局内存的访问

1.2.3. 数值稳定性

FlashAttention 使用在线 Softmax 算法保证数值稳定性,避免计算过程中的数值溢出问题。

1.3. FlashAttention 实现细节

1.3.1. 安全 Softmax

标准 Softmax 计算存在数值溢出风险,特别是当输入值较大时。
安全 Softmax 通过以下方式解决:

  1. 减去最大值:在计算指数前减去输入的最大值
  2. 数值稳定:确保计算过程不会产生数值溢出
def safe_softmax(x, dim=-1):
    # 减去最大值防止数值溢出
    x_max = torch.max(x, dim=dim, keepdim=True)[0]
    x_shifted = x - x_max
    return torch.softmax(x_shifted, dim=dim)

1.3.2. 在线 Softmax

在线 Softmax 是 FlashAttention 的关键技术,它允许在一次遍历中计算 Softmax,而不需要构建完整的注意力矩阵。

1.3.2.1. 传统 Softmax 的问题

传统 Softmax 即便采用安全写法,也需要对输入进行多次遍历:

  1. 第一次遍历:找到最大值
  2. 第二次遍历:计算指数和(分母)
  3. 第三次遍历:输出归一化结果

1.3.2.2. 在线 Softmax 的解决方案

在线 Softmax 通过递推关系在一次遍历中完成计算:

m_i = max(m_{i-1}, x_i)
d_i = d_{i-1} * exp(m_{i-1} - m_i) + exp(x_i - m_i)

其中:

  • m_i:截至当前的最大值
  • d_i:截至当前的代理分母

1.3.3. FlashAttention 的分块计算

FlashAttention 将 Q、K、V 矩阵分成小块,在 SRAM 中完成计算:

  1. 加载块数据:将 Q、K、V 的块加载到 SRAM
  2. 局部计算:在 SRAM 中完成局部注意力计算
  3. 更新全局状态:使用递推公式更新全局状态
  4. 处理下一块:继续处理下一块数据

1.3.4. FlashAttention 的最终实现

FlashAttention 的最终实现结合了在线 Softmax分块计算

def flash_attention(q, k, v, block_size, softmax_scale=1.0):
    # q: (1, d),k/v: (N, d)
    # 初始化全局状态(必须是张量,才能参与后续的逐块更新)
    m = torch.full((1,), -float('inf'))   # 运行最大值
    d = torch.zeros(1)                    # 代理分母
    o = torch.zeros(1, q.size(-1))        # 代理输出

    # 分块处理 K 和 V
    for i in range(0, k.size(0), block_size):
        k_block = k[i:i+block_size]
        v_block = v[i:i+block_size]

        # 计算局部注意力分数
        scores = torch.matmul(q, k_block.transpose(-2, -1)) * softmax_scale

        # 在线更新最大值,旧状态按 exp(m - new_m) 缩放
        block_max = scores.max(dim=-1)[0]
        new_m = torch.maximum(m, block_max)
        correction = torch.exp(m - new_m)

        # 更新代理分母和输出
        p = torch.exp(scores - new_m)
        d = d * correction + p.sum(dim=-1)
        o = o * correction.unsqueeze(-1) + torch.matmul(p, v_block)

        m = new_m

    # 用真实分母归一化,得到最终输出
    return o / d.unsqueeze(-1)

1.4. PagedAttention

1.4.1. PagedAttention 的原理

PagedAttention 是 vLLM 的另一项关键技术,它通过分页管理解决显存碎片化问题:

  1. 块管理:将 KV Cache 分成固定大小的块
  2. 虚拟内存:使用块表管理非连续的显存
  3. 按需分配:根据需要分配和释放块

1.4.2. 块表结构

PagedAttention 使用块表结构管理非连续的显存:

块表 = [
    [block_id_1, block_id_2, ...],  # 序列 1 的块
    [block_id_3, block_id_4, ...],  # 序列 2 的块
    ...
]

1.4.3. 前缀缓存

PagedAttention 支持前缀缓存,共享相同前缀的序列可以复用 KV Cache:

  1. 前缀识别:识别序列的共同前缀
  2. 缓存共享:多个序列共享同一前缀的 KV Cache
  3. 内存节省:显著减少显存使用

1.5. vLLM 中的 Attention 实现

1.5.1. Prefill 阶段

Prefill 阶段处理输入序列,生成 KV Cache:

def prefill_attention(q, k, v, block_tables):
    # 存储 KV Cache
    store_kvcache(k, v, k_cache, v_cache, slot_mapping)

    # 计算 Attention
    o = flash_attn_varlen_func(
        q, k, v,
        max_seqlen_q=max_seqlen_q,
        cu_seqlens_q=cu_seqlens_q,
        max_seqlen_k=max_seqlen_k,
        cu_seqlens_k=cu_seqlens_k,
        softmax_scale=scale,
        causal=True,
        block_table=block_tables,
    )
    return o

1.5.2. Decode 阶段

Decode 阶段逐个生成 token,使用缓存的 KV Cache:

def decode_attention(q, k_cache, v_cache, block_tables):
    # 添加序列长度维度
    q = q.unsqueeze(1)

    # 计算 Attention
    o = flash_attn_with_kvcache(
        q, k_cache, v_cache,
        cache_seqlens=context_lens,
        block_table=block_tables,
        softmax_scale=scale,
        causal=True,
    )
    return o

1.5.3. 块表管理

vLLM 使用块表管理非连续的显存:

class BlockTable:
    def __init__(self, block_size, num_blocks):
        self.block_size = block_size
        self.num_blocks = num_blocks
        self.free_blocks = list(range(num_blocks))
        self.allocated_blocks = {}

    def allocate(self, seq_id, num_blocks_needed):
        # 分配块
        blocks = []
        for _ in range(num_blocks_needed):
            if not self.free_blocks:
                raise RuntimeError("No free blocks available")
            block_id = self.free_blocks.pop()
            blocks.append(block_id)

        self.allocated_blocks[seq_id] = blocks
        return blocks

    def free(self, seq_id):
        # 释放块
        if seq_id in self.allocated_blocks:
            blocks = self.allocated_blocks.pop(seq_id)
            self.free_blocks.extend(blocks)

1.6. 性能优化

1.6.1. 内核融合

vLLM 使用内核融合技术减少内存访问:

  1. 多步融合:将多个计算步骤融合为一个 CUDA 内核
  2. 减少中间结果:避免存储中间结果到全局内存
  3. 提高计算密度:增加每个内存访问的计算量

1.6.2. 内存池管理

vLLM 使用内存池管理显存:

  1. 预分配:预先分配大块显存
  2. 块管理:将显存分成固定大小的块
  3. 快速分配:避免频繁的显存分配和释放

1.6.3. 批处理优化

vLLM 通过批处理提高吞吐量:

  1. 动态批处理:将不同长度的序列组合成批次
  2. 连续批处理:新序列可以立即加入批次
  3. 前缀缓存:共享相同前缀的序列可以复用计算

1.7. 总结

vLLM 通过以下技术显著提高了 LLM 推理性能:

  1. FlashAttention:通过 I/O 感知计算和在线 Softmax 减少内存访问
  2. PagedAttention:通过分页管理解决显存碎片化问题
  3. 前缀缓存:共享相同前缀的序列可以复用 KV Cache
  4. 内核融合:减少内存访问次数,提高计算效率
  5. 动态批处理:提高 GPU 利用率,增加吞吐量

这些技术的结合使得 vLLM 成为了目前最主流的 LLM 推理框架,显著提高了 LLM 推理的性能和效率。