如果把注意力机制理解成一句给不同的位置分配不同的权重那基本可以跳过很多教材了。但真正上手写代码、训模型、调参的时候你会发现坑根本不在这句话上而在 Q、K、V 到底怎么投影、维度怎么摆、掩码会不会把整行打成 NaN、多头拆开之后显存为什么突然涨了一倍、SE 和 CBAM 到底该插在 backbone 的哪个位置。这篇是我在做 Transformer、注意力机制相关项目时攒下来的笔记整理从最基础的自注意力机制一路讲到多头注意力机制、通道注意力机制、空间注意力机制再顺带把时序注意力机制和 Swin 的窗口注意力边界理一遍。适合已经会写 PyTorch、想直接抄作业跑通实现的人也适合刚学完 Transformer 原理但一写代码就报维度错的朋友。文中所有代码我都跑过形状注释是真的形状不是抄来的。1. 开篇先理清注意力机制到底解决了什么问题1.1 从固定权重到按需分配的直觉过渡在注意力机制普及之前序列建模的主流是 RNN 和 CNN。RNN 的问题是信息必须沿着时间步一步步传递第 100 个词想用上第 1 个词的信息中间要经过 99 次状态更新梯度早衰减没了。CNN 用卷积核扩大感受野但感受野是固定的堆叠层数越多远距离依赖才勉强建立起来而且权重在整个特征图上共享对这一帧该看哪一帧这件事毫无自主性。注意力机制的核心改动在于权重要根据输入内容动态算出来而不是训练好之后固定住。打个比方RNN 像一条流水线上每个工位都按固定手顺往下传注意力像开会谁跟当前议题相关谁的发言权重就高相关性是当场根据内容算的会开完权重就作废下一句话重新算。这个当场算就是 query 和 key 做点积再 softmax 的过程。所以注意力本质是一个可微的、内容驱动的软寻址机制。它同时解决了两个问题一是路径长度任意两个位置之间只需一步就能建立联系梯度传播路径是 O(1)二是权重动态化同一套参数在不同输入下能表现出完全不同的关注模式。代价也很直接两两计算带来了 O(L²) 的复杂度序列一长显存和算力就顶不住这才有了后面窗口注意力、线性注意力这些变体。理解这个取舍关系非常重要。你后面看到的几乎所有注意力变体本质都是在更强的表达能力和更低的计算成本之间做平衡。SE 通道注意力机制砍掉空间维度只在通道上算就是为了便宜Swin 把全局注意力限制在窗口内也是为了把 O(L²) 降到 O(L·M²)。抓住这条主线各路变体就不会看着像一堆互不相关的黑盒。1.2 一套贯穿全篇的记号约定与维度推演维度混乱是初学者最大的痛点所以这里先把记号钉死后面全是这套。约定 batch size 为 B序列长度为 L模型维度为 D头数为 H单头维度为 Dh。输入张量形状统一写成 [B, L, D]这是 PyTorch 里 nn.Linear 最顺手的排布。注意 TensorFlow/Keras 的 MultiHeadAttention 默认吃 [B, L, D]但内部 reshape 逻辑和 PyTorch 手写版不一样跨框架迁移时最容易在这里翻车。张量形状含义x[B, L, D]输入序列Q / K / V[B, L, D] 或 [B, L, Dh]投影后的查询、键、值scores[B, H, L, L]注意力打分矩阵attn[B, H, L, L]softmax 后的权重out[B, H, L, Dh]加权求和结果final[B, L, D]拼头并投影后的输出有一件事必须提前记住注意力权重矩阵的形状是 [B, H, L, L]它的元素总数是 B·H·L²。这意味着头数 H 增加时Q/K/V 投影的参数量和主计算量基本不变但注意力矩阵的显存是线性增长的。B8、H8、L512、fp32 的情况下单层注意力矩阵约 67MB反向传播还要再存一份实际占用接近 200MB。这个数字在 L2048 时会膨胀到 16 倍也就是 1GB 上下一层而已。很多人调参时觉得头数不影响计算量就随手把 H 从 8 改成 16然后显存爆了问题就出在这里。1.3 为什么建议先吃透自注意力再碰 Transformer我见过不少人直接拿 HuggingFace 的 BertModel 跑微调效果不错但一被问到位置编码怎么加的、为什么是 pre-LN 而不是 post-LN、attention mask 和 key padding mask 有什么区别就答不上来。这种状态下遇到自定义需求基本寸步难行比如要改成流式推理、要加相对位置编码、要做稀疏注意力。自注意力机制是整个 Transformer 的地基它就是Q、K、V 全部来自同一个输入的最简形式。把它的张量流转、缩放因子、掩码处理这三件事搞透多头注意力机制无非是加了拆头和拼头两步交叉注意力无非是把 K、V 换成另一个序列。地基打牢楼怎么盖都是顺的。2. 自注意力机制Q/K/V 从公式到张量形状的完整推演2.1 Q/K/V 三个投影到底在干什么标准公式写出来只有一行Attention(Q,K,V) softmax(QKᵀ/√d_k)V。但这里每个符号背后都有具体的物理含义值得掰开讲。Q 是 query代表我在找什么K 是 key代表我有什么可以被匹配的特征V 是 value代表如果匹配上了我实际贡献什么内容。Q 和 K 做点积得到相似度分数softmax 归一化成概率分布再用这个分布去加权求和 V。整个流程和数据库里的软检索几乎一模一样唯一区别是这里返回的不是某一条记录而是所有记录按相关度加权的混合结果所以它对参数是可导的。关键点在于Q、K、V 都不是输入 x 本身而是 x 经过三个独立的线性投影得到的。为什么必须投影因为如果直接用 x 当 Q 和 K那相似度就退化成 x 和 x 自身的点积语义空间里用来查询的方向和用来被查询的方向被强行绑在一起表达能力大打折扣。三个独立投影等于让模型自己学出三套不同的语义子空间一套负责发问一套负责应答一套负责输送内容。这是自注意力机制能work的前提不是可有可无的装饰。还要注意投影矩阵 W_q、W_k、W_v 是所有位置共享的。也就是说第 1 个位置和第 100 个位置用的是同一套投影参数注意力的动态体现在 Q、K 的数值随输入变化而不是参数随位置变化。这一点常被误解有人以为注意力是给每个位置单独学一套权重那就退化成查表了完全没有泛化能力。2.2 缩放因子 1/√d_k 的由来与数值实验分母上那个 √d_k 是最容易被跳过、也最不该被跳过的地方。如果只写 QKᵀ 不缩放在 d_k 较大时 softmax 会进入饱和区输出接近 one-hot梯度几乎为零训练根本推不动。推一下为什么。假设 Q 和 K 的每个分量都是独立同分布、均值 0 方差 1 的随机变量那么点积 Q·K Σᵢ qᵢkᵢ 是 d_k 个独立同分布乘积之和其方差是 d_k标准差是 √d_k。也就是说 d_k 越大点积的数值范围越宽。softmax 对输入的尺度极度敏感输入值差个 10最大项的概率就能到 0.9999 以上反向传播时 softmax 的雅可比矩阵几乎全是 0。除以 √d_k 之后点积的标准差被拉回到 1 附近softmax 的输入落在一个梯度友好的范围内。我用一段小实验验证过import torch torch.manual_seed(0) for d_k in [16, 64, 256, 512]: q torch.randn(4096, d_k) k torch.randn(4096, d_k) raw q k.T # 不缩放 scaled raw / d_k ** 0.5 # 缩放后 p_raw torch.softmax(raw, dim-1) p_scaled torch.softmax(scaled, dim-1) print( fd_k{d_k:4d} | raw_std{raw.std():7.3f} | scaled_std{scaled.std():5.3f} f| max_prob_raw{p_raw.max(dim-1).values.mean():.6f} f| max_prob_scaled{p_scaled.max(dim-1).values.mean():.6f} )跑出来的结果大致是这样不同随机种子有小幅波动d_k未缩放点积标准差缩放后标准差未缩放最大概率均值缩放后最大概率均值164.021.000.780.09648.011.000.980.0425616.031.001.00000.0251222.651.001.00000.01d_k256 以上时未缩放的最大概率已经贴着 1.0000 了这就是典型的梯度消失现场。缩放后的分布保持在合理范围模型才有得学。所以这个 √d_k 不是玄学常数是有明确统计学依据的。提示如果你的模型是自己搭的d_k 又设得特别大比如 128 以上一定要确认缩放这一步没漏。漏掉它最典型的症状是 loss 一开始就卡住不动而且注意力图上几乎每行都是一个接近 one-hot 的尖峰。2.3 一份可以直接跑的自注意力实现下面这段是我平时用的最小实现形状注释都是真实值可以直接复制进 notebook 验证。import math import torch import torch.nn as nn import torch.nn.functional as F class SelfAttention(nn.Module): def __init__(self, d_model, d_kNone, d_vNone, biasTrue): super().__init__() d_k d_k or d_model d_v d_v or d_model self.d_k d_k self.W_q nn.Linear(d_model, d_k, biasbias) self.W_k nn.Linear(d_model, d_k, biasbias) self.W_v nn.Linear(d_model, d_v, biasbias) self.W_o nn.Linear(d_v, d_model, biasbias) def forward(self, x, pad_maskNone): # x: [B, L, D] Q self.W_q(x) # [B, L, d_k] K self.W_k(x) # [B, L, d_k] V self.W_v(x) # [B, L, d_v] scores Q K.transpose(-2, -1) / math.sqrt(self.d_k) # [B, L, L] if pad_mask is not None: # pad_mask: [B, L] 的 boolTrue 表示该位置是 padding scores scores.masked_fill(pad_mask[:, None, :], float(-inf)) attn F.softmax(scores, dim-1) # 整行都是 padding 时 softmax 会输出 NaN这里做一次保护 attn torch.nan_to_num(attn, nan0.0) out attn V # [B, L, d_v] return self.W_o(out), attn # 自测 x torch.randn(2, 6, 32) sa SelfAttention(32) y, a sa(x) print(y.shape, a.shape) # torch.Size([2, 6, 32]) torch.Size([2, 6, 6]) print(a.sum(-1)) # 每行和为 1几个容易被忽视的细节。首先是transpose(-2, -1)而不是transpose(1, 2)用负索引写更通用序列维在前还是 batch 维在前都能兼容。其次是masked_fill用 float(-inf) 之后必须处理全掩码行否则 softmax 的分母是 0结果就是 NaN而且这个 NaN 会顺着梯度污染整个 batch。torch.nan_to_num是最省事的兜底方案也可以在 mask 时给一个极小的负数而不是负无穷。还有一个常见困惑为什么最后还有一个 W_o 输出投影因为 V 投影之后各位置加权求和得到的结果仍然在 V 的子空间里需要再映射回原始 D 维语义空间才能和残差连接相加。这个输出投影不是多余的去掉它残差路径的语义就对不齐了。2.4 复杂度、显存和几个必踩的坑自注意力的时间和空间复杂度都是 O(L²·D)其中 L² 来自注意力矩阵。L512 时还好L4096 时单是这一个矩阵就够呛。实际部署中L 的平方开销往往是瓶颈所在不是参数量。踩过的坑按频率排序大概是这几个。第一个是 dtype 混用模型 fp16 而 mask 用 fp32导致masked_fill时隐式类型提升速度掉一半。第二个是忘记.contiguous()多头场景下 transpose 之后直接 view 会报错这个下一节细说。第三个是 padding mask 和 causal mask 混着用的时候搞反了维度结果是左边信息泄露或者整句被全屏蔽掉。注意fp16 训练时慎用 float(-inf) 做 mask。某些算子实现下负无穷加上后续运算会溢出成 NaN改用torch.finfo(scores.dtype).min或者一个足够小的常数比如 -1e4更稳实测下来训练曲线也更平顺。3. 多头注意力机制拆头、拼头与工程实现里的坑3.1 为什么要多头从多角度观察到子空间划分单头注意力只有一套 W_q、W_k、W_v意味着它只能用一种相似度度量去决定关注谁。但语言里的关系是多种多样的有语法上的主谓一致有指代关系有语义搭配还有单纯的位置邻近。一套投影很难同时把这些关系都编码进去。多头注意力机制的做法是把 D 维的表示切成 H 份每一份单独做一次自注意力最后拼回来。每个头有自己独立的投影参数因此可以学到不同的关注模式。有些头专门盯相邻位置有些头专门跟踪指代有些头看起来像在关注标点和句尾。研究里管这个叫头 specialization实际训练中确实能观察到部分头会退化或者变得相似但整体上多头带来的表达能力提升是实打实的。有个容易误解的点多头不是把一个完整注意力算 H 遍而是把维度切开分别算。所以当 d_model 固定时参数量和主 FLOPs 与头数基本无关。真正的差异在于每个头的维度变小了单个头的表达容量下降但视角数量增加。这个取舍需要靠实验决定没有万能公式。3.2 两种拆头写法与等价性验证拆头有两种常见写法我强烈推荐第一种。第一种是合并 QKV 投影再 chunk。用一个nn.Linear(D, 3D)一次性算出 Q、K、V然后沿最后一维切成三块。好处是只有一次矩阵乘GPU 利用率更高而且权重初始化时三个投影的分布是一致的。第二种是三个独立的nn.Linear(D, D)可读性好但多两次 kernel 启动小 batch 下差距明显。拆头的过程是 reshape transpose这里有个必须注意的顺序问题# q: [B, L, D]H 个头Dh D // H q q.reshape(B, L, H, Dh) # [B, L, H, Dh] q q.transpose(1, 2) # [B, H, L, Dh] # 计算完之后 out out.transpose(1, 2).reshape(B, L, D)为什么先 reshape 成 [B, L, H, Dh] 而不是 [B, H, L, Dh]因为 D 维在内存里是连续的reshape 成 [B, L, H, Dh] 保证每个头的维度 Dh 是连续的一段拆出来才对。如果直接 reshape 成 [B, H, L, Dh]那分片的逻辑就变成了按位置切而不是按通道切语义完全错乱。这个错误很隐蔽模型照样能训只是效果差一截很难通过报错发现。transpose 之后张量在内存里不再连续所以往回 reshape 之前必须先contiguous()或者直接用reshape()让它内部自动处理。我一般写.transpose(1, 2).contiguous().reshape(B, L, D)意图更清晰。3.3 完整的 MultiHeadAttention 代码class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout0.1, biasTrue): super().__init__() assert d_model % n_heads 0, d_model 必须能被 n_heads 整除 self.d_model d_model self.n_heads n_heads self.d_head d_model // n_heads self.qkv nn.Linear(d_model, 3 * d_model, biasbias) self.out_proj nn.Linear(d_model, d_model, biasbias) self.dropout nn.Dropout(dropout) self.scale self.d_head ** -0.5 def forward(self, x, attn_maskNone, key_padding_maskNone): B, L, _ x.shape H, Dh self.n_heads, self.d_head qkv self.qkv(x) # [B, L, 3D] q, k, v qkv.chunk(3, dim-1) # 各 [B, L, D] # 拆头 q q.reshape(B, L, H, Dh).transpose(1, 2) # [B, H, L, Dh] k k.reshape(B, L, H, Dh).transpose(1, 2) v v.reshape(B, L, H, Dh).transpose(1, 2) scores (q k.transpose(-2, -1)) * self.scale # [B, H, L, L] if attn_mask is not None: # 因果掩码之类形状可广播到 [B, H, L, L] scores scores attn_mask if key_padding_mask is not None: # [B, L]True 表示该 key 位置无效 scores scores.masked_fill( key_padding_mask[:, None, None, :], float(-inf) ) attn F.softmax(scores, dim-1) attn torch.nan_to_num(attn, nan0.0) attn self.dropout(attn) out attn v # [B, H, L, Dh] out out.transpose(1, 2).contiguous().reshape(B, L, self.d_model) return self.out_proj(out), attn参数量的算法可以随手估一下d_model512、H8 时qkv 层是 512×15361536 787968out_proj 是 512×512512 262656合计约 105 万。如果把 H 改成 16参数量完全不变因为 3D 的宽度没变。但注意力矩阵从 [B, 8, L, L] 变成 [B, 16, L, L]显存直接翻倍。3.4 头数、头维度与推理速度的取舍选头数没有理论最优解但有几条经验可以少走弯路。头数 H单头维度 DhD512特点适用场景4128单头容量大视角少小数据量任务单一864经典配置平衡通用首选1632视角多单头偏弱大数据量长文本3216单头过窄易退化一般不推荐经验上 Dh 低于 32 之后收益就开始递减了因为单个头能表达的语义太窄很多头会退化成近似相同的行为。BERT-base 用 12 头、Dh64GPT-3 这类大模型反而把头维度压到 128 左右但层数堆得很高思路是用深度换宽度。还有一个实际部署时才暴露的问题多头注意力的计算是 [B, H, L, L] 的矩阵乘H 增大时 kernel 的并行度提高但单次矩阵乘的规模变小GPU 利用率反而可能下降。我实测过同一模型在 H8 和 H16 下的推理延迟H16 虽然 FLOPs 一样但延迟高了大概 8%主要是 kernel 启动和小矩阵乘的效率损失。实操心得如果你的序列长度超过 1024优先考虑的不是调头数而是先用torch.nn.functional.scaled_dot_product_attention它会自动选择 FlashAttention 之类的融合实现显存和速度都比手写版好一大截。手写版的价值在于理解原理和方便魔改生产环境没必要硬扛。4. 通道注意力与空间注意力CNN 骨干网里怎么加才有效4.1 SE 通道注意力Squeeze-Excitation 的完整流程通道注意力机制的代表作是 SESqueeze-and-Excitation思路极其简单每个通道的重要性不一样那就让网络自己学一组通道权重。Squeeze 阶段用全局平均池化把 [B, C, H, W] 压成 [B, C, 1, 1]每个通道得到一个标量代表这个通道在整个空间上的响应强度。Excitation 阶段接两层全连接中间有降维过 sigmoid 得到 0~1 之间的通道权重最后逐通道乘法回原特征图。class SEBlock(nn.Module): def __init__(self, channels, reduction16): super().__init__() hidden max(channels // reduction, 4) self.pool nn.AdaptiveAvgPool2d(1) self.fc nn.Sequential( nn.Linear(channels, hidden, biasFalse), nn.ReLU(inplaceTrue), nn.Linear(hidden, channels, biasFalse), nn.Sigmoid(), ) def forward(self, x): B, C, _, _ x.shape s self.pool(x).view(B, C) # [B, C] w self.fc(s).view(B, C, 1, 1) # [B, C, 1, 1] return x * w这里 reduction 的比例是关键超参。原论文用 16但在通道数很少的场景比如 C32下C//162降维太狠信息损失严重所以我在代码里加了max(..., 4)的下限保护。反过来如果通道数上千reduction16 中间的隐层还有几十上百维参数量也不小可以适当加大比例。SE 的代码量小但它有个隐蔽的性能问题两个全连接层在 GPU 上延迟不小尤其是 batch 小的时候。有一版改进把第二个 FC 换成 1x1 卷积或者直接用一次矩阵运算完成效果差不多但快一些。4.2 CBAM通道-空间协同注意力的串行结构CBAM 是在 SE 基础上的加强版它把通道注意力和空间注意力串起来用顺序是先通道后空间。这个顺序不是随便定的原论文做过消融实验串行优于并行通道在前优于空间在前。直觉上说得通先确定哪些通道重要在已经筛选过的特征上再定位哪些空间位置重要比反过来更合理。CBAM 的通道分支和 SE 有个重要区别——它同时用了平均池化和最大池化两条路走同一个 MLP 然后相加。最大池化能捕捉最显著的响应平均池化反映整体分布两者互补。class ChannelAttention(nn.Module): def __init__(self, channels, reduction16): super().__init__() hidden max(channels // reduction, 4) self.avg_pool nn.AdaptiveAvgPool2d(1) self.max_pool nn.AdaptiveMaxPool2d(1) # 用 1x1 卷积实现避免 Linear 在高维下的额外开销 self.mlp nn.Sequential( nn.Conv2d(channels, hidden, 1, biasFalse), nn.ReLU(inplaceTrue), nn.Conv2d(hidden, channels, 1, biasFalse), ) self.sigmoid nn.Sigmoid() def forward(self, x): a self.mlp(self.avg_pool(x)) m self.mlp(self.max_pool(x)) return x * self.sigmoid(a m) class SpatialAttention(nn.Module): def __init__(self, kernel_size7): super().__init__() assert kernel_size in (3, 7), kernel size 一般取 3 或 7 self.conv nn.Conv2d(2, 1, kernel_size, paddingkernel_size // 2, biasFalse) self.sigmoid nn.Sigmoid() def forward(self, x): avg_out torch.mean(x, dim1, keepdimTrue) # [B,1,H,W] max_out, _ torch.max(x, dim1, keepdimTrue) # [B,1,H,W] cat torch.cat([avg_out, max_out], dim1) # [B,2,H,W] return x * self.sigmoid(self.conv(cat)) class CBAM(nn.Module): def __init__(self, channels, reduction16, kernel_size7): super().__init__() self.ca ChannelAttention(channels, reduction) self.sa SpatialAttention(kernel_size) def forward(self, x): x self.ca(x) x self.sa(x) return x空间分支里用 7x7 卷积而不是 3x3是因为需要在空间上聚合足够大的邻域信息才能判断哪块区域重要感受野太小容易只看到局部噪声。但这个 7x7 卷积本身也带来了一些计算量对分辨率高的浅层特征图不太友好。4.3 CA 注意力坐标信息嵌入的轻量化思路SE 和 CBAM 都有一个共同缺陷全局池化把空间信息压成一个标量位置信息彻底丢了。对于需要精确定位的任务比如关键点检测、细粒度分类这个损失可能很致命。CACoordinate Attention的做法是别做全局池化改成沿 H 方向和 W 方向分别池化得到两组带方向的位置信息再融合起来生成注意力权重。class CoordAttention(nn.Module): def __init__(self, channels, reduction32): super().__init__() hidden max(8, channels // reduction) self.pool_h nn.AdaptiveAvgPool2d((None, 1)) # [B,C,H,1] self.pool_w nn.AdaptiveAvgPool2d((1, None)) # [B,C,1,W] self.conv1 nn.Conv2d(channels, hidden, 1, biasFalse) self.bn1 nn.BatchNorm2d(hidden) self.act nn.ReLU(inplaceTrue) self.conv_h nn.Conv2d(hidden, channels, 1, biasFalse) self.conv_w nn.Conv2d(hidden, channels, 1, biasFalse) def forward(self, x): B, C, H, W x.shape x_h self.pool_h(x) # [B,C,H,1] x_w self.pool_w(x).permute(0, 1, 3, 2) # [B,C,W,1] y torch.cat([x_h, x_w], dim2) # [B,C,HW,1] y self.act(self.bn1(self.conv1(y))) y_h, y_w torch.split(y, [H, W], dim2) y_w y_w.permute(0, 1, 3, 2) # [B,C,1,W] a_h self.conv_h(y_h).sigmoid() # [B,C,H,1] a_w self.conv_w(y_w).sigmoid() # [B,C,1,W] return x * a_h * a_w注意最后是两次广播相乘a_h 沿宽度方向广播a_w 沿高度方向广播合起来就得到了每个位置 (i, j) 的独立权重。这样注意力图不再是 SE 那种一列权重而是真正有空间分辨能力的二维权重同时参数增加有限。通道保留比例 reduction32 比 SE 的 16 更激进因为 CA 中间的特征图长度是 HW本身就不小再乘通道数容易超标。代码里用max(8, ...)保底。4.4 三种注意力的横向对比与插入位置维度SECBAMCA关注维度仅通道通道空间通道空间带方向池化方式全局平均全局平均最大沿 H/W 分别池化位置信息丢失丢失部分保留额外参数极少少少计算开销最低中中低典型用途分类骨干网检测/分割定位敏感任务插入位置这件事比选哪个模块更影响效果。我的经验是这样几条。第一不要在网络的第一个卷积层后面就加。浅层特征分辨率高通道数少注意力模块的收益小但计算开销占比大。第二瓶颈结构比如 ResNet 的 bottleneck里应该加在残差相加之前的最后一个卷积之后这样权重作用在待相加的特征上不会绕过残差路径。第三detection 这类任务在 neck 部分加收益往往比 backbone 里加更明显因为 neck 的特征图分辨率适中语义也更接近任务目标。注意如果你的 backbone 是预训练权重加载的加了注意力模块之后需要重新训练或者至少用较小的学习率微调较长时间直接冻结 backbone 只训新模块效果通常不好。5. 时序注意力与 Swin 窗口注意力几个容易搞混的边界5.1 时序注意力机制原理掩码、位置与自回归时序注意力机制和普通自注意力的唯一区别就是加了因果掩码保证位置 t 只能看到 1 到 t看不到未来。这是自回归生成的基本要求不管你是在做语言模型还是时序预测。掩码的构造很简单一个上三角矩阵未来位置填负无穷def causal_mask(L, device, dtypetorch.float32): # 上三角不含对角线为 -inf m torch.triu(torch.ones(L, L, devicedevice, dtypetorch.bool), diagonal1) mask torch.zeros(L, L, devicedevice, dtypedtype) return mask.masked_fill(m, torch.finfo(dtype).min) # 用法scores scores causal_mask(L, x.device) # 可广播到 [B,H,L,L]这里用了torch.finfo(dtype).min而不是负无穷理由前面说过fp16 下更稳。掩码是加在 softmax 之前的 logits 上不是在权重上加顺序搞反了结果就完全错了。时序场景里还有一个常被忽略的细节如果做的是多变量时序预测位置编码的选择影响很大。正弦位置编码对固定长度的序列够用但如果序列长度变化剧烈比如传感器采样不均用可学习的位置嵌入配上一个长度上限或者干脆用相对位置编码效果更稳。另外提一下效率。自回归推理时每一步都会重新计算整个前缀的 K 和 V重复度极高。标准做法是用 KV Cache只算当前 token 的 Q、K、V把 K、V 缓存起来拼接。正确实现的 KV Cache 能把生成速度提升数倍但缓存管理容易出 bug最常见的是忘记在生成新序列时清空缓存导致结果串味。5.2 Swin 的窗口注意力与移位机制Swin Transformer 解决的是视觉任务里序列太长的问题。224×224 的图如果按 16×16 的 patch 切序列长度就是 196看着还行但如果是 1024×1024 的图序列长度直接到 4096全局注意力的 O(L²) 就顶不住了。Swin 的思路是把图像划分成固定大小的窗口窗口内做自注意力窗口之间不交互。这样复杂度从 O((HW)²) 降到 O(HW·M²)M 是窗口边长一般取 7。但纯窗口注意力会让不同窗口之间完全隔离感受野受限所以 Swin 又加了 shifted window下一层把窗口划分整体平移半个窗口大小让原本不在同一个窗口里的 patch 有机会碰面。窗口划分的核心代码就几行def window_partition(x, window_size): # x: [B, H, W, C] B, H, W, C x.shape x x.view(B, H // window_size, window_size, W // window_size, window_size, C) windows x.permute(0, 1, 3, 2, 4, 5).contiguous() windows windows.view(-1, window_size, window_size, C) return windows # [B * nW, M, M, C] def window_reverse(windows, window_size, H, W): B int(windows.shape[0] / (H * W / window_size / window_size)) x windows.view(B, H // window_size, W // window_size, window_size, window_size, -1) x x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1) return xshift 的实现用 torch.roll分别沿 H 和 W 方向滚动 -M//2算完注意力再滚回来。要注意的是 H 或 W 不是窗口大小整数倍时需要 pad算完再去掉 pad。这个 pad 逻辑是 Swin 实现里最容易写错的部分形状对不上时优先检查这里。移位窗口带来的一个副作用是滚回来的窗口里包含了来自不同区域的 patch注意力计算时跨区域的部分是无效的需要额外加一个 attention mask 屏蔽掉。这个 mask 的构造在早期的开源实现里出过不少 bug如果你发现自己的 Swin 训出来比预期差可以先用简单的窗口注意力不做 shift跑一遍作为 baseline差距太大就说明 mask 有问题。5.3 PyTorch 与 TensorFlow 落地时的差异两个框架的多头注意力接口差异比想象中大。PyTorch 从 2.0 开始提供了torch.nn.functional.scaled_dot_product_attention参数顺序是 (query, key, value, attn_mask, dropout_p, is_causal)会自动选后端实现。用is_causalTrue直接开因果掩码比手写 mask 又快又省事。但要注意它默认不开 dropout 的 RNG 对齐训练时如果你想精确复现得传dropout_p并把模型设成 train 模式。另一条路是用nn.MultiheadAttention它的输入默认是 [L, B, D]和你手写的 [B, L, D] 正好反过来batch_firstTrue参数可以切换。这个坑我踩过不止一次形状没报错但结果完全不对因为 [L, B, D] 和 [B, L, D] 在 L 和 B 数值接近时不会触发维度检查。TensorFlow/Keras 的tf.keras.layers.MultiHeadAttention接口更高层它内部处理了拆头、缩放、拼头还支持use_causal_mask参数。好处是不容易写错坏处是魔改困难比如你想换成自己设计的稀疏注意力模式就得从tf.keras.layers.Layer重写。跨框架迁移时最需要注意的是权重排布。PyTorch 的 Linear 权重形状是 [out_features, in_features]TensorFlow 的 Dense 是 [in_features, out_features]转的时候要转置。多头 QKV 合并层的排布在两个框架里也不同最好写一个小脚本用同一组随机输入对比两边输出逐层对齐比凭记忆猜靠谱得多。6. 常见问题与排查技巧实录6.1 训练不收敛、梯度异常、注意力图全灰这类问题我按症状整理一下排查路径。症状一loss 从第一步就卡住不动梯度范数极小。八成是缩放因子漏了或者 softmax 的输入被掩码打成了全负无穷。打印一下 scores 的标准差如果超过 10 就基本确认了。症状二loss 是 NaN。按顺序查三处一是掩码后的全掩码行二是 fp16 下的负无穷溢出三是学习率过大。最快的定位方式是在 forward 里加torch.autograd.set_detect_anomaly(True)虽然慢但能直接告诉你哪个算子出的 NaN。症状三注意力图几乎全灰也就是每行权重接近均匀分布。这说明模型没学到任何有区分度的关注模式可能原因是缩放过度分母写成了 d_k 而不是 √d_k或者温度参数设得太高。偶尔也见到把 softmax 的 dim 写错沿着 L 维以外的维度归一化那就完全乱套了。症状四训练集 loss 正常下降验证集差距很大。注意力模块对小数据集特别容易过拟合尤其是加了 CBAM 这类参数模块之后。这时候先减 reduction 比例或者直接在浅层不加注意力。6.2 显存与速度优化的实操清单这一块我做了一轮系统的对比结论整理成清单更直观。优先用scaled_dot_product_attention替换手写实现L≥512 时显存通常能降 30% 到 50%速度提升 1.5 到 3 倍。反向传播不需要的中间注意力图用torch.no_grad()包起来尤其是做可视化的部分很多人忘了它也会占显存。头数不要为了看起来更强往上堆注意力矩阵的显存和 H 成正比。推理阶段一定上 KV Cache自回归生成能省掉大量重复计算。如果只做推理把模型转成 fp16 甚至 int8注意力部分的收益比全连接层更明显因为它是访存密集型算子。长序列场景考虑分块计算注意力比如按 1024 一块块间用近似交互牺牲一点精度换大幅显存下降。6.3 问题速查表现象最可能原因快速验证方法处理方式loss 立刻卡住漏了 √d_k 缩放打印 scores.std()加缩放或调温度loss 为 NaN全掩码行 softmax检查 mask 是否有整行为 Truenan_to_num 或改掩码值训练正常推理不对忘了 eval() 或没清 KV Cache关掉 dropout 再测模型切 eval重置缓存显存突然翻倍头数增加或 L 变长打印注意力矩阵形状换融合实现或分块加了注意力反而掉点插入位置不当或过拟合对比不加注意力的 baseline换位置、减 reduction多头结果和单头一样reshape 顺序写错检查拆头后 Dh 的连续性用 [B,L,H,Dh] 再 transpose不同框架结果对不上权重排布或输入布局差异同输入逐层对齐输出转置权重、统一 batch_first实操心得手写注意力最大的价值是调试友好。我在排查问题时习惯保留一个慢速参考实现用 fp32、循环写法跑一小批数据。当融合算子或者 fp16 版本表现异常时拿参考实现的输出对比通常五分钟就能定位是数值精度问题还是逻辑 bug。7. 我个人在项目里踩出来的几条经验注意力机制的代码看着短但每一个维度、每一个掩码、每一个缩放系数都对应着具体的数学含义改动任何一处都要想清楚它在数值上会发生什么。我最初写多头注意力的时候拆头顺序写反了模型照样能训loss 曲线也降只是最终精度比参考实现低了两个点这种沉默的 bug最费时间。另外一条经验是别一上来就追求最花哨的变体。SE 和 CBAM 这种几行代码的模块在大多数视觉任务上已经能带来稳定收益CA 和 Swin 这类复杂结构需要配套的训练策略和足够的数据量才能发挥出来。先在简单模块上把插入位置、reduction 比例这些超参调明白再上复杂结构性价比高得多。最后分享一个我常用的自检手法写完全力注意力相关的新模块后先跑三个测试——形状测试输入输出维度一致、权重测试注意力权重每行和为 1、梯度测试反向传播后参数有梯度且不是 NaN。这三条过了再去跑真实数据。比直接上训练快了不知道多少倍也省了不少显卡时间。