ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

注意力机制原理与PyTorch实战:从QKV到多头注意力

注意力机制原理与PyTorch实战:从QKV到多头注意力 这段时间一直在折腾NLP项目在多个文本任务里反复锤炼“注意力机制”这个核心模块从最初读Transformer源码时的懵圈到后来能在新闻处理、在线问诊等真实场景里独立完成代码实战和效果调优中间踩过的坑不少所以这次把“从原理到代码实战”整条链路拆开来讲。这篇文章适合三类人刚接触NLP、想搞懂QKV到底是什么的新手已经会调库但想理解内部细节的工程师以及打算在业务项目里对注意力做魔改的开发者。我会从最基础的动机讲起给出手写的PyTorch实现最后聊聊训练时的坑和工程上的取舍保证你读完能直接抄作业也能举一反三。1. 先想清楚注意力机制到底在解决什么问题1.1 从固定向量到动态对齐注意力机制出现之前的痛点很多文章一上来就甩公式我觉得不对得先从问题出发。早期做机器翻译主流方案是Encoder-Decoder结构。Encoder把源句子整个读一遍最后一个时间步的隐状态当作一个“压缩包”所有的语义信息全靠这个固定向量传给Decoder。短句还好长句子要命句子越长这个向量就越像把十本书的内容塞进一个快递箱装不下也理不清。于是Decoder往往在长句子的后半段就开始遗忘翻出来的内容丢三落四。注意力机制解决的就是这个“信息瓶颈”。核心思想很直白Decoder在生成每个词的时候不要只盯着固定向量而是去源句子里面动态地找和当前词语对齐的部分。比如翻译“I love you”生成“爱”这个字时模型应该把注意力集中在“love”这个词上而不是把整句话平均看待。这种动态打分的机制就是软对齐也叫注意力分布。我在真实项目里体会很深。做新闻信息抽取时输入是一篇几百字的长文要让模型输出事件类型和关键人物。如果不用注意力机制底层依赖的CNN或RNN很容易被长句干扰接入注意力后模型在抽取“人物”时会明显倾向于关注“发表”“指出”“吕梁”这类承载主体信息的词准确率高出一大截。1.2 软对齐机制用一句话解释注意力在做什么注意力机制本质上就是加权求和。每一个输出位置对输入的所有位置算一个“匹配分数”然后通过softmax变成归一化的权重再按权重把输入对应的内容加权累加起来。匹配度高的位置拿到更大权重匹配度低的位置拿到近似0的权重。打个生活化的比方你在图书馆找一本关于“深度学习”的书。你不会把每一层书架都整排整排地仔细读一遍而是先看标签目光会自动落在有“机器”“神经网络”“AI”这些关键词的区域。你的目光就是Query书架上的标签就是Key标签对应的书就是Value。视线停留时间长的书其实就被赋予了高注意力权重。这个类比对应到公式上就是经典三件套Query你想找的东西代表“当前需要什么信息”。Key输入里每个位置的“标签”代表“我这里有什么”。Value输入里每个位置的实际内容代表“被提取的信息本身”。注意力分数就是Query和Key的匹配度最后按照匹配度对Value做加权求和。理解了这套设定后面所有的变体包括自注意力、多头注意力、交叉注意力都是在这个基础上加限制、换来源、改视野。2. 从机器翻译场景拆解Attention的核心公式2.1 Q、K、V到底分别是什么注意力机制里最劝退初学者的就是Q、K、V这三个字母其实没那么玄。以最早那篇《Neural Machine Translation by Jointly Learning to Align and Translate》为例Decoder在生成第t个目标词时会把Decoder当前的隐状态当作Query把Encoder的所有隐状态当作Key和Value。Query和每一个Key计算相似度得到一组权重再用权重去加权求和所有Value得到一个上下文向量供当前词解码使用。从实现角度看这三个角色未必是同一个来源。比如在机器翻译的decode阶段Query来自目标语言侧Key和Value都来自源语言侧这就是交叉注意力的雏形。到了自注意力里三者都来自同一句话每个词既当查询者又当被查询者。但不管是哪种场景底层都需要做投影用三个可学习的权重矩阵Wq、Wk、Wv把原始输入映射成Q、K、V。投影的目的是把原始embedding空间变换到若干不同的语义空间让匹配和提取各司其职。我见过很多人直接拿输入的embedding当作Q、K、V来算即QKVx。这样不是完全不行比如简单的文本分类还能跑但表达力严重受限。因为模型没法学会“在不同维度上用不同的方式去比较和抽取”这会影响后面层数的加深以及多头的效果。到Transformer里几乎都是输入x先各自过一层线性变换再做attention这是标配。再强调一个点Q、K、V的维度要一致吗不必要。Q和K的最后一维必须一致因为要做点积Value的最后一维可以和它们不一样因为加权求和的输出维度只和Value的维度有关。工程实现里通常让三者维度一致统一用d_model方便拼装但理解上要分清这个区别。2.2 点积之后为什么要除以sqrt(dk)这个点几乎逢面试必问。Q和K点积之后的结果会随维度变大而膨胀导致softmax的输入值过大进入梯度饱和区。具体来说如果Q和K每个元素都是均值0、方差1的独立随机变量那么它们点积结果的均值还是0但方差等于维度dk。dk越大点积分布的方差越宽分数很容易出现很大的正数或负数。想象一下softmax函数输入从-3到3这段范围梯度还算正常一旦分数变成30、50这种极端值softmax就会输出一个接近one-hot的分布绝大多数位置权重趋近于0单个位置的权重趋近于1。这不仅让模型“盯死”某个位置失去了灵活对齐的能力更重要的是梯度变得很平反向传播时更新信号弱训练不稳定。除以sqrt(dk)之后点积结果的方差被拉回1附近softmax曲线正好落在正常工作区间。数学推导很直接假设Q和K独立同分布每个分量均值0方差1那么单个维度的点积项方差是1dk个维度累加后方差是dk除以sqrt(dk)把方差重新归一化到1。这就是scaled dot-product attention里“scaled”的由来。实际操作中你甚至可以直接除以一个可学习的温度参数效果也能调出来。但固定用sqrt(dk)更省心因为它不引入额外参数而且被证明在多种任务上都能保持训练稳定。我调试模型时经常会检查attention分数的量级如果发现softmax分布过于尖锐或过于平坦优先检查是否做了scale往往会有奇效。2.3 软注意力与硬注意力的取舍注意力机制从决策方式上分软硬两种。硬注意力指的是从输入序列里“选出一个”位置忽略其他位置类似argmax操作。问题是argmax不可导训练时没法直接用梯度回传常见做法是用强化学习来近似优化过程繁琐且不稳定。凡是要在真实业务里快速落地我都不建议碰硬注意力。软注意力则对所有位置做加权求和权重分布在0到1之间整体可微可以直接端到端训练。我们平时说的“带softmax的注意力”就是软注意力也是Transformer标配。它表面上比硬注意力“浪费”了一些算力因为每个输出位置都要看完全部输入但换来了梯度稳定和表达灵活这笔交易非常划算。需要补充的是masked attention就是一种软注意力的变体。我们不是把所有位置的权重限制为0而是通过把某些位置的分数设为负无穷让softmax自动把它们压到趋近0形式上依然是软注意力。后面讲到padding mask和causal mask时会具体写代码这里先记住mask是在softmax之前对分数做掩膜而不是在softmax之后粗暴地把权重设为0否则归一化就不对了。3. 从零手写Attention代码实现与维度推演3.1 先写一个通用点积注意力函数贴一段我最常用的基础实现刻意不用PyTorch封装好的nn.MultiheadAttention就是为了看清楚每一步在做什么。输入输出都设计成可以直接嵌入自定义模型的形状。import torch import torch.nn.functional as F import math def scaled_dot_product_attention(Q, K, V, maskNone): Q: [batch_size, seq_len_q, d_k] K: [batch_size, seq_len_k, d_k] V: [batch_size, seq_len_k, d_v] mask: [batch_size, seq_len_q, seq_len_k] 或可以广播的形状 d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, V) return output, attn_weights这段代码的关键点就三个。第一K.transpose(-2, -1)把Key的序列长度维度交换到前面才能让Q的每个查询位置和K的所有位置做点积。第二除法缩放放在softmax之前顺序不能错。第三mask用masked_fill(mask 0, -1e9)把需要屏蔽的位置填成负无穷而不是填0。填0的话softmax之后依然会有一定权重达不到真正屏蔽的目的。注意负数填充值我习惯用-1e9只要足够小就行。如果你用的是FP16混合精度训练建议直接填torch.finfo(scores.dtype).min或者-65504这类在FP16范围内的极小值。我早期在混合精度训练里吃过亏填了-1e9在某些极端情况下会溢出出nan换成-65504之后问题消失。3.2 自注意力实现从词向量到上下文表示自注意力就是让一句话里的每个词都能根据整句话的语义更新自己的表示。比如“苹果”这个词旁边是“华为”还是“富士”注意力权重会告诉模型应该朝哪个方向理解。我写一个用于实际调试的最小自注意力模块。import torch import torch.nn as nn import math class SelfAttention(nn.Module): def __init__(self, d_model, d_kNone, d_vNone): super().__init__() self.d_model d_model self.d_k d_k if d_k is not None else d_model self.d_v d_v if d_v is not None else d_model self.W_q nn.Linear(d_model, self.d_k) self.W_k nn.Linear(d_model, self.d_k) self.W_v nn.Linear(d_model, self.d_v) def forward(self, x, maskNone): Q self.W_q(x) # [batch, seq_len, d_k] K self.W_k(x) # [batch, seq_len, d_k] V self.W_v(x) # [batch, seq_len, d_v] scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn_weights F.softmax(scores, dim-1) context torch.matmul(attn_weights, V) return context, attn_weights在这个实现里输入x的形状是[batch, seq_len, d_model]输出context的形状是[batch, seq_len, d_v]。如果你想让输出和输入形状一致最常见的做法是把d_v设成d_model这也解释了为什么很多模型里d_model从头到尾都不变。每个token得到的新向量实际上是全序列所有token的Value的加权混合因此每个token都“看”到了上下文信息而且看的范围是整个序列。这里我想多说一句维度推演的事。很多初学者写代码跑通靠凑一到自己改结构就崩。务必在纸上推一遍x进入W_q后变成[batch, seq_len, d_k]转置K变成[batch, d_k, seq_len]两个矩阵乘出来是[batch, seq_len, seq_len]这就是注意力矩阵softmax沿着最后一个维度做保证每一行的权重加起来等于1最后和V相乘V是[batch, seq_len, d_v]结果是[batch, seq_len, d_v]。每一步都对得上代码基本不会出维度错误。3.3 Maskpadding mask和causal maskmask是注意力里最容易被忽视又最容易写错的部分。它在NLP里有两个典型场景。第一个是padding mask。一批文本长短不一要pad到相同长度才能组成batchpad位置是无效信息。如果不对这些位置处理模型会把“补零”当成真实的词参与计算还可能在生成任务里输出奇怪的填充尾巴。做法是在计算scores之后把pad位置对应的score填成-1e9softmax后这些位置自然趋近0。通常你构建的mask是[batch, seq_len]值为1表示有效位置、0表示pad位置。送到上面的函数里需要先扩展维度到[batch, 1, seq_len]才能和[batch, seq_len, seq_len]做广播。第二个是causal mask也叫因果掩码用在自回归解码器里。生成目标序列时当前位置只能看到自己和左侧的内容不能看到右侧未来的词。实现方式是构造一个上三角矩阵对角线右上方的位置全部填成0让softmax把未来位置压掉。# causal mask 示例 seq_len 5 causal_mask torch.tril(torch.ones(seq_len, seq_len)).bool() # 实际上注意力的 scores 形状可能是 [batch, seq_len, seq_len] # 用 causal_mask.unsqueeze(0) 广播即可写mask时最容易犯三个错mask位置写反把有效位置屏蔽了mask没有扩展到多头维度导致部分头正常工作、部分头输出全等于mask对应值以及用了masked_fill(mask, -1e9)但mask的布尔方向搞错。排查方法很简单把scores和mask分别打印出来检查被屏蔽位置在softmax前的值是否足够小。这个方法救了我很多次比我盯着代码看半天管用得多。4. 从单头到多头Multi-Head Attention的代码实战4.1 多头注意力的动机单头注意力的问题在于表达单一它只能学到一种匹配模式。某种情况下适合按语义相似度做匹配另一种情况下可能适合按距离远近做匹配单头模型没办法在同一时刻同时兼顾多种模式。多头注意力的做法是把Q、K、V投影到h个不同的子空间每个子空间独立计算注意力再把结果拼接起来最后过一个线性层。可以类比成开会时请了好几个领域的专家有人专盯语法结构有人专盯指代关系有人专盯情感色彩每个专家的视角不同汇总到一起才完整。深度学习里这种“多视角”设计很常见类似CNN里用多个卷积核提取不同特征。我实际看注意力可视化时也发现确实有的头关注相邻词有的头关注远距离核心词有的头关注标点和停顿位置。多头并不是随意设置的头数h和每个头的维度d_k需要满足整除关系常见做法是d_k d_model / h。比如d_model是512h是8每个头的维度就是64。头数太小表达力不够头数太大单头维度太低每个头学到的匹配过于碎片化。实践中8头或16头比较常见。4.2 PyTorch实现多头注意力直接上代码注释里写清楚每个view和transpose的作用这是最容易秃头的一段。import torch import torch.nn as nn import math class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads self.W_q nn.Linear(d_model, d_model) self.W_k nn.Linear(d_model, d_model) self.W_v nn.Linear(d_model, d_model) self.W_o nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def forward(self, query, key, value, maskNone): batch_size query.size(0) Q self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn_weights self.dropout(F.softmax(scores, dim-1)) context torch.matmul(attn_weights, V) context context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) return self.W_o(context)这里关键的一步是分头。输入query原本是[batch, seq_len, d_model]先经过W_q保持形状不变然后view(batch_size, -1, self.num_heads, self.d_k)把最后一维拆成num_heads和d_k两段再用transpose(1, 2)把num_heads换到第二个维度最终形状变成[batch, num_heads, seq_len, d_k]。这样矩阵乘法时每个头独立操作互不干扰。算完注意力后要把头拼回去所以先transpose(1, 2)还原再contiguous().view重新展平成[batch, seq_len, d_model]。如果你用的是PyTorch自带的nn.MultiheadAttention记得设置batch_firstTrue否则输入输出形状都和常规TensorFlow风格不一样容易传错。自带的实现还支持add_bias_kv、kdim、vdim这些参数适合快速验证模型但真要自己魔改结构还是自己写一版更顺手。我自己在项目里两种方案都试过兜底原型用自带模块正式迭代时换成自写的可控性更强。4.3 交叉注意力业务项目中的常客多头注意力本身是通用框架稍微改一下输入来源就变成了交叉注意力。交叉注意力的核心是Query来自一个序列Key和Value来自另一个序列。举个例子做机器翻译时Query是目标语言的解码状态Key和Value是源语言编码器的输出模型在每个目标词上动态地在源句里寻找对齐信息。在RAG类系统里更常见。用户输入一个query问题系统检索出若干相关文档片段然后把query当作Query把文档片段当作Key和Value注意力模块负责在文档里抽取出真正和用户问题相关的信息。我之前做新闻问答系统就是这么设计的底层用BM25或向量检索召回若干篇新闻上层用交叉注意力把query和新闻段落融合效果比简单的向量拼接好很多。交叉注意力的实现代码和自注意力几乎一样只是forward的query、key、value传参不同。你需要留意的只有一点mask的形状要跟着key的序列长度走因为注意力矩阵的行数是query的长度列数是key的长度。很多人在这个上面踩坑mask维度写错轻则广播报错重则数据泄漏。5. 位置编码别让Attention搞乱词语顺序5.1 不带位置信息的自注意力无法区分语序自注意力机制本身是“集合操作”它天然不考虑词语的前后顺序。对模型来说“我打你”和“你打我”在进入自注意力之前如果不做任何额外处理完全是同一组token。可语义就差之千里了。这就是Transformer要引入位置编码的根本原因。注意这里说的位置编码和时序注意力不是一个东西。时序注意力通常指在序列维度上学习一个权重向量用于筛选哪些时间步更重要而位置编码是往输入里注入顺序信号。两者解决的问题不同一个是对“位置重要性”做加权一个是给“位置身份”做标记。有人可能会问RNN为什么不处理这个问题因为RNN天生按时间步顺序读取位置信息隐含在循环结构里。CNN可以通过不同大小的卷积核覆盖局部窗口隐含部分局部位置信息。而Transformer的多头注意力让每个token直接连到所有位置结构上缺乏内在的顺序感所以必须显式地补上位置信息否则模型对语序完全无感。5.2 正弦位置编码的直觉与代码Transformer原版用的是正弦位置编码公式长这样PE(pos, 2i) sin(pos / 10000^(2i / d_model)) PE(pos, 2i1) cos(pos / 10000^(2i / d_model))其中pos是token在序列里的位置i是维度下标。不同维度拥有不同的频率维度号越小频率越低。这样每个位置都有一个独一无二的编码向量同时不同位置之间的相对关系可以通过向量的线性变换近似表达这让模型在一定程度上能感知“两个词相隔多远”。我给出一个简单的PyTorch实现频率计算用指数形式避免溢出import torch import math def sinusoidal_position_encoding(max_len, d_model): pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp( torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model) ) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) return pe.unsqueeze(0) # [1, max_len, d_model]使用时把这个矩阵加到输入embedding上即可也就是x embedding pe。注意embedding和pe的形状要兼容通常是[batch, seq_len, d_model]和[1, seq_len, d_model]做广播加法。5.3 位置编码在实际工程中的选型正弦位置编码的好处是外推性相对较好模型即使在训练时没见过特别长的序列也能在推理时对超出部分给出一个合理的位置向量。我在早期做长文本抽取时确实依赖过这个特性。但后来实测发现如果训练时最大长度只有128推理长度拉到256效果还是会明显下滑。外推不是万能的想要真正处理超长文本还是得靠RoPE这类相对位置编码方案或者干脆做分段处理。可学习位置编码是另一个思路直接用nn.Embedding(max_len, d_model)去训练一个位置向量表。这种方案胜在灵活模型可以根据具体任务自己调整位置表示数据量足够时效果不会比正弦差甚至更好。缺点是它无法平滑外推训练没见过的位置id直接表现拉跨。我的选择标准很粗暴序列长度固定或比较短直接用可学习位置编码文本长度波动大或者推理时可能超过训练长度用RoPE或者Alibi这类专门优化过外推的编码。还有一个被忽略的小细节位置编码是加在词向量上而不是拼在特征维度上。加和操作会损失一部分显式的坐标信息但好处是不增加额外参数和复杂度模型可以靠后面的层把这些信息重新解码出来。工程上如果算力紧张且序列不是特别长可以先尝试可学习位置编码迭代最快。6. 实战答疑训练中常见的Attention问题与排查6.1 模型不收敛先查这四件事注意力机制虽然强大但它不是魔法训练起来比RNN更挑剔。我总结过一套排查顺序几乎能覆盖大多数不收敛问题。第一查mask。mask的位置反了或者形状错了模型会看到一堆无效信息loss在早期就会表现异常。验证办法是写一个小测试用例把padding序列和正常序列单独测一下查看attention权重是否把pad位置压到接近0。第二查scale。没有除以sqrt(dk)或者dk写错都会导致softmax进入饱和区。打印scores的均值方差正常情况应该大致落在0和1附近而不是几十上百。第三查学习率和warmup。Transformer这类模型对学习率非常敏感尤其刚训练时直接用太大的学习率很容易把分布打崩。常见做法是先warmup几百步把学习率从小到大逐步提升再用衰减策略。如果不做warmup前期loss可能会先上升一截再下降不明真相的人会以为模型坏了。第四查数据长度分布。如果pad比例太高比如所有样本都pad到512但实际长度只有20模型把大量计算花在了无效位置上训练效率极低甚至学到一堆无用模式。这时候要么做动态padding要么做bucketing按长度分桶再padding效率能翻几倍。6.2 注意力分布过于均匀是怎么回事有时候训练loss降到一定程度就停滞打印attention矩阵发现权重都差不多每个位置都分到一小杯水。这种情况说明模型没有学到有效的区分模式。我遇到最多的原因是Q和K的语义空间重叠太大匹配分数区分度不足。处理手段有几个。第一增大d_model给模型更多的参数空间来分离语义。第二调整初始化让W_q和W_k初始权重差异更大一点。第三适当降低dropout太强的dropout会抹平注意力差异让人感觉“谁都重要谁也都不重要”。第四检查是不是输入特征太弱比如只是稀疏的one-hot没有经过充分的embedding学习。如果你是做新闻分类这类类间差异较小的任务注意力分布均匀的问题会更突出。我试过在Transformer底层用预训练模型的输出替代随机embedding均匀现象会缓解很多因为预训练表示本身已经带上了丰富的语义区分信息。还有一种实用招数是在loss里加一个辅助的attention正则项比如鼓励某些头关注局部窗口、某些头关注全局不过我一般最后才用这招。6.3 长序列带来的显存与复杂度优化注意力机制的复杂度是序列长度的平方512个token还好到2048就非常吃显存。很多业务场景又偏偏需要长文本比如整篇新闻、聊天记录、病程记录这时候优化就是硬需求。我整理过一个表格按投入产出比从高到低排列方案核心思路适用场景截断和分段把长文切成段落分别建模大多数分类问题简单直接抽取关键句先用小模型选热点句再进Transformer事件抽取、问答窗口注意力每只关注局部窗口窗口外忽略语言建模、生成稀疏注意力部分头看全局部分头看局部长文本通用任务降维注意力先压缩序列再计算资源紧张的在线服务实际项目里我常驻的策略是“先粗后细”先用廉价方式筛选出和任务最相关的句子再对筛选结果做全量注意力。比如做新闻事件抽取时先让TF-IDF或一个简单的CNN分类器找出包含事件主体、时间、地点的句子抽出来后再跑Transformer这样显存开销大减精度还有提升。很多人一上来就想上Longformer除非你是学术研究或者预算充足否则业务场景优先做“减法”往往更划算。7. 一些工程上的取舍和个人体会7.1 用注意力权重做初步排查注意力权重最实用的功能之一就是做模型诊断。训练完一个文本分类模型后我会挑几个预测正确和预测错误的样本把最后一层的attention权重可视化出来画成热力图。预测正确的样本通常能明显看到模型聚焦于和标签强相关的词上预测错误的样本往往注意力飘到了无关词甚至标点上。我在做在线问诊文本抽取项目时就用这招发现过一个有意思的问题模型抽“主诉症状”时经常把“没有不适”里的“不适”当成正例。看热力图才发现模型过度关注了“不适”这个词本身而没有关注前面的否定词“没有”。后来在数据里增加了否定短语的覆盖症状抽取的F1值从0.76涨到0.84。需要提醒的是注意力权重不应该被当作严格的可解释性证据。学术界已经有不少论文指出attention weights和feature importance之间不能直接划等号它只能作为“模型在看什么”的初步线索。但用来排查明显的错误模式和做bad case分析这个手段非常有效性价比极高。7.2 我的Attention调参路径最后分享一条我在实战中反复验证的调参路径不一定适用所有任务但可以当作起点。模型结构方面d_model建议从256或512起跳再根据数据量加减。num_heads常用8或16要保证d_model能整除。dropout在大部分任务里设0.1不会错但如果数据量很小考虑dropout降到0.05甚至去掉。如果数据规模真的很小比如几千条不建议从零训练Transformer直接拿一个轻量预训练模型做迁移把注意力模块当作抽取特征的工具效果会好得多。训练策略方面Learning Rate建议配合warmupwarmup steps占整体训练步数的5%到10%。如果用Adam注意把epsilon设置到1e-8或者1e-6某些PyTorch版本默认的epsilon在不同精度下表现差异明显。FP16训练时还要额外注意损失缩放策略和极端值的掩膜处理就是我在前面提到过的负无穷填充值问题。我自己在新闻处理和医疗文本抽取这两类任务上反复做对比结论是注意力机制本身不是银弹真正决定效果上限的是数据质量和任务设计。把QKV理解透彻能让你在写代码时少走弯路也让你在遇到模型不收敛、训练波动、效果不佳时知道从哪里下手而不是病急乱投医。这也是这篇文章最想传达的东西——你手里那套注意力代码远没有看上去那么神秘拆开揉碎以后每一个算子都有它的存在理由。
返回列表