导读:本期聚焦于弦宿​创作的《Flash Attention与Paged Attention如何共同优化大模型推理?》,敬请观看详情。注意力矩阵的显存占用会随序列长度平方增长,这在长文本推理和批量服务中很快就会触碰显存上限。标准注意力计算会把完整的 N×N 分数矩阵写入高带宽显存,N 一旦到达数万,仅中间结果就可能超过数十GB。Flash Attention 改变计算顺序,将查询、键、值分块加载到片上 SRAM,通过在线 Softmax 维护运行统计量,避免物化完整注意力矩阵。Paged Attention 则针对解码阶段的 KV 缓存管理,借鉴操作系统的分页机制,把每层 KV 缓存切分成固定大小的物理块,通过页表映射允许非连续存储,从而减少碎片并提升并发吞吐。两者并不冲突:Flash Attention 降低单次自注意力的显存峰值,Paged Attention 改善多请求推理时的显存分配效率。理解这两个优化点的适用场景和实现思路,能帮助工程师在长上下文、高并发部署中更合理地选择推理框架和模型参数。

Transformer 模型在自然语言处理和多模态任务中的核心是自注意力机制。给定查询 Q、键 K 和值 V,标准计算会先得到注意力分数矩阵 S = QK^T,再经过 Softmax 后与 V 相乘。这个过程在序列长度为 N 时需要维护 N×N 的中间矩阵,其显存和访存开销成为长序列推理的瓶颈。以 N=4096、FP16 精度为例,单个注意力分数矩阵约占 32MB,多层堆叠后仅中间结果就可能占用数 GB。为了突破这一限制,业界出现了从不同层面切入的优化方案,其中 Flash Attention 聚焦计算过程中的 IO 效率,Paged Attention 聚焦推理服务中的 KV 缓存内存管理。下面从这两类机制展开分析。

Flash Attention与Paged Attention如何共同优化大模型推理?

一、标准注意力的显存与访存问题

标准自注意力的前向过程可以拆成三步:先计算 Q 和 K 的点积得到分数矩阵 S,再对每一行做 Softmax 得到权重矩阵 P,最后用 P 对 V 加权求和得到输出 O。当序列长度为 N 时,S 和 P 的形状都是 N×N。这个平方级的显存增长会迅速吃掉 GPU 的可用空间。以 N=2048 为例,一个注意力头在 FP16 下的 S 矩阵需要 2048×2048×2 字节,约 8MB;N 增加到 4096 时,矩阵大小直接变为约 32MB。多头注意力会进一步放大这个数值,而大型语言模型的层数通常有几十层,因此原始注意力在长序列场景下几乎不可用。

除了显存占用,访存模式也是性能瓶颈。GPU 的高带宽显存 HBM 容量大但带宽相对片上 SRAM 慢得多。标准实现会把完整的 S 和 P 矩阵写入 HBM,再从 HBM 读取出来参与后续计算。这种反复搬运并没有增加计算量,却消耗了大量时间。Flash Attention 的核心观察是:注意力计算并不需要在 HBM 中保存完整的 S 和 P,只需要最终输出 O 和用于反向传播的某些统计量。因此可以通过分块和重算来绕过完整矩阵的物化。

二、Flash Attention的分块计算与在线Softmax

Flash Attention 将 Q、K、V 切分成较小块,每次只把当前需要的块从 HBM 加载到 SRAM 中,在片上完成该块对应的注意力计算,再更新输出。普通 Softmax 需要先知道整行的最大值与所有元素的和,而分块之后某一行的数据被拆到多个块里,无法一次性获得全局最大值与求和。Flash Attention 使用在线 Softmax 技巧维护两个运行状态:当前所见的最大值 m 和累加和 l。当处理一个新的 K、V 块时,先计算局部分数,再用新的最大值修正之前累积的分母,并把旧输出按新旧最大值之差进行缩放。

这种方法在数学上等价于原始 Softmax,但避免了完整 N×N 矩阵的写入和读取。为了支持反向传播,Flash Attention 不缓存中间注意力矩阵,而是选择在前向时保存 Softmax 的归一化统计量,在反向时重算 S 和 P。虽然增加了一些重算代价,但整体显存和带宽收益远大于额外计算,因此训练和推理都能受益。PyTorch 从 2.0 开始集成 Flash Attention,调用方式如下:

import torch
import torch.nn.functional as F

# q, k, v 的形状为 (batch, heads, seq_len, head_dim)
# scaled_dot_product_attention 会根据硬件和输入自动选择 Flash Attention 内核
out = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0, is_causal=True)

如果想了解内部逻辑,下面是一个简化后的分块注意力前向伪代码:

def flash_attention_forward(Q, K, V, scale, block_size=128):
    # Q, K, V: [N, d]
    N = Q.shape[0]
    O = torch.zeros_like(Q)
    m = torch.full((N,), -float('inf'), device=Q.device)
    l = torch.zeros(N, device=Q.device)

    for start in range(0, N, block_size):
        end = min(start + block_size, N)
        Kj = K[start:end]      # [B, d]
        Vj = V[start:end]      # [B, d]
        S = Q @ Kj.T * scale   # [N, B]

        m_new = torch.maximum(m, S.max(dim=-1).values)
        p = torch.exp(S - m_new[:, None])
        l_new = torch.exp(m - m_new) * l + p.sum(dim=-1)

        O = O * torch.exp(m - m_new)[:, None] + p @ Vj
        m = m_new
        l = l_new

    O = O / l[:, None]
    return O

这个实现本质上就是 Flash Attention 想要达到的数值等价效果。它没有产生 N×N 的完整中间矩阵,而是通过循环逐个块更新输出。实际工程实现还会结合张量核心、共享内存分配和寄存器优化,使计算效率接近理论峰值。

三、Paged Attention的KV缓存分页管理

在自回归解码阶段,模型每生成一个新 token 都会依赖之前所有 token 的 Key 和 Value。为了避免重复计算,推理框架会把每一层的 KV 向量保存为缓存,后续步骤直接从缓存读取。问题在于,不同请求的生成长度不同,有的请求可能先结束,有的请求持续生成很长内容。如果为每个请求预分配一段连续的 KV 缓存,短请求结束后会留下无法利用的空洞,新请求又需要新的连续空间。随着并发请求增加,显存碎片会越来越严重,导致可用空间远低于实际剩余空间。

Paged Attention 借鉴操作系统的分页机制,将 KV 缓存划分成固定大小的块,例如每块包含 16 或 32 个 token。每个请求只保留一个页表,记录自己的逻辑位置对应哪些物理块。物理块可以离散地分布在显存中,请求之间不再要求连续分配。当请求长度增加时,只需要分配新的物理块并更新页表。请求结束后,整页可以立即回收给新请求使用。这种管理方式几乎消除了显存碎片,提升了显存利用率和系统吞吐。

下面是一个简化的 KV 缓存块分配器伪代码,体现页表管理的基本思路:

class PagedKVCache:
    def __init__(self, num_blocks, block_size, num_heads, head_dim):
        self.num_blocks = num_blocks
        self.block_size = block_size
        self.num_heads = num_heads
        self.head_dim = head_dim
        # 每个物理块同时保存 K 和 V,形状: [num_blocks, 2, num_heads, block_size, head_dim]
        self.physical_blocks = torch.zeros(num_blocks, 2, num_heads, block_size, head_dim)
        self.free_blocks = list(range(num_blocks))

    def allocate(self, num_tokens):
        needed = (num_tokens + self.block_size - 1) // self.block_size
        if len(self.free_blocks) < needed:
            raise RuntimeError('No enough KV cache blocks')
        blocks = [self.free_blocks.pop() for _ in range(needed)]
        return blocks

    def release(self, blocks):
        self.free_blocks.extend(blocks)

实际系统中的页表还会记录每个块的有效 token 数,并在 attention 计算时根据页表把物理块组织成逻辑序列,再送给 Flash Attention 内核。vLLM 是较早实现 Paged Attention 的推理框架,它可以在接近零浪费的情况下管理动态增长的 KV 缓存,使单一 GPU 能同时服务更多请求。

四、两者协同与工程落地

Flash Attention 与 Paged Attention 的优化目标不同,但可以在推理系统中协同工作。Flash Attention 解决的是单次注意力计算中的显存与带宽压力,让长序列推理的内存峰值显著下降;Paged Attention 解决的是多请求并发时 KV 缓存的组织方式,减少碎片对显存总量的浪费。前者更偏算子级优化,后者更偏系统级内存管理。两者结合后,推理框架可以在有限的显存下容纳更大的 batch size 和更长的上下文窗口。

从工程选择角度看,是否需要两者同时启用取决于部署场景。单用户长文本生成通常首先受益于 Flash Attention,因为它直接降低单次计算的高峰显存;高并发在线服务则更依赖 Paged Attention 的碎片治理。如果一个框架使用 Paged Attention 但底层 attention 仍采用传统算子,虽然缓存分配更紧凑,单个请求的长序列计算依然可能触发显存高峰。因此主流推理方案通常同时集成两种机制。下表简要对比三者的定位:

方案主要目标核心手段适用阶段
标准注意力正确计算输出完整物化 N×N 矩阵训练与推理
Flash Attention降低单次计算 IO分块与在线 Softmax训练与推理
Paged Attention减少 KV 缓存碎片分页映射管理缓存块高并发推理

在实践中,开发人员可以为 PyTorch 训练和推理直接启用 Flash Attention,并在使用 vLLM、TensorRT-LLM 等推理框架时关注 Paged Attention 的块大小与显存占用。理解这两种优化手段的边界,有助于在模型超长上下文、批处理服务和低显存部署之间找到合适的平衡。

Flash AttentionPaged Attention注意力机制优化修改时间:2026-10-01 19:06:53

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/1001/64359.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。