从零实现 vLLM(1.3):如何加速 Attention 计算
1. 从零实现 vLLM(1.3):如何加速 Attention 计算
vLLM 是目前最主流的 LLM 推理框架,它主要解决了 LLM 推理时的内存瓶颈问题。
本文主要介绍 vLLM 如何加速 Attention 计算。
1.1. Attention 计算过程
1.1.1. Attention 的基本原理
Attention 机制是 Transformer 架构的核心,它允许模型在处理序列时动态关注不同位置的信息。
基本的 Attention 计算过程如下:
- 计算查询(Query)、键(Key)和值(Value):通过线性变换从输入中生成 Q、K、V 矩阵
- 计算注意力分数:通过 Q 和 K 的点积计算注意力分数
- 应用 Softmax:对注意力分数进行归一化
- 加权求和:使用归一化后的注意力分数对 V 进行加权求和
1.1.2. 标准实现的问题
标准的 Attention 实现存在以下问题:
- 内存占用大:需要存储完整的注意力矩阵,空间复杂度为 O(N²)
- 计算效率低:需要多次访问 GPU 全局内存
- 显存碎片化:不同长度的序列导致显存利用率低
1.2. FlashAttention 原理
1.2.1. 核心思想
FlashAttention 的核心思想是避免构建完整的注意力矩阵,通过分块计算和在线更新策略,显著减少内存访问次数。
1.2.2. I/O 感知计算
FlashAttention 通过以下方式实现 I/O 感知计算:
- 分块处理:将 Q、K、V 矩阵分成小块,在 SRAM 中完成计算
- 在线更新:使用递推公式在线更新中间结果
- 减少全局内存访问:最大程度减少对慢速 GPU 全局内存的访问
1.2.3. 数值稳定性
FlashAttention 使用在线 Softmax 算法保证数值稳定性,避免计算过程中的数值溢出问题。
1.3. FlashAttention 实现细节
1.3.1. 安全 Softmax
标准 Softmax 计算存在数值溢出风险,特别是当输入值较大时。
安全 Softmax 通过以下方式解决:
- 减去最大值:在计算指数前减去输入的最大值
- 数值稳定:确保计算过程不会产生数值溢出
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.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 中完成计算:
- 加载块数据:将 Q、K、V 的块加载到 SRAM
- 局部计算:在 SRAM 中完成局部注意力计算
- 更新全局状态:使用递推公式更新全局状态
- 处理下一块:继续处理下一块数据
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 的另一项关键技术,它通过分页管理解决显存碎片化问题:
- 块管理:将 KV Cache 分成固定大小的块
- 虚拟内存:使用块表管理非连续的显存
- 按需分配:根据需要分配和释放块
1.4.2. 块表结构
PagedAttention 使用块表结构管理非连续的显存:
块表 = [
[block_id_1, block_id_2, ...], # 序列 1 的块
[block_id_3, block_id_4, ...], # 序列 2 的块
...
]
1.4.3. 前缀缓存
PagedAttention 支持前缀缓存,共享相同前缀的序列可以复用 KV Cache:
- 前缀识别:识别序列的共同前缀
- 缓存共享:多个序列共享同一前缀的 KV Cache
- 内存节省:显著减少显存使用
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 使用内核融合技术减少内存访问:
- 多步融合:将多个计算步骤融合为一个 CUDA 内核
- 减少中间结果:避免存储中间结果到全局内存
- 提高计算密度:增加每个内存访问的计算量
1.6.2. 内存池管理
vLLM 使用内存池管理显存:
- 预分配:预先分配大块显存
- 块管理:将显存分成固定大小的块
- 快速分配:避免频繁的显存分配和释放
1.6.3. 批处理优化
vLLM 通过批处理提高吞吐量:
- 动态批处理:将不同长度的序列组合成批次
- 连续批处理:新序列可以立即加入批次
- 前缀缓存:共享相同前缀的序列可以复用计算
1.7. 总结
vLLM 通过以下技术显著提高了 LLM 推理性能:
- FlashAttention:通过 I/O 感知计算和在线 Softmax 减少内存访问
- PagedAttention:通过分页管理解决显存碎片化问题
- 前缀缓存:共享相同前缀的序列可以复用 KV Cache
- 内核融合:减少内存访问次数,提高计算效率
- 动态批处理:提高 GPU 利用率,增加吞吐量
这些技术的结合使得 vLLM 成为了目前最主流的 LLM 推理框架,显著提高了 LLM 推理的性能和效率。