Skip to content

Attention 与 Transformer

目录

1. 学习目标

  • 理解 Query、Key、Value(Q/K/V)如何完成基于内容的加权聚合;
  • 能从张量维度推导 Scaled Dot-Product Attention 与 Multi-Head Attention;
  • 能解释 Mask、残差连接、LayerNorm、前馈网络和位置表示各自解决的问题;
  • 能分析标准 Attention 的时间、空间复杂度以及长上下文瓶颈;
  • 能手写最小因果 Self-Attention,并排查 Mask 方向、Softmax 维度和 NaN 等故障。

2. 面试结论

2.1 30 秒回答

Attention 的本质是“用 Query 与各 Key 的相关性计算权重,再对 Value 加权求和”。Self-Attention 让同一序列的每个位置直接聚合其他位置;Multi-Head Attention 把表示拆到多个子空间并行计算。Transformer 用 Attention 建模跨位置依赖,用前馈网络做逐位置非线性变换,再通过残差、归一化和位置表示保证可训练性与顺序信息。标准 Self-Attention 的分数矩阵是 n×n,因此长序列通常面临 O(n2) 的计算与显存压力。

2.2 一分钟复述版

给定输入 XRB×n×d,分别乘投影矩阵得到 Q、K、V,再拆成 h 个头,每头维度 dh=d/h。每个头计算 S=QK/dh,叠加 Padding 或 Causal Mask,经 Softmax 得到权重,再乘 V。各头拼接并做输出投影,形状回到 [B,n,d]。除 Attention 外,一个 Transformer Block 还包含残差、LayerNorm 和位置前馈网络。因果语言模型用上三角 Mask 禁止当前位置读取未来 Token。工程上除了公式,还要关注 Mask 的布尔语义、全屏蔽行导致 NaN、n2 注意力矩阵导致 OOM,以及训练和推理使用优化内核时的精度与兼容性。

3. 面试官为什么问

  • 原理深度:是否理解相关性、加权聚合和缩放,而不是背诵 Q/K/V 名称;
  • 张量能力:能否推导每一步形状并解释多头如何拆分与合并;
  • 训练理解:是否知道残差、归一化、位置表示和 FFN 缺一不可;
  • 工程能力:能否处理 Mask、数值稳定、显存、长上下文和高效 Attention 内核;
  • 系统连接:能否把 Attention 连接到 KV Cache、吞吐、延迟和模型架构选择。

4. 概念与边界

小白先这样理解:班会上的“多路听取意见”

班级要定春游方案,班长此刻关心“下雨怎么办”,会根据每位同学的擅长判断这次该听谁更多:爱看天气的同学权重高,只谈游戏的权重低,最后按权重汇总他们的实际建议。若平行分组讨论天气、交通和预算,就像从多个角度同时收集信息。

生活角色 → 技术概念: 当前议题是 Query,同学的“擅长标签”是 Key,他们的实际建议是 Value,匹配后的关注比例是 Attention 权重,平行讨论组对应 Multi-Head。

类比边界: Attention 是可微的数值计算,不是人类投票或意识;权重高不能直接证明因果关系。Transformer 还包含 FFN、残差、归一化、位置机制和 Mask,不能用班会故事取代公式与张量边界。

4.1 是什么

Attention 是一个可微分的信息路由机制:查询位置根据内容从一组 Key/Value 中选择并聚合信息。Transformer 是以 Attention 为核心、同时包含前馈网络、残差、归一化和位置机制的序列架构。

4.2 不是什么

  • Q、K、V 不是三份固定语义标签,它们是由输入和可学习矩阵投影得到的角色表示;
  • Attention 权重不等同于可靠的人类可解释性或因果归因;
  • Transformer 不等同于只有 Self-Attention,FFN 往往也占大量参数与计算;
  • Multi-Head 不保证每个头自动学习出人类可命名的功能;
  • FlashAttention 等方法主要优化精确 Attention 的 IO/内存访问,并不等于把数学复杂度普遍降为线性。

4.3 类型与适用边界

类型Q 来源K/V 来源分数形状典型用途
Self-Attention当前序列当前序列[B,h,n,n]编码上下文
Causal Self-Attention当前序列当前序列且屏蔽未来[B,h,n,n]自回归生成
Cross-Attention目标序列外部/源序列[B,h,nq,nk]编码器—解码器、多模态融合

4.4 Transformer 主要形态

  • Encoder-only:通常允许双向关注,适合理解、分类和编码;
  • Decoder-only:使用因果 Mask,自回归预测下一个 Token,是常见 LLM 形态;
  • Encoder-decoder:Encoder 编码源序列,Decoder 同时做因果 Self-Attention 和 Cross-Attention,适合条件生成。

这些是架构边界,不代表某一形态只能完成单一任务。

5. 原理剖析

5.1 直觉理解

读“银行提高了利率,因为它担心通胀”时,模型处理“它”需要从其他位置寻找相关信息。Query 表示“当前要找什么”,Key 表示“每个位置可被怎样匹配”,Value 表示“匹配后实际取回什么信息”。Attention 不移动原始 Token,而是为每个位置重新组合上下文表示。

教学插图:Q、K、V 的动态信息聚合

一个 Query 与多个 Key 形成不同强度的连接,再聚合对应 Value 得到输出

替代文本: 一排 Token 中的当前位置产生 Query;Query 与可见位置的 Key 建立粗细不同的匹配连接,再按照这些权重聚合对应 Value,形成当前位置新的上下文表示。

读图结论: Q/K 匹配决定“从哪些位置取多少信息”,真正被加权汇入输出的是 V;一个位置通常会同时聚合多个可见 Token,而不是只复制最相关的一个 Token。

插图省略了 Head、缩放和 Mask 的精确张量细节;Decoder-only 模型仍必须由 Causal Mask 屏蔽未来位置,而且 Attention 权重不能直接当作因果解释。

5.2 张量维度与核心公式

设输入 XRB×n×d,头数为 h,且 dh=d/h。为便于说明,投影总维度仍为 d

Q=XWQ,K=XWK,V=XWV

其中 WQ,WK,WVRd×d,投影后先得到 [B,n,d],再重排为:

Q,K,VRB×h×n×dh

每个头的缩放点积 Attention:

S=QKdh+M,SRB×h×n×nA=softmax(S,dim=1),O=AV

M 是 Mask:允许的位置加 0,被屏蔽的位置通常加一个足够小的负值。O 的形状为 [B,h,n,dh],拼接为 [B,n,d],再乘 WORd×d

除以 dh 是因为当 Q、K 各维近似独立、均值为 0、方差为 1 时,点积方差随 dh 增长。缩放可避免 Logit 过大使 Softmax 过度饱和、梯度变小。

5.3 Transformer Block

图 1:Pre-LN Decoder-only Transformer Block 数据流替代文本: 输入先经 LayerNorm、因果多头 Self-Attention 和残差,再经 LayerNorm、FFN 和第二次残差,输出到下一层。

图表加载中…

读图结论: Attention 负责跨位置混合信息,FFN 负责每个位置内部的非线性变换,残差和归一化共同维持深层网络的优化稳定性。

原始 Transformer 论文采用 Post-LN 形式;许多后续架构使用 Pre-LN 或其变体。面试中应先说明具体架构,不能把某一种顺序说成所有 Transformer 的唯一实现。

5.4 Mask 的语义

  • Causal Mask:对查询位置 i,屏蔽所有 j>i 的 Key,防止训练时偷看未来;
  • Padding Mask:屏蔽批次中补齐的无效 Key;
  • 两种 Mask 需要正确合并,并匹配框架 API 的布尔约定;不同 PyTorch API 的布尔 Mask 语义可能不同,迁移时必须查目标版本官方文档;
  • 若某一查询行的所有 Key 都被屏蔽,Softmax 可能出现 NaN,入口与 Mask 构造应避免该状态或做明确兜底。

5.5 复杂度与关键假设

忽略常数并设 Q/K/V 总维度为 d

  • Q/K/V 与输出线性投影约为 O(Bnd2)
  • 分数 QK 与加权 Value 约为 O(Bn2d)
  • 朴素实现显式保存注意力分数/概率,内存约为 O(Bhn2)
  • FFN 若中间维度为 dff,计算约为 O(Bnddff)
  • Cross-Attention 的分数成本为 O(Bnqnkd)

“Self-Attention 是 O(n2)”只描述序列长度相关的核心项;当 n 较短而 d 很大时,投影或 FFN 也可能是主要计算。FlashAttention 通过分块和 IO-aware 算法避免把完整注意力矩阵反复写入高带宽内存,在不改变标准 Attention 结果定义的前提下显著降低中间内存访问;具体收益依硬件、形状、精度和实现而定,不能脱离实测给固定倍数。

6. 实现与代码

6.1 最小可运行 PyTorch 示例

运行环境:Python 3.10+、PyTorch 2.x。示例实现单头因果 Self-Attention,重点验证形状、Mask 和概率,不追求生产性能。

python
import math
import torch
import torch.nn as nn


class CausalSelfAttention(nn.Module):
    def __init__(self, d_model: int):
        super().__init__()
        self.q_proj = nn.Linear(d_model, d_model, bias=False)
        self.k_proj = nn.Linear(d_model, d_model, bias=False)
        self.v_proj = nn.Linear(d_model, d_model, bias=False)
        self.out_proj = nn.Linear(d_model, d_model, bias=False)

    def forward(self, x: torch.Tensor):
        # x: [B, n, d]
        q, k, v = self.q_proj(x), self.k_proj(x), self.v_proj(x)
        scores = q @ k.transpose(-2, -1) / math.sqrt(x.size(-1))  # [B, n, n]
        n = x.size(1)
        blocked = torch.ones(n, n, dtype=torch.bool, device=x.device).triu(1)
        scores = scores.masked_fill(blocked, float("-inf"))
        weights = torch.softmax(scores, dim=-1)
        output = self.out_proj(weights @ v)  # [B, n, d]
        return output, weights


torch.manual_seed(0)
x = torch.randn(2, 4, 8)
layer = CausalSelfAttention(d_model=8)
y, attention = layer(x)

assert y.shape == (2, 4, 8)
assert attention.shape == (2, 4, 4)
assert torch.allclose(attention.sum(-1), torch.ones(2, 4), atol=1e-6)
assert torch.count_nonzero(attention[:, 0, 1:]) == 0  # 第一个位置不能看未来
print("output shape:", tuple(y.shape))
print("sample attention:\n", attention[0])

6.2 关键实现说明

  1. k.transpose(-2, -1)[B,n,d] 变为 [B,d,n],得到 [B,n,n] 分数;
  2. triu(1) 生成严格上三角屏蔽区,主对角线保留;
  3. Softmax 必须沿 Key 维,即最后一维执行;
  4. 教学代码是单头、显式注意力矩阵;生产代码优先评估框架官方 scaled_dot_product_attention 等优化实现;
  5. 多头实现还需 view/transpose[B,h,n,dh],并确认 d 可被 h 整除。

6.3 边界条件与验证

  • 输入长度为 1 时,注意力权重应为 1;
  • 加入 Padding Mask 后,所有查询对 Padding Key 的权重应为 0;
  • 比较手写结果与同精度、同 Mask 的官方函数结果;
  • 训练与评估分别检查 Dropout,不能在评估阶段误保留 Attention Dropout;
  • 混合精度下检查 NaN/Inf,并记录触发形状和输入范围。

6.4 技术栈与横向选型

Attention 与 Transformer 的公式不依赖特定框架。最小实验可用 PyTorch 显式计算检查形状和 Mask;生产参考栈才考虑优化 Kernel、模型生态与分布式运行。

技术点 ID技术点/环节类型采用方案链路职责版本/证据边界
TP-AT-01Scaled Dot-Product Attention 执行框架/KernelPyTorch SDPA 作生产参考,显式 MatMul + Softmax 作教学基准校验 Mask 后执行 QK计分、Softmax、Dropout 和 V 聚合SDPA 后端会随软硬件与版本选择;语义和数值必须用固定输入对齐
TP-AT-02Transformer 模型组装框架/库PyTorch + Transformers 作参考栈管理模型配置、权重加载、Block 组装、生成与检查点架构类型、权重格式和精度需共同锁定;不宣称对任意硬件更快
TP-AT-03长序列注意力内存优化Kernel/基础设施优先评估框架 SDPA 的可用后端;FlashAttention 作显式候选减少中间注意力矩阵的显存读写,在语义一致前提下降低内存压力是否可用取决于 GPU、dtype、Mask 形式和安装组合;性能需本机基准
技术点 ID候选方案优点缺点/代价适用场景不适用场景选择结论与依据
TP-AT-01PyTorch scaled_dot_product_attention统一 API 可按条件选后端,减少手写数值错误后端选择受版本、硬件和 Mask 约束,需观测实际路径常规训练和推理、需框架维护兼容性研究某个中间矩阵或定制非标准算子生产默认参考,但先以手写基准验证语义
TP-AT-01显式 MatMul + Mask + Softmax每个中间值可观察,最适合教学和正确性对齐会物化注意力矩阵,通常更占内存且难获得融合优化小张量单元测试、形状排错、算子教学长序列高吞吐生产请求保留为黄金对照,不作默认热路径
TP-AT-02PyTorch + Transformers预训练模型、配置与社区工具链丰富高层封装可能隐藏 Mask、Cache 和生成默认值主流预训练模型微调、推理和应用研发必须用 JAX 变换或特定 TPU 工程栈的项目应用工程参考选择,读取目标版本配置并做端到端回归
TP-AT-02JAX + Flax函数变换、编译与显式并行表达能力强权重生态、调试与团队技能成本可能更高TPU 或 JAX 原生训练平台、大规模研究团队只维护 PyTorch 部署链路且无转换收益只在硬件、并行和团队能力证明收益时选择
TP-AT-03框架 SDPA 优化后端无需直接维护额外 Kernel 接口,升级路径较集中对实际选用的后端控制较间接希望保留框架兼容性的通用负载必须固定某一 Kernel 特性或需特定调优先测默认路径,只有内存或延迟证据不达标才引入额外依赖
TP-AT-03FlashAttentionIO-aware 算法可减少中间读写,长序列常有明显内存价值硬件、dtype、编译和 Mask 支持有约束,升级需回归长序列 GPU 训练或推理,且支持矩阵组合已验证CPU、不支持的 GPU 或需特殊 Mask 语义把它当条件性 Kernel 候选,以数值一致性和本机基准决策

6.5 架构与技术调用流程

图:架构|Transformer Block 与执行后端边界

替代文本: 模型配置与权重进入 Transformer Block,Block 内的 LayerNorm、QKV 投影、Attention Kernel、残差和 FFN 组成模块边界;Kernel Dispatcher 根据设备、精度和 Mask 选执行后端,数值对照与性能基准分别验收正确性和效率。

图表加载中…

读图结论: Transformer 语义由模块结构、权重和 Mask 决定,优化 Kernel 只是可替换的执行层,不应改变模型语义。

架构图将“模型结构”与“Kernel 选型”分开:先用参考路径验证正确性,再在相同输入、dtype 和 Mask 下比较速度与显存。

图:技术调用流程|Attention 计算与 Kernel 降级时序

替代文本: Block 先校验张量与 Mask,Dispatcher 尝试选用优化 Kernel;条件不支持时降级到框架参考后端,语义或数值异常则立即失败,成功时才返回残差输出。

图表加载中…

读图结论: “后端不支持”可降级,“Mask 错误或数值异常”必须阻断;两类失败不能用同一重试逻辑处理。

时序图展示了执行层的条件分支。生产排障时应记录实际 Kernel、设备、dtype、Mask 类型和形状,否则无法复现“同代码不同性能”。

7. 实际项目案例

示例项目,非真实仓库实现;所有性能和质量结论必须通过后续实测获得。

7.1 背景、目标与约束

为内部代码审查助手构建 Decoder-only 模型推理链路。输入包含系统规则、代码片段和问题,要求模型逐 Token 输出分析。约束是代码上下文长、不同请求长度差异大、不能读取被 Padding 的位置、线上显存有限。

7.2 架构与调用链

  1. Tokenizer 产生 input_ids、Padding Mask 和位置;
  2. Prefill 阶段对完整 Prompt 做因果 Self-Attention;
  3. 每层保存已计算的 K/V,Decode 阶段新 Query 只与已有 K/V 计算;
  4. 服务层按长度和显存预算调度请求,流式返回 Token;
  5. 记录请求长度、生成长度、TTFT、逐 Token 延迟、KV Cache 占用和 OOM。

7.3 方案选择与实现难点

  • 使用因果 Self-Attention 保证自回归约束;代码补全任务若需要特殊双向上下文,应明确采用与模型训练匹配的格式或架构;
  • 长请求先做 Token 预算与截断策略,不能等到 Attention 分配 n2 中间量时才失败;
  • 优化内核只能在支持的设备、dtype、形状和 Mask 条件下启用,应保留正确性基线;
  • KV Cache 优化推理重复计算,但不会减少 Prefill 的完整上下文 Attention,也会占用随序列增长的显存。

7.4 异常处理、监控与测试

现象可能根因解决方式验证证据
训练指标异常好,但自回归生成崩坏Causal Mask 方向反了,训练时偷看未来修正上三角屏蔽并锁定统一 Mask 构造函数可预测序列上未来权重为 0,逐步生成与全量因果前向一致
输出或梯度出现 NaN某个查询行全屏蔽,或低精度 Logit 溢出保证至少一个有效 Key,检查有限值并在必要时切换稳定内核/精度空输入、极值与长序列回归中全部张量有限
Padding 长度改变真实 Token 输出只应用 Causal Mask,遗漏 Padding Mask 或布尔语义用反合并两类 Mask,并按目标 API 写真值测试同样本在不同 Padding 长度下的有效位置输出满足容差
长请求 OOMn2 中间量、KV Cache 或并发超预算入口做 Token/KV 准入,启用合适高效内核、限流或拒绝长度—批次压测无进程崩溃,峰值显存低于门槛
延迟剧烈抖动长短请求混批、优化内核静默回退长度感知调度并记录实际内核、dtype 与批次形状按长度和内核切片的 P95/P99 回到既定门槛

7.4.1 故障演练:长请求触发 OOM 并拖垮同批短请求

  • 现象与影响:少量长上下文请求使显存峰值越界,Worker 重启,同批短请求超时并造成队列重试放大。
  • 定位证据:关联输入长度、batch 形状、Prefill/Decode 阶段、KV Cache 占用、实际 Attention Kernel、OOM 日志与队列深度。
  • 根因:入口只限制字符数,没有按 Token、并发和 KV Cache 做容量准入;长短请求混批放大 (n^2) 中间量和尾延迟。
  • 临时止损:拒绝或截断超预算请求,降低并发并按长度拆批;关闭导致不稳定的优化内核,回退已验证路径。
  • 长期修复:建立 Token/KV 容量模型、长度感知调度和每租户预算;对支持矩阵内的高效 Attention Kernel 做数值与性能双门禁。
  • 回归验证:在目标硬件上压测长度与批次矩阵,验证无进程崩溃、输出容差通过,TTFT、逐 Token 延迟、P99 与峰值显存均在门槛内。
  • 防复发:监控长度分桶、KV 使用率、Kernel 回退、OOM、重启和重试率;新模型或上下文窗口升级先执行容量回归。

7.5 结果与复盘

上线前必须以同一输入验证基线与优化内核的容差、因果性、Padding 不变性和流式结果一致性。性能只报告目标硬件和固定配置下的 TTFT、逐 Token 延迟、吞吐、峰值显存及分位数,不引用无上下文的宣传倍数。

8. 方案权衡与常见误区

8.1 适用与不适用场景

  • 全局 Attention 适合需要任意位置直接交互的序列;上下文极长且局部性强时,可评估滑窗、稀疏、分层或检索增强,但要验证丢失全局依赖的代价;
  • Decoder-only 适合自回归生成;双向编码或严格条件生成可能更适合其他形态;
  • Attention 权重可以辅助调试,但不能单独作为模型决策的因果解释。

8.2 替代与优化方案

方案主要收益主要代价
FlashAttention/融合 SDPA降低中间 IO 与显存,保持精确 Attention 定义有设备、dtype、形状和版本约束
滑窗 Attention把每层关注范围限制在局部单层缺失远距离直接连接
稀疏/块稀疏 Attention减少部分 n2 计算模式设计、内核和质量验证复杂
Cross-Attention清晰分离查询和外部上下文需要额外 K/V 与架构支持
RAG 缩短上下文只把相关证据送入模型依赖检索质量,可能漏召回

8.3 常见错误回答

  • “除以 d”却没有说明实际是每头维度 dh
  • 只说 Attention 是 O(n2),忽略投影、FFN、批大小、头数与内存;
  • 把 Causal Mask 说成删除未来 Token,而非在分数/内核中禁止关注;
  • 认为多头只是重复计算同一份 Attention;
  • 认为 KV Cache 把整个生成复杂度变成常数;每个新 Token 仍需读取/关注不断增长的历史 K/V。

8.4 生产环境风险

  • Mask API 语义变化或混用:目标版本上写小型真值测试,不靠记忆;
  • 优化内核静默回退:记录设备、dtype、形状与实际执行路径;
  • 长上下文尾延迟/OOM:入口做 Token 和 KV 预算,采用长度感知调度;
  • 低精度数值问题:在代表性长序列上做有限值、容差和质量回归;
  • 错误归因:输出差异可能来自位置编码、Tokenizer、采样或 Cache,不应一律归咎于 Attention。

9. 面试题与参考答案

问题 1:Q、K、V 分别起什么作用?

  • 难度:基础;
  • 考察点:Attention 直觉与角色;
  • 合格答案要点:Q 发起匹配,K 被匹配,V 是按权重聚合的信息;
  • 优秀答案加分项:强调它们来自可学习投影,同一输入可承担不同角色;
  • 常见错误:把 Q/K/V 当固定保存的三类词义;
  • 可继续追问:Self-Attention 与 Cross-Attention 的来源有何不同?

问题 2:为什么点积要除以 dh

  • 难度:中级;
  • 考察点:方差、Softmax 饱和和梯度;
  • 合格答案要点:点积方差随维度增大,缩放稳定 Logit 尺度;
  • 优秀答案加分项:写出独立单位方差假设下点积方差约为 dh
  • 常见错误:说“为了归一化到 0~1”;
  • 可继续追问:为什么不是除以 dh

问题 3:Self-Attention 的复杂度如何推导?

  • 难度:中级;
  • 考察点:张量维度与瓶颈判断;
  • 合格答案要点:分数和聚合为 O(Bn2d),朴素分数内存为 O(Bhn2)
  • 优秀答案加分项:补充投影 O(Bnd2) 与 FFN,并说明不同形状下主导项不同;
  • 常见错误:只背 O(n2) 而不能说分数矩阵形状;
  • 可继续追问:FlashAttention 改变了哪个瓶颈?

问题 4:线上出现长输入 OOM,如何系统排查?

  • 难度:高级;
  • 考察点:端到端工程诊断;
  • 合格答案要点:确认输入/生成长度、批大小、dtype、Attention 中间量、KV Cache、并发与内核路径;
  • 优秀答案加分项:区分 Prefill 和 Decode,提出准入预算、降级、压测与分位数监控;
  • 常见错误:直接减小模型而不定位显存构成;
  • 可继续追问:如果降低批大小后仍 OOM,下一步如何分解验证?

10. 递进追问

  1. 基础概念:Self-Attention、Cross-Attention 和 Causal Attention 有什么区别?
  2. 公式推导:从 [B,n,d] 输入推到 [B,h,n,n] 权重矩阵,每一步形状是什么?
  3. 实现细节:Causal Mask 与 Padding Mask 如何组合,Softmax 应沿哪一维?
  4. 数值边界:为什么全屏蔽行可能产生 NaN,怎样写测试发现它?
  5. 工程权衡:FlashAttention、滑窗 Attention 和 RAG 缩短上下文分别解决什么瓶颈?
  6. 系统复盘:如果优化内核上线后输出有小幅差异且 P99 延迟下降,你用什么质量和性能证据决定是否保留?

11. 实践任务

  • [ ] 最小实现:运行本文单头代码,再扩展到多头并打印所有中间形状;
  • [ ] 正确性对比:与 PyTorch 官方 scaled_dot_product_attention 在相同 Mask 下比较;
  • [ ] 故障注入:反转 Mask、改错 Softmax 维度、制造全屏蔽行,记录现象与根因;
  • [ ] 复杂度实验:逐步增加序列长度,记录运行时间和峰值内存,只比较同一硬件与配置;
  • [ ] 架构实验:实现一个 Pre-LN Block,验证输入输出形状与梯度有限;
  • [ ] 面试口述:不用看文档,在白板上完成 Q/K/V 公式和维度推导。

12. 相关知识与参考资料

12.1 相关知识

12.2 一手参考资料

以下资料均于 2026-07-10 访问;PyTorch API 与可用内核会演进,应以目标安装版本文档为准。

  1. Vaswani et al., Attention Is All You Need,Transformer 原始论文;
  2. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness,IO-aware 精确 Attention 论文;
  3. Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning,FlashAttention-2 论文;
  4. PyTorch, torch.nn.functional.scaled_dot_product_attention 官方文档
  5. PyTorch, torch.nn.MultiheadAttention 官方文档
  6. Ba et al., Layer Normalization,LayerNorm 原始论文。

13. 简明总结

一句话记忆: Attention 用 Q 匹配 K、再按权重汇总 V;Transformer 用它混合位置,并靠 FFN、残差、归一化和位置机制组成可训练的序列模型。

  • 单头分数为 QK/dh,多头权重形状是 [B,h,n,n]
  • Causal Mask 禁止看未来,Padding Mask 排除补齐位置,Mask 语义必须按具体 API 验证;
  • 标准 Attention 的序列相关计算为 O(Bn2d),朴素注意力矩阵内存为 O(Bhn2)
  • 项目易错点是 Mask 方向、Softmax 维度、全屏蔽 NaN、优化内核回退和长上下文 OOM;
  • 面试不仅要背公式,还要能推维度、拆复杂度并连接到 KV Cache 与服务指标。