跳到主要内容
极客日志极客日志面向AI+效率的开发者社区
首页博客我的书AI学习GitHub 精选镜像AI 生图工具UI配色美学关于
搜索内容 / 工具 / 仓库 / 镜像...⌘K搜索
注册
博客列表
PythonAI算法

LLM 核心技术:Attention 机制的实现与优化

详细解析了 LLM 中 Attention 机制的核心原理及多种优化方案。首先阐述了 Multi-Head Attention (MHA) 的计算复杂度与显存瓶颈,随后介绍了 MQA 和 GQA 如何通过共享 KV 矩阵降低推理显存占用。接着分析了 Sliding Window Attention (SWA) 对长文本复杂度的优化。在底层实现方面,重点讲解了 FlashAttention 如何利用 SRAM 分块计算减少 HBM 访问,以及 PagedAttention 如何通过虚拟内存管理解决 KV Cache 的碎片化和显存浪费问题。这些优化技术共同提升了大模型的训练效率和推理吞吐量。

莫名其妙发布于 2025/2/6更新于 2026/9/1155 浏览
LLM 核心技术:Attention 机制的实现与优化

LLM 核心技术:Attention 机制的实现与优化

背景介绍

在大型语言模型(LLM)的架构中,Attention 机制是核心组件,决定了模型处理序列数据的能力。随着上下文窗口(Context Window)的不断扩展,传统的 Attention 实现方式面临着计算复杂度高、显存占用大等挑战。本文深入探讨 Multi-Head Attention (MHA) 的原理及其多种优化方案,包括 MQA、GQA、SWA、FlashAttention 和 PagedAttention,旨在提升模型训练与推理的性能。

Multi-Head Attention (MHA)

原理与计算流程

MHA 的目标在于重构文本中的 Token Embedding 表示,使其能够捕捉上下文语义相关性和位置相关性。其计算过程主要包含以下步骤:

  1. Embedding Lookup:输入文本长度为 n(n 个 token),经过 Embedding Table 后,每个 token 返回一个大小为 (1, d) 的向量。对于长度为 n 的文本,生成 Embedding Matrix,大小为 (n, d),其中 d 为 Embedding 维度。
  2. 线性映射:Embedding Matrix X 进入 MHA 层后,通过线性变换生成 Query (Q)、Key (K)、Value (V)。假设 Head 数量为 h,每个 Head 的维度为 k 或 v(通常 k=v=d/h)。
    • Q 维度:(h, n, k)
    • K 维度:(h, n, k)
    • V 维度:(h, n, v)
  3. Attention 计算:对每个 Head 执行 Softmax(QK^T / sqrt(d_k)) * V。由于 n 个 Token 两两交互,时间复杂度为 O(n^2)。

复杂度分析

MHA 的整体计算复杂度与上下文长度 n 的二次方成正比,与模型规模 d 的二次方成正比。公式如下:

$$ \text{Complexity} = O(n^2 \cdot d) $$

增大 Context 长度会带来计算复杂度的二次方增长,这限制了长文本的处理能力。同时,在自回归推理过程中,为了加速解码,需要缓存之前生成的 Key 和 Value (KV Cache),导致显存占用随序列长度线性增加。

推理优化:MQA 与 GQA

Multi-Query Attention (MQA)

在标准 MHA 中,每个 Head 都有独立的 K 和 V 矩阵。在推理时,GPU 显存占用会随着预测 Token 数目增加而累积。

MQA 通过在不同 Head 间共享 K 和 V 矩阵来优化。即所有 Head 使用同一组 Key 和 Value,仅 Query 独立。这使得存储的 K/V 矩阵数量从 2h 降低为 2 个。虽然显著降低了显存占用并提高了推理速度,但可能因信息压缩导致精度略有下降。

Group Query Attention (GQA)

GQA 是对 MQA 的改进,它在 Head 之间进行分组。一个 Group 内的多个 Head 共享一组 K 和 V,不同 Group 之间则独立。这种方式在保持接近 MHA 效果的同时,大幅减少了 KV Cache 的大小。

GQA 实现逻辑
# 初始化投影层
self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)

# 前向传播中的重复操作
# key_states 和 value_states 初始形状为 (batch, num_key_value_heads, seqlen, head_dim)
# 需要通过 repeat_kv 将其扩展到 (batch, num_attention_heads, seqlen, head_dim)
key_states = repeat_kv(key_states, self.num_key_value_groups)
value_states = repeat_kv(value_states, self.num_key_value_groups)
内存开销对比
  • MHA 显存开销:batch * max_seq_len * n_heads * head_dim * sizeof(half) * 2
  • GQA 显存开销:batch * max_seq_len * n_kv_heads * head_dim * sizeof(half) * 2

其中 n_heads / n_kv_heads 即为 Group 大小。使用 GQA 可将 KV Cache 显存降低到 MHA 的 1/group 水平,极大缓解了访存密集型计算的瓶颈。

窗口注意力:Sliding Window Attention (SWA)

为了进一步降低 Attention 与 Context Length 的依赖关系,SWA 限制了每个 Token 只能关注前 W 个 Token。这将时间复杂度从 $O(n^2)$ 降低至 $O(n \cdot w)$。

在 SWA 中,注意力的传递通过层数增加向后延伸。每一层注意力层允许信息传递 W tokens。例如,对于 16k 序列长度和 4k 滑动窗口,通过 4 层即可实现整个序列信息的传递。这种机制特别适合长文本场景,但需注意长距离依赖可能丢失的问题。

底层算子优化:FlashAttention

FlashAttention 旨在解决 Attention 计算过程中频繁访问 HBM(High Bandwidth Memory)的问题。它利用 GPU 中 SRAM(片上高速缓存)速度快但容量小的特点,将 Attention 计算 Block 化,直接在 SRAM 中进行。

核心优化点

  1. IO 感知计算:减少 HBM 读写次数。传统方法需要将中间结果写入 HBM,FlashAttention 通过分块计算避免此步骤。
  2. 分块 Softmax:为了保证分块计算的 Softmax 值与原值一致,采用了在线 Softmax 算法,结合重计算(Recomputation)策略。
  3. 并行计算:将 Q、K、V 切分为 Block,提高 Operation 的处理效率。

FlashAttention 不仅提升了训练速度,也显著改善了推理延迟,是目前主流的大模型框架(如 PyTorch FSDP)默认启用的优化方案。

内存管理优化:PagedAttention

PagedAttention 解决了 Attention 计算过程中的内存分配问题,特别是针对 KV Cache 的动态变化特性。

传统 KV Cache 问题

  1. 显存占用大:大型模型单个序列的 KV Cache 可能占用高达 GB 级显存。
  2. 动态变化:序列长度不可预测,难以预分配。
  3. 内存碎片化:静态批处理策略下,请求结束后剩余空间浪费严重(内部碎片);连续内存分配要求导致外部碎片。

PagedAttention 优势

PagedAttention 引入了虚拟内存的概念,允许 KV Cache 在非连续的物理内存块中存储。系统维护一个页表来映射虚拟地址到物理地址。

  • 非连续存储:解决了连续内存分配造成的空间浪费。
  • 动态分配:支持小空间内存的动态分配,适应不同长度的序列。
  • 高吞吐量:通过减少碎片和优化内存利用率,可以实现更大的 Batch Size 和更高的吞吐量。

总结

Attention 机制的优化是大模型性能提升的关键路径。从算法层面的 MQA/GQA 减少参数量和显存,到系统层面的 FlashAttention 优化 IO,再到内存管理的 PagedAttention 解决碎片问题,这些技术共同推动了 LLM 向更长上下文、更低成本的方向发展。在实际工程落地中,应根据硬件资源和业务需求选择合适的优化组合。


参考文献

  • [1] FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv preprint.
  • [2] PagedAttention: Virtual Memory for Efficient LLM Inference. arXiv preprint.

目录

  1. LLM 核心技术:Attention 机制的实现与优化
  2. 背景介绍
  3. Multi-Head Attention (MHA)
  4. 原理与计算流程
  5. 复杂度分析
  6. 推理优化:MQA 与 GQA
  7. Multi-Query Attention (MQA)
  8. Group Query Attention (GQA)
  9. GQA 实现逻辑
  10. 初始化投影层
  11. 前向传播中的重复操作
  12. keystates 和 valuestates 初始形状为 (batch, numkeyvalueheads, seqlen, headdim)
  13. 需要通过 repeatkv 将其扩展到 (batch, numattentionheads, seqlen, headdim)
  14. 内存开销对比
  15. 窗口注意力:Sliding Window Attention (SWA)
  16. 底层算子优化:FlashAttention
  17. 核心优化点
  18. 内存管理优化:PagedAttention
  19. 传统 KV Cache 问题
  20. PagedAttention 优势
  21. 总结

更多推荐文章

查看全部
  • Python 量化交易实盘部署与风险管理实战
  • Flutter 技术优势显著但市场普及度为何滞后?
  • 中国人工智能大模型技术白皮书:核心技术与应用前景解读
  • 构建并微调大型语言模型实现文本分类任务
  • 基于 StructBERT 的零样本中文文本分类方案与 WebUI 实现
  • Agent 框架大比拼:19 种主流框架优劣分析
  • LLM Agent 工作流 Prompt 设计精解:规划、反思与工具调用
  • Python 环境安装与配置 Pandas 库指南
  • IntelliJ IDEA 运行 JUnit 报 NoSuchMethodError 错误排查
  • Stable Diffusion LoRA 模型高效微调实战指南
  • 腾讯游戏 2026 年 Q1 财报解读:AI 驱动增长与全球化布局
  • AI 音频工具合集
  • 微信小程序全局配置 window 属性详解及常见误区
  • LeetCode 142. 环形链表 II(C 语言实现)
  • MySQL 错误 1130 解决方案:Host 不被允许连接服务器
  • Claude Code 高级编程技巧实战项目详解
  • Whisper-large-v3 功能测评:多语言语音识别真实表现
  • 斯坦福 2025 AI 指数报告深度解读:从技术突破到产业扩散
  • MATLAB 实现基于多目标粒子群算法(MOPSO)的无人机三维路径规划
  • 使用 Python 和 Flask 构建简易 TODO 任务管理系统

相关免费在线工具

  • 加密/解密文本

    使用加密算法(如AES、TripleDES、Rabbit或RC4)加密和解密文本明文。 在线工具,加密/解密文本在线工具,online

  • RSA密钥对生成器

    生成新的随机RSA私钥和公钥pem证书。 在线工具,RSA密钥对生成器在线工具,online

  • Mermaid 预览与可视化编辑

    基于 Mermaid.js 实时预览流程图、时序图等图表,支持源码编辑与即时渲染。 在线工具,Mermaid 预览与可视化编辑在线工具,online

  • 随机西班牙地址生成器

    随机生成西班牙地址(支持马德里、加泰罗尼亚、安达卢西亚、瓦伦西亚筛选),支持数量快捷选择、显示全部与下载。 在线工具,随机西班牙地址生成器在线工具,online

  • Gemini 图片去水印

    基于开源反向 Alpha 混合算法去除 Gemini/Nano Banana 图片水印,支持批量处理与下载。 在线工具,Gemini 图片去水印在线工具,online

  • curl 转代码

    解析常见 curl 参数并生成 fetch、axios、PHP curl 或 Python requests 示例代码。 在线工具,curl 转代码在线工具,online