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拆开看三步,我用一个 $n=3$、$d_k=2$ 的小例子手算一遍,理解会牢得多。设:
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_{i=1}^{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_k=512$ 时前 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}$。$h=8$、$d_{model}=512$ 时每个头是 64 维。有人会问,切成 8 份每个头只有 64 维,表达能力不是变弱了吗?单头确实弱,但 8 个头关注不同子空间,总容量并没减少,而且参数量和多头共享投影的情况是一致的。
class MultiHeadAttention(torch.nn.Module): def __init__(self, d_model=512, num_heads=8, dropout=0.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, mask=None): 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, dtype=torch.bool, device=device))因果掩码的常见错误是用float('-inf')填充时把整行都填满,导致 softmax 出现全 $-\infty$,输出 NaN。稳妥做法是保证对角线至少为 0,或者在 softmax 之前检查“每行是否至少有一个非掩码位置”。我在调试一个流式解码任务时,就因为在序列全 padding 的 batch 上触发过这个 NaN,排查了半天。
3. 视觉注意力模块三兄弟:SE、CBAM、CA
视觉里的注意力模块和 NLP 的 attention 名字一样,思路却不同。图像任务里大家想要的往往不是“序列位置之间互相看”,而是“让网络自己判断哪些通道重要、哪些空间区域重要”。于是就有了 SE、CBAM、CA 这一系列轻量模块,它们通常插在 backbone 的残差分支上,参数量增加极小,却能稳定涨点。
3.1 SE 通道注意力:把全局信息压成一个权重向量
SE(Squeeze-and-Excitation)是通道注意力的起点。流程是三步:先对每个通道做全局平均池化(Squeeze),把 $H \times W$ 压成一个标量;再经过两层全连接加激活,中间有个降维比例 $r$(常用 16);最后用 Sigmoid 得到 $C$ 维权重,乘回原特征。
class SEBlock(torch.nn.Module): def __init__(self, channels, reduction=16): super().__init__() self.pool = torch.nn.AdaptiveAvgPool2d(1) self.fc = torch.nn.Sequential( torch.nn.Linear(channels, channels // reduction), torch.nn.ReLU(inplace=True), 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$ 参数,$C=512$ 时超过 26 万,而加个 $r=16$ 后降到约 3.3 万。实测下来这个降维对手持设备的推理延迟也很友好。要注意的是AdaptiveAvgPool2d(1)会丢掉所有空间位置信息,SE 对“目标在哪里”是无感的,它只知道“哪些通道整体活跃”。
3.2 CBAM:通道与空间的串联组合
CBAM(Convolutional Block Attention Module)在 SE 的基础上加了一个空间注意力分支,顺序是先通道后空间。通道分支和 SE 略有不同:它同时用平均池化和最大池化,各自过一个小 MLP 后相加再 Sigmoid。空间分支则是在通道维度上做平均池化和最大池化,得到两个 $H \times W$ 图,拼成 2 通道后用一个大核卷积(通常是 7×7)压成 1 通道,再 Sigmoid。
class SpatialAttention(torch.nn.Module): def __init__(self, kernel_size=7): super().__init__() self.conv = torch.nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False) self.sigmoid = torch.nn.Sigmoid() def forward(self, x): avg_out = torch.mean(x, dim=1, keepdim=True) max_out, _ = torch.max(x, dim=1, keepdim=True) cat = torch.cat([avg_out, max_out], dim=1) return x * self.sigmoid(self.conv(cat))用大核 7×7 而不是 3×3,是为了让空间注意力有更大的感受野,能覆盖到整块目标区域而不是局部纹理。代价是插在浅层高分辨率特征上时计算量不小。我给的经验是:CBAM 放在 backbone 的 stage3 之后收益最明显,放 stage1 高分辨率处性价比低。
3.3 CA 注意力:把坐标信息塞回通道注意力里
CA(Coordinate Attention)针对的就是 SE 丢空间信息这个短板。它的做法是把全局池化拆成两个方向:沿宽度方向做池化得到 $C \times H \times 1$,沿高度方向做池化得到 $C \times 1 \times W$。这两份特征分别编码了“在哪一行”和“在哪一列”的坐标信息。然后拼接、过共享卷积降维、再拆开,各自通过卷积恢复通道数并 Sigmoid,最后作为权重乘回原特征。
这个设计的好处是,权重不再是单一通道标量,而是带方向的位置敏感权重。做遥感图像里的细长目标检测时,我在同一 backbone 上对比过 SE 和 CA,CA 在小目标召回上大约高出 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 里都插 CBAM,FLOPs 涨了 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, device='cuda', dtype=torch.float16) k = torch.randn(8, 16, 1024, 64, device='cuda', dtype=torch.float16) v = torch.randn(8, 16, 1024, 64, device='cuda', dtype=torch.float16) with torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=False, enable_mem_efficient=False): out = F.scaled_dot_product_attention(q, k, v, is_causal=True)用sdp_kernel上下文管理器强制走 Flash 后端,如果后端不可用会直接报错,而不是静默回退到慢速实现。这一点很重要:静默回退会让你以为优化生效了,实际上还在跑数学实现。
维度上我实测过一组数据,序列长度 4096、batch 8、头数 16、头维 64,FP16 下标准实现峰值显存约 4.2 GB,SDPA 走 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)" # 安装 Triton(Windows 上通常用社区维护的预编译包) 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, N+1, N+1),取第一个样本、对多头求平均 w = attn_weights[0].mean(dim=0) # (N+1, N+1) 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 | 整行被 mask,softmax 全 $-\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 后单头表达力下降明显,收益开始变负。我做消融时 $h=32$、头维 16 的配置比 $h=8$ 低了约 0.8 个点。
关于 dropout,注意力权重上的 dropout(attention 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 错误在这个规模下都会暴露得清清楚楚,修好之后直接放大配置,基本不会再有结构性问题。这套“缩小到玩具规模再放大”的流程,我在多头自注意力、交叉注意力、视觉注意力模块上都用过,屡试不爽。