
1. 注意力机制的直觉拆解与数学骨架注意力机制attention这个词我第一次真正被它绊住是读 Transformer 论文的时候。之前做序列任务脑子里全是 RNN 那套“上一步的隐状态传给下一步”的流水线思路看到 attention 直接把所有时间步拉平做加权求和第一反应是“这不是把顺序信息丢了吗”。后来才想明白注意力机制本质上只解决一件事当模型需要输出某个位置的表示时它应该回头去看输入序列里的哪些部分以及看多重。这个“看多重”就是权重而权重怎么算、算完怎么用决定了整套机制的脾气和适用边界。这篇笔记不打算复述教材而是把我自己从手写单头 attention、到多头、到视觉里的 SE/CBAM/CA、再到部署阶段折腾 Flash Attention 和 Sage Attention 的整条链路捋一遍。中间会穿插公式手算、PyTorch 实现、显存实测以及在 ComfyUI 这类推理前端里装加速库踩过的坑。适合已经会写nn.Linear、但被各种 attention 变体名字绕晕的人也适合只想知道“这东西到底在算什么”的初学者。1.1 从查字典说起Q、K、V 到底是什么把 attention 类比成查字典最省事。你手上有个词要去查Query查询字典里每个词条有个标题Key键标题下面有释义Value值。你拿手上的词和每个词条的标题做匹配匹配度高的词条它的释义就更多进入你的理解里。整个过程的输出不是某一个词条的释义而是所有释义按匹配度加权混合的结果。映射回张量假设输入序列长度是 $n$每个位置的向量维度是 $d$。经过三个不同的线性层得到 $Q \in \mathbb{R}^{n \times d_k}$、$K \in \mathbb{R}^{n \times d_k}$、$V \in \mathbb{R}^{n \times d_v}$。注意这里的“三个线性层”不是装饰它们是让模型自己学会“我该拿什么去查”“我该用什么被查”“我该返回什么内容”的关键参数。很多人第一次实现会偷懒让 Q、K、V 都等于同一个输入张量结果发现模型也能训就以为线性投影可省。这是个误区不投影的话匹配度和内容被绑死成同一个空间的度量模型失去自由度在自注意力里 Q/K/V 同源不同投影是保证表达力的前提。我试过在一个小规模机器翻译任务里把投影去掉BLEU 直接掉了 3 个点以上。1.2 缩放点积注意力的公式与手算过程标准写法是Attention(Q, K, V) softmax(Q K^T / sqrt(d_k)) V拆开看三步我用一个 $n3$、$d_k2$ 的小例子手算一遍理解会牢得多。设import torch import torch.nn.functional as F Q torch.tensor([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]) K torch.tensor([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]) V torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) scores Q K.T # 第一步相似度打分 scores scores / (2 ** 0.5) # 第二步缩放 weights F.softmax(scores, dim-1) # 第三步归一化成权重 out weights V # 第四步加权求和第一步 $QK^T$ 得到的是 $3 \times 3$ 的分数矩阵第 $i$ 行第 $j$ 列表示第 $i$ 个 Query 和第 $j$ 个 Key 的内积。内积越大代表方向越一致、越“像”。第二步除以 $\sqrt{d_k}$第三步沿最后一维 softmax让每一行的权重和为 1。第四步拿权重去乘 V得到每个位置的新表示。这里有个容易忽略的细节softmax 是按行做的也就是每个 Query 各自拥有一套对全部 Key 的分配方案。这意味着 attention 的输出长度和 Query 数量一致而和 Key/Value 数量可以不同——这正是解码器里交叉注意力的基础Query 来自解码器当前步Key/Value 来自编码器全部输出。1.3 为什么必须除以 sqrt(d_k)一次数值实验教科书说“防止内积过大导致 softmax 梯度消失”听起来抽象。实际做一次实验就懂了。假设 Q、K 的每个分量都独立服从均值 0、方差 1 的分布那么内积 $\sum_{i1}^{d_k} q_i k_i$ 的方差就是 $d_k$。$d_k 64$ 时分数的标准差约是 8$d_k 512$ 时标准差约 22.6。分数跨度一大softmax 就会变成近似 one-hot 的分布最大值位置拿到接近 1 的权重其余接近 0。后果是反向传播时除最大值位置外的梯度几乎为零参数更新停滞。除以 $\sqrt{d_k}$ 恰好把方差压回 1 附近让分布保持“柔软”。我用同一份数据、同一组初始化只改是否缩放跑 200 步看 loss 曲线不缩放的版本在 $d_k512$ 时前 50 步几乎不动缩放版本稳定下降。注意如果你自己在实现里用了自定义的 Q/K 初始化放大了初始方差缩放因子要相应调整。我遇到过一次把 K 的初始化标准差设成 0.1 而非默认值配合不缩放反而收敛更快原因就是初始分数尺度被压小了。别把公式当教条看实际数值范围。2. 自注意力与多头注意力Transformer 的心脏自注意力机制self-attention指的是 Q、K、V 全部来自同一个序列序列里每个位置都去和包括自己在内的所有位置做匹配。这带来一个直接后果任意两个位置之间的信息传递路径长度是 1不再像 RNN 那样随距离线性增长。长距离依赖在梯度上变得可训练这是 Transformer 能取代循环结构的核心原因。但代价也很明显计算复杂度是 $O(n^2 d)$内存同样是 $O(n^2)$因为要显式存那个 $n \times n$ 的分数矩阵。序列长度从 512 涨到 4096分数矩阵的元素数量涨了 64 倍。这就是后来 Flash Attention 这类 IO 感知算法出现的直接动机后面第 4 章细讲。2.1 位置编码自注意力丢掉的顺序信息怎么补回来自注意力对输入做的是集合式的加权聚合打乱位置顺序输出只是跟着换行语义上完全等价。所以必须显式注入位置信息。主流做法有两类绝对位置编码正弦函数或可学习 embedding和相对位置编码RoPE、ALiBi。正弦编码的形式是 $PE_{(pos, 2i)} \sin(pos / 10000^{2i/d})$偶数维用 sin奇数维用 cos。选这个形式的原因是它满足一个漂亮的性质位置 $pos k$ 的编码可以表示成位置 $pos$ 编码的线性变换模型理论上能学到“相对位移”。可学习 embedding 更简单直接nn.Embedding(max_len, d_model)缺点是无法外推到训练时没见过的长度。RoPE 现在在大模型里更常见。它的做法不是加在输入上而是把 Q、K 按二维一组做旋转旋转角度和位置相关。这样一来两个位置的内积只依赖它们的相对距离。我在小模型上对比过三种方案训练长度 512、测试长度 1024 的外推场景下正弦编码性能衰减最严重RoPE 相对最稳。2.2 多头注意力机制原理拆开、并行、再拼回去多头自注意力机制原理其实一句话能说完把 $d_{model}$ 维的 Q/K/V 切成 $h$ 份每份独立做一次注意力最后把 $h$ 个输出拼接再过一个线性层。切分让每个头在自己的子空间里关注不同的模式——有的头盯语法依赖有的头盯相邻词有的头几乎只关注自己注意力图上看是对角线。维度关系必须记牢$d_k d_v d_{model} / h$这样拼接后总维度才回到 $d_{model}$。$h8$、$d_{model}512$ 时每个头是 64 维。有人会问切成 8 份每个头只有 64 维表达能力不是变弱了吗单头确实弱但 8 个头关注不同子空间总容量并没减少而且参数量和多头共享投影的情况是一致的。class MultiHeadAttention(torch.nn.Module): def __init__(self, d_model512, num_heads8, dropout0.1): super().__init__() assert d_model % num_heads 0 self.d_model d_model self.h num_heads self.d_k d_model // num_heads self.w_q torch.nn.Linear(d_model, d_model) self.w_k torch.nn.Linear(d_model, d_model) self.w_v torch.nn.Linear(d_model, d_model) self.w_o torch.nn.Linear(d_model, d_model) self.dropout torch.nn.Dropout(dropout) def forward(self, q, k, v, maskNone): B, Lq, _ q.shape Lk k.shape[1] # 投影 拆头(B, L, d_model) - (B, h, L, d_k) Q self.w_q(q).view(B, Lq, self.h, self.d_k).transpose(1, 2) K self.w_k(k).view(B, Lk, self.h, self.d_k).transpose(1, 2) V self.w_v(v).view(B, Lk, self.h, self.d_k).transpose(1, 2) scores Q K.transpose(-2, -1) / (self.d_k ** 0.5) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn self.dropout(torch.softmax(scores, dim-1)) out attn V # (B, h, Lq, d_k) out out.transpose(1, 2).contiguous().view(B, Lq, self.d_model) return self.w_o(out)这段代码里有两个坑。第一个是transpose之后必须contiguous()才能view否则报错或者得到错误的内存布局。第二个是 mask 的形状PyTorch 广播要求它能匹配 $(B, h, L_q, L_k)$我习惯在调用前把 mask 统一成 $(B, 1, 1, L_k)$省掉一堆 shape 报错。2.3 掩码与 seq2seq 解码器的通用注意力模块在 seq2seq 里注意力模块有两种角色。编码器自注意力不加掩码所有位置互相可见解码器自注意力必须加因果掩码位置 $i$ 只能看到 $j \le i$否则训练时模型会偷看未来的 token训练 loss 好看但推理时完全崩掉。交叉注意力的 Query 来自解码器Key/Value 来自编码器输出掩码只针对编码器侧的 padding。写一个通用 decoder attention module 时我建议把三件事参数化is_causal、q_len/kv_len、key_padding_mask。这样同一份代码能覆盖编码器自注意力、解码器自注意力、交叉注意力三种场景。def build_causal_mask(L, device): # 下三角为 True表示允许被看见 return torch.tril(torch.ones(L, L, dtypetorch.bool, devicedevice))因果掩码的常见错误是用float(-inf)填充时把整行都填满导致 softmax 出现全 $-\infty$输出 NaN。稳妥做法是保证对角线至少为 0或者在 softmax 之前检查“每行是否至少有一个非掩码位置”。我在调试一个流式解码任务时就因为在序列全 padding 的 batch 上触发过这个 NaN排查了半天。3. 视觉注意力模块三兄弟SE、CBAM、CA视觉里的注意力模块和 NLP 的 attention 名字一样思路却不同。图像任务里大家想要的往往不是“序列位置之间互相看”而是“让网络自己判断哪些通道重要、哪些空间区域重要”。于是就有了 SE、CBAM、CA 这一系列轻量模块它们通常插在 backbone 的残差分支上参数量增加极小却能稳定涨点。3.1 SE 通道注意力把全局信息压成一个权重向量SESqueeze-and-Excitation是通道注意力的起点。流程是三步先对每个通道做全局平均池化Squeeze把 $H \times W$ 压成一个标量再经过两层全连接加激活中间有个降维比例 $r$常用 16最后用 Sigmoid 得到 $C$ 维权重乘回原特征。class SEBlock(torch.nn.Module): def __init__(self, channels, reduction16): super().__init__() self.pool torch.nn.AdaptiveAvgPool2d(1) self.fc torch.nn.Sequential( torch.nn.Linear(channels, channels // reduction), torch.nn.ReLU(inplaceTrue), torch.nn.Linear(channels // reduction, channels), torch.nn.Sigmoid() ) def forward(self, x): b, c, _, _ x.shape w self.pool(x).view(b, c) w self.fc(w).view(b, c, 1, 1) return x * w瓶颈结构的设计意图是控制参数。如果不降维两层全连接是 $C^2$ 参数$C512$ 时超过 26 万而加个 $r16$ 后降到约 3.3 万。实测下来这个降维对手持设备的推理延迟也很友好。要注意的是AdaptiveAvgPool2d(1)会丢掉所有空间位置信息SE 对“目标在哪里”是无感的它只知道“哪些通道整体活跃”。3.2 CBAM通道与空间的串联组合CBAMConvolutional Block Attention Module在 SE 的基础上加了一个空间注意力分支顺序是先通道后空间。通道分支和 SE 略有不同它同时用平均池化和最大池化各自过一个小 MLP 后相加再 Sigmoid。空间分支则是在通道维度上做平均池化和最大池化得到两个 $H \times W$ 图拼成 2 通道后用一个大核卷积通常是 7×7压成 1 通道再 Sigmoid。class SpatialAttention(torch.nn.Module): def __init__(self, kernel_size7): super().__init__() self.conv torch.nn.Conv2d(2, 1, kernel_size, paddingkernel_size // 2, biasFalse) self.sigmoid torch.nn.Sigmoid() def forward(self, x): avg_out torch.mean(x, dim1, keepdimTrue) max_out, _ torch.max(x, dim1, keepdimTrue) cat torch.cat([avg_out, max_out], dim1) return x * self.sigmoid(self.conv(cat))用大核 7×7 而不是 3×3是为了让空间注意力有更大的感受野能覆盖到整块目标区域而不是局部纹理。代价是插在浅层高分辨率特征上时计算量不小。我给的经验是CBAM 放在 backbone 的 stage3 之后收益最明显放 stage1 高分辨率处性价比低。3.3 CA 注意力把坐标信息塞回通道注意力里CACoordinate Attention针对的就是 SE 丢空间信息这个短板。它的做法是把全局池化拆成两个方向沿宽度方向做池化得到 $C \times H \times 1$沿高度方向做池化得到 $C \times 1 \times W$。这两份特征分别编码了“在哪一行”和“在哪一列”的坐标信息。然后拼接、过共享卷积降维、再拆开各自通过卷积恢复通道数并 Sigmoid最后作为权重乘回原特征。这个设计的好处是权重不再是单一通道标量而是带方向的位置敏感权重。做遥感图像里的细长目标检测时我在同一 backbone 上对比过 SE 和 CACA 在小目标召回上大约高出 1 到 2 个点。代价是多了一次 concat 和两个额外卷积延迟增加大约 5% 到 8%。3.4 三种模块的选型对照模块关注维度是否含空间信息额外参数适用场景SE通道无极少约 $2C^2/r$分类任务、算力受限的移动端CBAM通道 空间有但不含坐标方向少空间分支仅 98 参数通用检测/分割插在 stage3 后CA通道 方向坐标有含 H/W 方向略多于 SE细长目标、需要位置敏感权重的任务通道-空间协同注意力机制这个说法本质就是 CBAM 这类模块的设计哲学通道决定“关注什么特征”空间决定“关注哪里”两者串联或并联。如果非要并联得注意两个分支输出的尺度要归一化到同一量级否则相乘会放大某一侧的影响。实操心得给已有 backbone 加这些模块时先用torchsummary或thop算一遍 FLOPs 增量再决定插几层。我见过有人在 ResNet 每个 bottleneck 里都插 CBAMFLOPs 涨了 40%精度只涨 0.3 个点性价比极低。4. 工程落地从 Flash Attention 到一键部署理论清楚了真正上手做项目时瓶颈几乎永远在显存和带宽而不是算力。标准 attention 的 $n \times n$ 分数矩阵要反复在 HBM 和 SRAM 之间搬GPU 的算力单元大部分时间在等数据。这就是 Flash Attention 系列要解决的问题也是为什么现在大模型里的 attention 都要重写成 fused kernel。4.1 Flash Attention 的核心思路与版本演进Flash Attention 的关键点有两个。一是分块tiling把 Q、K、V 切成小块每块能塞进 GPU 的 SRAM在片上算完再写回避免把完整分数矩阵写进 HBM。二是重计算recomputation反向传播时不存中间注意力矩阵而是重新算一遍用算力换显存。数学上还要处理一个麻烦softmax 需要全局最大值才能数值稳定但分块计算时看不到全局。解决方案是 online softmax维护一个运行最大值和运行分母每处理一个新块就更新一次最后统一归一化。这个技巧让分块结果和一次性计算结果完全等价。版本上Flash Attention 1 和 2 主要针对 Ampere 及之后的架构2 代把非矩阵乘的操作也搬到了片上减少了约一半的非矩阵乘开销。Flash Attention 3 针对 Hopper 的异步特性做了流水线重叠。Flash Attention 4 面向新一代架构重点是用更底层的张量核心指令和异步流水把矩阵乘的效率再往上推。对日常使用者来说不用太关心具体代际关键是知道“长度超过 1024 且显存吃紧时开 fused attention 基本是稳赚”。4.2 PyTorch 里的 SDPA先别急着装第三方库从 PyTorch 2.0 开始torch.nn.functional.scaled_dot_product_attention已经内置了后端选择会自动在 Flash、Memory-Efficient 和数学实现之间挑。我的建议是先用它跑不满了再考虑第三方。import torch import torch.nn.functional as F q torch.randn(8, 16, 1024, 64, devicecuda, dtypetorch.float16) k torch.randn(8, 16, 1024, 64, devicecuda, dtypetorch.float16) v torch.randn(8, 16, 1024, 64, devicecuda, dtypetorch.float16) with torch.backends.cuda.sdp_kernel(enable_flashTrue, enable_mathFalse, enable_mem_efficientFalse): out F.scaled_dot_product_attention(q, k, v, is_causalTrue)用sdp_kernel上下文管理器强制走 Flash 后端如果后端不可用会直接报错而不是静默回退到慢速实现。这一点很重要静默回退会让你以为优化生效了实际上还在跑数学实现。维度上我实测过一组数据序列长度 4096、batch 8、头数 16、头维 64FP16 下标准实现峰值显存约 4.2 GBSDPA 走 Flash 后降到约 1.1 GB单步前向时间从约 18 ms 降到约 6 ms。长度 512 时两者差距不到 20%所以短序列没必要折腾。4.3 Sage Attention 与 Triton 的安装要点Sage Attention 走的是另一条路用量化降低 attention 中 QK 计算的位宽把输入量化到 INT8 做矩阵乘再反量化回去做 PV。它的卖点是精度损失可控的前提下速度比标准实现更快尤其适合视频生成这类长序列、高显存的场景。Triton 是它依赖的 kernel 编译框架。安装顺序建议是先装匹配 CUDA 版本的 PyTorch再装 Triton最后装 Sage Attention。顺序错了很容易出现版本冲突。# 以 Windows 环境为例先确认 torch 与 CUDA 版本匹配 python -c import torch; print(torch.__version__, torch.version.cuda) # 安装 TritonWindows 上通常用社区维护的预编译包 pip install triton-windows # 安装 Sage Attention pip install sageattention在 ComfyUI 里启用时启动参数加--use-sage-attention。如果启动后没有报错但速度没变化多半是没有真正挂载上去日志里搜sage关键字确认。我第一次装的时候在虚拟环境里装错了 Python 版本对应包表面上 import 成功实际调用时回退到了默认实现白折腾一晚上。注意量化 attention 对精度敏感的模型比如某些需要精细文本还原的文生图流程可能有可见影响。建议先用固定随机种子跑一组对照图肉眼确认细节没有明显退化再正式用。4.4 DINOv1 attention map把注意力画出来看做可解释性分析时DINOv1 的 attention map 是个很好的观测窗口。DINO 训练时用自蒸馏最后一层自注意力里 CLS token 对其他 patch 的权重往往会呈现出清晰的目标轮廓不需要任何分割标签。导出方式是从注意力模块里取出 softmax 后的权重矩阵取 CLS 那一行去掉 CLS 自身后 reshape 回 $H \times W$再归一化到 0 到 255 存成灰度图。import torch import numpy as np from PIL import Image def dump_attention_map(attn_weights, num_patches_per_side, out_path): # attn_weights: (B, heads, N1, N1)取第一个样本、对多头求平均 w attn_weights[0].mean(dim0) # (N1, N1) cls_row w[0, 1:] # 去掉 CLS 自身 grid cls_row.reshape(num_patches_per_side, num_patches_per_side) grid (grid - grid.min()) / (grid.max() - grid.min() 1e-6) img (grid.cpu().numpy() * 255).astype(np.uint8) Image.fromarray(img).resize((224, 224), Image.BICUBIC).save(out_path)几个观察经验多头平均通常比单个头更干净浅层的注意力图比较散深层的才聚成目标形状patch size 越小轮廓越精细但计算量越大。如果图上是均匀噪点先检查是不是把 softmax 之前的 logits 直接拿来用了或者归一化时除到了接近零的极差。5. 踩坑记录与排查清单不管是在 NLP 还是视觉任务里attention 相关的报错和“效果不对”往往有几类固定模式。这一章按我实际遇到过的问题整理成速查表附带排查路径。5.1 典型报错与定位方法现象可能原因排查动作输出全是 NaN整行被 masksoftmax 全 $-\infty$检查 mask 是否每行至少有一个有效位loss 不下降忘了除以 $\sqrt{d_k}$ 或缩放写错打印 scores 的标准差应在 1 附近显存随长度平方暴涨未启用 fused attention检查 SDPA 后端是否真的走了 Flash训练好但推理崩解码器缺因果掩码训练时偷看未来检查掩码是否为下三角换头数后 shape 报错$d_{model}$ 不能被头数整除加断言或改为非均匀分头注意力图全是对角线层数太浅或学习率过大观察深层注意力调小 lr多头输出和单头几乎一样各头初始化太接近未分化检查是否共享了投影层参数NaN 这个问题我想多说两句。它最容易在混合精度训练里出现。FP16 的动态范围窄$-\infty$ 经过某些 kernel 会变成 NaN。稳妥做法是用float(-inf)之前先把 mask 转成加性掩码或者直接用torch.finfo(dtype).min代替负无穷。我在一个 batch 里混有全 padding 样本时踩过这个坑后来在数据加载阶段直接过滤掉全 padding 样本问题就消失了。5.2 参数选择与调试的几个经验值关于头数不是说越多越好。$d_{model}512$ 时8 头每头 64 维是常见配置16 头每头 32 维在部分任务上略有提升但头维低于 32 后单头表达力下降明显收益开始变负。我做消融时 $h32$、头维 16 的配置比 $h8$ 低了约 0.8 个点。关于 dropout注意力权重上的 dropoutattention dropout和残差上的 dropout 作用不同。前者防止某些位置被过度关注后者防止整体过拟合。小数据集上我把 attention dropout 设到 0.1残差 dropout 设 0.1 到 0.2效果比只调一个更稳。关于学习率预热attention 的 Q/K 投影层对初始学习率比较敏感。前 4000 步线性预热、峰值 lr 在 1e-4 到 3e-4 之间是我在小规模训练里比较通用的设置。跳过预热直接上大 lr经常出现前几百步 loss 剧烈震荡甚至直接发散。5.3 一个容易忽略的细节初始化Q、K 的投影层如果用默认的 Xavier 或 Kaiming 初始化在深层堆叠时注意力分数会逐层放大。有的实现会对 Q、K 用更小的初始化标准差让初始阶段的注意力分布接近均匀。判断是否需要调整的方法很简单训练第一步打印第一层和最后一层的注意力权重熵如果最后一层熵明显低于第一层说明分数尺度已经在累积放大。熵的计算是 $-\sum p \log p$。均匀分布时熵最大等于 $\log n$接近 one-hot 时熵接近 0。用这个指标监控注意力是否过早“硬化”比盯着 loss 曲线更直观。我在一个 12 层的小模型上用过这个方法发现第 9 层之后熵急剧下降把该层 Q/K 初始化标准差调小一半后最终指标提升了约 1.2 个点。5.4 长度外推时的注意事项训练长度 512、推理长度 2048 这种场景除了位置编码要选可外推的方案注意力本身的分布也会变化。序列变长后同样的温度下 softmax 分布会更平因为 Key 数量变多分母变大。有的实现会在长序列推理时对 logits 乘一个略大于 1 的缩放系数来补偿但这个系数最好在验证集上搜不要拍脑袋。我自己在这个环节的体会是与其在推理时打补丁不如在训练阶段就用可变长度采样让模型见过不同长度下的分布。把训练时的长度从固定 512 改成 256 到 1024 之间随机采样外推到 2048 时的性能衰减比固定长度训练小了一半左右代价是训练时每步的计算量略有波动但总体可接受。最后再分享一个小技巧调试 attention 时把序列长度设成 4、头数设成 2、维度设成 8手动打印每一步的张量形状和数值比在大模型上盲猜快得多。所有 shape 错误在这个规模下都会暴露得清清楚楚修好之后直接放大配置基本不会再有结构性问题。这套“缩小到玩具规模再放大”的流程我在多头自注意力、交叉注意力、视觉注意力模块上都用过屡试不爽。