资讯动态

手写Transformer:从注意力机制到代码实战

发布时间:2026/9/9 2:11:25 来源:尧图企业网站定制
初识Transformer很多人第一反应是打开那篇著名的《Attention Is All You Need》然后被一堆术语劝退。我自己当年也是这样什么QKV、多头注意力、位置编码、自回归每个词都认识连成句子就完全不知道在说什么。后来硬着头皮手推了一遍代码再把各个模块拆开重新组装才真正建立起了这东西其实不复杂的直觉。这篇博文就是想把当初那条弯路帮你省掉从一个初学者的视角把Transformer从名气到原理、从手写代码到实际部署一层层讲清楚。内容不追求严格的数学证明只求让你读完能知道Transformer到底在干什么、每一步为什么要这样做以及它为什么能横扫NLP、CV甚至时间序列预测。1. 初见的困惑为什么偏偏是Transformer取代了RNN和CNN1.1 从按顺序读到一眼看全文在Transformer出现之前处理序列数据的主流方案是RNN、LSTM、GRU这一类循环神经网络。它们的工作方式很像一个人在读书一个字一个字地看而且必须记住前面看过的内容才能理解当前这个字。这种按顺序读的机制有两个天然缺陷。第一个是无法并行。因为第t个时间步的隐状态依赖于第t-1个时间步的输出整个训练过程只能老老实实地串行计算。你给它一块再大的GPU算到第1000个字的时候仍然要等前面999个字的计算结果出来。模型规模一大训练成本就直线上升。第二个是长距离依赖问题。虽然LSTM和GRU通过门控机制缓解了梯度消失但信息在序列中每传递一步都会有一点点衰减。当序列长度超过几百甚至上千时靠循环传递的方式很难把很早之前的信息原封不动地保留到需要它的位置。文本理解、语音识别这类任务偏偏就是经常需要用到很久之前的上下文。CNN则走了另一条路它靠卷积核在局部窗口内提取特征天然适合图像这种空间结构。把CNN搬到文本上通过堆叠层数来扩大感受野虽然能并行计算但每扩大一次感受野都要多加几层网络效率低而且局部优先的归纳偏置和一句话里遥远单词之间的强关联并不完全匹配。Transformer的选择很直接我不循环了也不滑动窗口了一个序列进来我直接让任意两个位置之间两两交互。这种全局连接的方式配合矩阵运算让所有位置的计算天然就是并行的而且不管距离多远信息交互都是直接一步到位不存在中间衰减。1.2 一个类比全连接会议 vs 传纸条如果你想直观理解Transformer与传统序列模型的区别可以想象开一场全员会议。RNN的开会方式是一个接一个发言第一个人说完第二个人才能发言而且第一个人说的话只能通过会议纪要传递下去中途如果纪要写得不完整后面的人就拿不到那个信息。CNN的开会方式是分组讨论每桌只和隔壁桌交流想要让离得远的两个桌交换信息得通过更多轮分组讨论让信息一层一层传过去。Transformer的开会方式则是所有人围成一圈每个人都同时看全场的所有发言然后自己判断谁的话和我当前这句话关系最大我就重点关注谁。这个重点关注的权重不是事先定死的而是根据内容动态计算出来的也就是注意力机制。这个类比基本解释了Transformer所有架构设计背后想解决的事。理解了全局动态交互这个核心后面看QKV、多头注意力这些概念就知道它们只是实现这个目标的具体手段。2. 解剖一只麻雀从注意力机制到Encoder完整结构2.1 Q、K、V到底在图什么Transformer的所有核心都在一个叫Scaled Dot-Product Attention的公式里Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V很多人第一次看到这个公式就晕了Q、K、V是什么为什么要除以根号d_k为什么用softmax我用一个实际的场景来解释。假设你正在读一句话那只猫因为太胖没能跳上冰箱。读到猫这个字的时候你不确定它和胖冰箱跳之间到底是什么关系。注意力机制做的事就是给这句话里的每个词分配一个权重决定在理解猫的时候应该花多少精力去看胖、看冰箱、看跳。QQuery是当前这个词发出的提问可以理解成我正在找什么信息KKey是序列中每个词给自己的标签可以理解成我身上带有什么信息VValue是每个词实际携带的内容一旦确定了该关注谁就把对方的内容取出来用。Q乘以K的转置算的是我的提问和你身上的标签匹配不匹配得到一个原始分数。除以\sqrt{d_k}是为了防止维度变大时点积数值过大导致softmax落入梯度饱和区。softmax则把这些分数归一化成一组和为1的权重。最后用这些权重去加权求和所有的V就得到了当前词的输出向量。一开始你会觉得这套机制很抽象但亲手写一次代码就会发现整个过程就是几个矩阵乘法十几行就能写完完全不神秘。2.2 多头注意力让模型同时听多种意见既然注意力是开会那多头注意力Multi-Head Attention就是开了好几场不同的会每场会关注不同的重点。头数通常是8或16每个头有自己独立的Q、K、V权重矩阵所以每个头学到的交互模式是不同的。以苹果这个词为例。一个头可能在关注这个苹果的颜色另一个头在关注谁在吃它第三个头在关注它出现的时空背景。把这些头的输出拼起来再经过一个线性变换融合模型就能同时从多个角度理解同一个词。多头注意力还有一个好处它相当于给模型提供了多个子空间的表示。每一个头都有自己的Q、K、V矩阵随机初始化之后通过训练被推向不同的语义方向。我见过一些可视化实验一个头明显负责句法关系比如动词和主语另一个头负责指代消解代词指向谁分工非常清晰。这让模型的表达能力远超过单头注意力。2.3 位置编码打乱顺序为什么还认识其实不认识自注意力机制本身对序列顺序是完全无感的。你把我打你换成你打我三个词的两两交互结果完全一样但语义完全相反。所以Transformer必须具备一种把位置信息注入网络的手段这就是位置编码Positional Encoding。原版论文用的是三角函数编码公式是PE(pos, 2i) sin(pos / 10000^(2i/d_model)) PE(pos, 2i1) cos(pos / 10000^(2i/d_model))简单说每个位置被编码成一组正弦和余弦值不同维度对应的频率不一样。就好像给每个词贴上一个多维坐标频率较低的那些维度编码较长距离的位置关系频率较高的维度编码较近距离的关系。论文选三角函数是有道理的。一方面其值域在[-1,1]之间不会把嵌入向量的尺度带偏另一方面它不需要训练直接计算即可而且外推到比训练时更长的序列时表现也相对平稳。后来的研究中也有直接用可学习位置嵌入的做法比如BERT就用了可学习的绝对位置嵌入效果在固定长度下通常也不错只是外推能力差一些。很多初学者会忽略位置编码的重要意义以为它只是一个锦上添花的组件。实际上如果去掉位置编码Transformer在绝大多数序列任务上的表现会断崖式下降因为它真的就变成了一个对顺序毫无感知的词袋模型。2.4 残差连接、LayerNorm与前馈网络一个标准的Encoder Layer由四块组成多头注意力、残差连接、LayerNorm、前馈网络。它们按固定的顺序组合输入向量先送进多头注意力层得到输出后和原始输入相加残差连接再做LayerNorm归一化然后送进一个由两个线性层和ReLU激活函数组成的前馈网络再残差连接和归一化一次得到这一层的最终输出。残差连接解决的问题是深层网络的梯度传播。Transformer通常要堆6层甚至更多层Encoder没有残差连接的话深层网络很容易出现梯度消失训练根本推不动。加入了抄近道之后梯度可以从最后一层直接传回第一层。LayerNorm做的是在每一个样本的特征维度上做归一化让数据分布保持稳定。它和BatchNorm不同BatchNorm是在一个batch的所有样本上统计均值和方差而LayerNorm只在当前样本内部统计。对于变长序列和RNN/Transformer这类结构LayerNorm往往更稳定因为它不依赖batch size训练和预测时的行为一致。前馈网络则是给模型增加非线性变换能力和特征交互的空间。注意力层本质上是对信息做筛选和聚合而前馈网络对每个位置的向量独立做一次加工和变换。没有这一层整个模型会退化成浅层的线性变换表达能力大大受限。2.5 Encoder与Decoder的分工原版Transformer是Seq2Seq架构由Encoder和Decoder共同组成。Encoder负责把输入序列编码成一组富含上下文的表示向量Decoder则负责根据这些向量以及已经生成的历史输出逐字生成目标序列。Decoder和Encoder最大的区别有两个。第一Decoder里多了一个Masked Multi-Head Attention也就是在做自注意力时把未来位置的信息全部遮掉。因为解码过程是逐词生成的生成第t个词的时候不能偷看t1以及之后的词否则就相当于考试时偷看了答案。第二Decoder中间插入了一层Cross-Attention让解码器的每个位置都能和编码器输出的所有位置做注意力交互这样就能从输入序列中获取需要的信息。这两处设计是先理解再生成任务比如机器翻译的灵魂。理解了Encoder和Decoder的分工再去看GPT和BERT的差异就非常容易GPT就是只保留Decoder或者更准确地说是一个自回归的TransformerBERT则是只用Encoder双向编码的Transformer。两者架构同源但因为训练目标和侧重点不同演变成了两个不同的门派。3. 手写一个微型Transformer核心代码逐行拆解3.1 环境准备与数据约定看十遍架构图不如亲手写一遍。我建议你用PyTorch把Transformer的核心组件从零实现一遍不需要调库自己写多头注意力和TransformerBlock这样能彻底打消黑盒恐惧。如果你还没有安装PyTorch先执行pip install torch numpy我们做一个非常简单的任务来验证模型给定一串随机生成的整数序列让模型学习把序列中每个位置的数值乘以2再输出。这听起来简单但它能很好地验证自注意力、位置编码和前馈网络真的在起作用。如果你愿意后面也可以换成英文翻译的小数据集。代码约定是序列长度seq_len为20嵌入维度d_model为64头数heads为8因此每个头的维度d_k 64 / 8 8。层数num_layers为2。模型输入是一个形状为(batch_size, seq_len)的整数序列。3.2 位置编码实现首先实现位置编码。我们希望得到一个形状为(1, seq_len, d_model)的矩阵它将被加到词嵌入上。import numpy as np import torch import torch.nn as nn def get_position_encoding(seq_len, d_model): pe np.zeros((seq_len, d_model)) for pos in range(seq_len): for i in range(0, d_model, 2): pe[pos, i] np.sin(pos / (10000 ** (i / d_model))) pe[pos, i 1] np.cos(pos / (10000 ** (i / d_model))) pe torch.FloatTensor(pe).unsqueeze(0) return pe这段代码会生成一个位置编码矩阵。第一次跑的时候你可以把它打印出来看看会发现相邻位置的编码值非常接近而距离远的位置差异逐渐增大。这个连续的位置指纹被加到词嵌入上之后模型在处理该位置时就能感知到它的大致位置。一个常见的易错点是忘记用unsqueeze(0)扩展batch维度。如果没有这一步和词嵌入相加时会因为维度对不上而报错。另一个易错点是循环里跳步range(0, d_model, 2)的方式偶数维度用sin奇数维度用cos顺序别搞反。3.3 多头注意力实现多头注意力是核心中的核心。先算Q、K、V然后分头缩放点积注意力最后拼接输出并线性变换。class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads): super().__init__() assert d_model % n_heads 0 self.d_model d_model self.n_heads n_heads self.d_k d_model // n_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_out nn.Linear(d_model, d_model) def split_heads(self, x): batch_size, seq_len, _ x.size() x x.view(batch_size, seq_len, self.n_heads, self.d_k) return x.transpose(1, 2) def forward(self, x, maskNone): batch_size, seq_len, _ x.size() Q self.split_heads(self.w_q(x)) K self.split_heads(self.w_k(x)) V self.split_heads(self.w_v(x)) scores torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtypetorch.float32)) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn torch.softmax(scores, dim-1) out torch.matmul(attn, V) out out.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) return self.w_out(out)这里最关键的操作是scores torch.matmul(Q, K.transpose(-2, -1))。Q的维度是(batch_size, n_heads, seq_len, d_k)K转置后维度是(batch_size, n_heads, d_k, seq_len)两者相乘就得到了(batch_size, n_heads, seq_len, seq_len)的注意力分数矩阵。这个矩阵的第i行第j列代表序列第i个位置对第j个位置的注意力分数。mask用-1e9填充而不是0原因是在做softmax之前把需要屏蔽的位置设为一个极大的负数softmax之后这些位置的概率就会趋近于0等价于完全忽略它们。头数是8时d_k 8所以scores比较大时除以根号8约2.83防止softmax输出过于尖锐。这个细节在原版论文里是专门强调过的千万别省。3.4 前馈网络与TransformerBlock组装前馈网络结构很简单两个线性层夹一个ReLU。第一层通常把维度放大4倍第二层再投影回去。class FeedForward(nn.Module): def __init__(self, d_model, d_ff2048): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) def forward(self, x): return self.linear2(torch.relu(self.linear1(x)))接着组装一个TransformerBlock把多头注意力、残差、LayerNorm、前馈网络按照标准顺序串起来class TransformerBlock(nn.Module): def __init__(self, d_model, n_heads, d_ff): super().__init__() self.attn MultiHeadAttention(d_model, n_heads) self.norm1 nn.LayerNorm(d_model) self.ff FeedForward(d_model, d_ff) self.norm2 nn.LayerNorm(d_model) def forward(self, x, maskNone): attn_out self.attn(x, mask) x self.norm1(x attn_out) ff_out self.ff(x) x self.norm2(x ff_out) return x注意残差连接的加法位置先把注意力输出和原始输入相加再做LayerNorm。顺序是先加后归一化不是先归一化再加这点和BERT等预训练模型保持一致。LayerNorm放在残差之后而不是之前是原版论文采用的Post-LN结构虽然训练时对学习率更敏感但原始实现和大多数经典教程用的都是它。最后把词嵌入、位置编码、若干TransformerBlock以及输出层拼成完整模型class TinyTransformer(nn.Module): def __init__(self, vocab_size, d_model64, n_heads8, num_layers2, d_ff256, max_len100): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.position_encoding nn.Parameter(get_position_encoding(max_len, d_model), requires_gradFalse) self.blocks nn.ModuleList([ TransformerBlock(d_model, n_heads, d_ff) for _ in range(num_layers) ]) self.out nn.Linear(d_model, vocab_size) def forward(self, x, maskNone): seq_len x.size(1) x self.embedding(x) self.position_encoding[:, :seq_len, :] for block in self.blocks: x block(x, mask) return self.out(x)把位置编码设为requires_gradFalse是合理的因为原版位置编码不需要训练。max_len是预设的最大序列长度超过这个长度就需要重新生成位置编码这解释了为什么Transformer外推到更长序列有时会受限。3.5 训练一个翻倍序列的玩具模型训练代码和普通PyTorch模型没有区别。我用MSELoss做回归你也可以把输出离散化成分类任务。生成训练数据时构造随机整数序列目标是序列里每个数乘2。import torch.optim as optim torch.manual_seed(42) def generate_batch(batch_size32, seq_len20, vocab_size100): x torch.randint(1, vocab_size, (batch_size, seq_len)) y (x * 2).float() return x, y model TinyTransformer(vocab_size100, d_model64, n_heads8, num_layers2, d_ff256, max_len50) optimizer optim.Adam(model.parameters(), lr1e-3) loss_fn nn.MSELoss() for epoch in range(50): total_loss 0 for _ in range(200): x, y generate_batch() optimizer.zero_grad() pred model(x) loss loss_fn(pred, y) loss.backward() optimizer.step() total_loss loss.item() if epoch % 10 0: print(fepoch {epoch}, loss {total_loss / 200:.4f})第一次跑你可能会看到loss下降得很快因为翻倍是一个线性变换自注意力模型并不需要真的去学多么复杂的关系。如果换成预测序列下一项的任务loss下降就会平缓很多。跑完这个玩具模型之后你已经有了一种亲手造出Transformer的掌控感再去看高级优化技巧会轻松得多。4. 不止于文本从ViT到时间序列预测的跨界玩法4.1 Vision Transformer把图片切成单词很多人以为Transformer只能处理文本但2020年ViTVision Transformer的提出直接把Transformer推到了图像分类的前沿。它的核心操作非常朴素把一张224x224的图片切成16x16的小块每个小块拉平后做一次线性映射变成向量然后当作单词输入标准的TransformerEncoder。ViT的关键是图片本身是有空间结构的如果直接打乱顺序模型无法知道谁在上谁在下。因此ViT同样加入位置编码只不过这里编码的是图像块在原始图片中的位置。同时还会在最前面加一个特殊的[CLS]标记这个标记在最后一层对应的输出向量就是整个图片的全局表示用来接分类头。从结果上看ViT在超大规模数据集上性能非常强超过了很多精心设计的CNN模型。不过它也需要更大的数据量来训练因为CNN自带的相邻像素之间有关联的归纳偏置在ViT中被削弱了所有关系都得靠数据学出来。这也解释了为什么后来的Swin Transformer会重新引入窗口化的多层金字塔结构本质上是在Transformer里重新加入一些局部先验。4.2 高光谱图像与遥感任务Transformer的用武之地最近几年高光谱图像分类和无人机多模态感知等任务里Transformer的身影越来越常见。高光谱图像每个像素点都有几十到上百个波段同时包含空间信息和光谱信息这对传统CNN来说非常难处理。空间维度上CNN可以通过卷积核捕捉邻域纹理光谱维度上几十个波段的序列关系更像一维序列。Transformer尤其是Swin Transformer这类分层结构天然适合同时建模空间邻域和光谱序列。在相关研究中常采用的做法是先用一个小的PatchEmbed把高光谱图像分成空间块然后通过位置编码保留每个块的空间坐标接着输入Vits的Encoder进行特征提取。实际测试中这类模型往往比纯CNN有更高的分类精度尤其是在训练数据充足的情况下。但代价是显存消耗大、训练时间长因此很多工作都在做轻量化比如用Restormer这类轻量级Transformer做图像复原任务通过优化注意力计算方式大幅降低参数量。4.3 目标检测与序列预测Attention真的是万能钥匙吗Transformer在目标检测领域的代表作是DETR它把目标检测重新定义成一个集合预测问题。传统检测器需要锚框、NMS非极大值抑制这些复杂的后处理流程而DETR直接让模型输出一组目标框和类别由二分图匹配算法和真实标注做对齐。检测器里的可变形注意力进一步优化了收敛速度只让模型关注目标周围少数的关键采样点而不是全图每个位置。这种思路在RGB-T双模态行人检测这类任务中也很有价值因为两个模态之间的信息配对本来就是一个天然的注意力问题。时间序列预测就更直接了。一个时间点过去的若干步就是序列未来的若干步就是需要解码的目标。Transformer可以学习到时间点之间的长期依赖比如电力负荷预测中今天中午的用电量可能和昨天中午、上周一中午都有关系这种多尺度周期相关性正是多头注意力擅长捕捉的。但要说Attention是万能钥匙也不尽然。Transformer在处理长序列时注意力矩阵的复杂度是序列长度的平方比如序列长度从1000涨到10000计算量会暴涨100倍。所以后来出现了很多稀疏注意力、线性注意力、分块注意力方案本质上都是要在全量互动和计算效率之间做取舍。对初学者来说先记住一个判断原则如果任务中长距离依赖真的重要Transformer值得试试如果任务本身就只依赖局部上下文传统CNN或简单循环模型可能更高效。5. 从一个手写Transformer引发的实战踩坑清单5.1 坑一忘了缩放点积loss就像过山车这是新手上路最常见的坑。注意力分数的计算公式里很多人会觉得反正后面有softmax除不除\sqrt{d_k}无所谓。但实际上不缩放的话当d_k比较大时点积结果会变得非常大softmax的梯度会几乎消失训练初期loss震荡特别严重甚至完全无法收敛。我自己调试过的最小复现不除\sqrt{d_k}训练50轮后loss还在高位波动加上缩放之后20轮内就能稳定下降。所以这个细节不是理论洁癖而是实实在在影响训练稳定性的关键设计。5.2 坑二维度变换时shape对不上注意力矩阵到底该是几维写多头注意力最容易懵的是view和transpose的顺序。很多初学代码会写成x x.view(batch_size, n_heads, seq_len, d_k)这会把一个序列强行切成头维度的形状完全打乱了数据排列。正确顺序是先把嵌入维度切成n_heads和d_k得到(batch_size, seq_len, n_heads, d_k)再transpose成(batch_size, n_heads, seq_len, d_k)。如果你在调试时看到scores的形状变成(batch_size, seq_len, batch_size, seq_len)之类的东西基本就是这里搞错了。一个很有效的调试技巧是在多头注意力的每一步都打印当前张量的shape对比预期值很快就能定位到具体是哪一行出了问题。5.3 坑三训练与推理不一致——Decoder的Mask到底怎么加如果你实现了完整的Encoder-Decoder版本还会遇到一个非常隐蔽的坑训练阶段Decoder是可以一次性输入完整目标序列的通过Mask遮住未来位置来模拟逐词生成但推理阶段没有真实目标只能先输入一个起始符生成一个词拼到序列末尾再输入模型生成下一个词。很多初学者在训练结束后直接把训练好的模型像CNN那样一次forward拿结果会发现生成质量极差。原因是模型在训练时从未见过不完整序列被输入Decoder的状态。正确的推理循环应该是for i in range(max_generate_len): output model(decoder_input) next_token output[:, -1, :].argmax(dim-1) decoder_input torch.cat([decoder_input, next_token.unsqueeze(0)], dim1)同时decoder部分要使用因果Maskcausal mask保证每个位置只能看到它之前的词。5.4 坑四学习率不是随便选的Transformer对学习率特别敏感Transformer和很多CNN模型不一样它对学习率极其敏感。原始论文用的是预热衰减的学习率策略前几步比如4000步学习率从很小的值线性增长到一个峰值之后按步数的倒数开方衰减。原因在于训练初期如果不做预热模型参数还没有稳定过大的学习率会让attention分布剧烈震荡。如果你只是跑玩具模型固定学习率也不是完全不行但注意取值范围最好在1e-4到3e-4之间。用1e-3训练小Transformer训练起来会很飘这一点和CNN很不一样。如果你只是跑玩具模型固定学习率也不是完全不行但注意取值最好在1e-4到3e-4之间用1e-3训练小Transformer很容易飘。另一个实用技巧是搭配梯度裁剪将梯度范数限制在1.0以内能显著减少训练过程中的loss尖峰。5.5 坑五显存爆炸——序列长度是指数级成本Transformer的显存开销不是随序列长度线性增长的而是近似平方级的增长。当你把序列长度从256调到512注意力矩阵需要的显存大约是原来的4倍。所以训练长文本模型时常见的做法是限制最大序列长度超过部分做截断或者分块使用梯度累积变相增大batch size而不增加显存峰值使用混合精度训练把部分计算改成FP16显存占用几乎减半尝试Flash Attention这类IO优化的注意力实现减少显存占用同时加速计算。这个坑在视觉任务里尤其明显。ViT把图片切成一个个Patch之后序列长度不短比如384x384的图片切16x16的Patch序列长度是57624x24。叠加12层Transformer之后显存消耗非常快。我当时的做法是先在256x256的小分辨率上验证模型结构和训练流程确认没问题再上大图避免反复跑着跑着OOM。5.6 坑六位置编码外推失败模型对长度变化很脆弱位置编码是固定频率三角函数时模型本质上只学过长度不超过max_len的位置表示。如果测试时突然给它一个比训练长度更长的序列那些位置完全没有对应的编码模型性能会明显下降。解决这类问题的方向之一是用RoPE旋转位置编码或ALiBi这类相对位置编码。它们不直接把位置编码加到词嵌入上而是通过旋转矩阵或偏置项注入相对位置信息外推能力更强。如果你只是做普通项目也可以像我在玩具模型里那样把max_len设成训练数据最大长度的1.5倍留出余量。6. 往前走手头的模型该往哪个方向改进6.1 从玩具到真实任务预训练和微调才是主流写完微型Transformer后你一定想知道下一步该做什么。如果把Transformer比作发动机那通用的Transformer架构就是一台能转起来的引擎但它还不懂任何具体知识。真实应用中的做法是在超大规模语料上做预训练让引擎学会语言的通用规律然后再在下游任务上微调。这么说吧预训练像是一个人从小到大读了海量的书学会了语法、常识、推理模式微调则是这些知识的基础上专门训练他完成某项具体工作比如翻译法律文书或者写代码。所以你不需要自己从零在任务数据上训练一个大Transformer除非你做的是特别特殊的领域且数据量极大否则更现实的路线是加载开源预训练权重然后在小数据上做微调。目前HuggingFace平台把整个过程变得非常简单只需要几行代码就能加载一个BERT或GPT系列模型。6.2 轻量化TransformerRestormer和Swin为什么能省参数初学者常有一种误解以为Transformer就是又大又慢。实际上通过设计更高效的注意力算法Transformer也能做到非常轻量。Restormer就是一个很好的例子它用在图像复原任务上参数量和显存占用远低于ViT但重建效果很好。它的核心技巧是把多头注意力从全空间维度降到通道维度。普通自注意力在HW个位置上做两两交互复杂度是(HW)^2。Restormer只在每个位置的通道维度上做注意力用跨通道的方式建模信息空间复杂度大幅降低。这个思路和Swin Transformer的窗口注意力有异曲同工之妙都是通过限制注意力的作用范围来压缩计算量。如果你想尝试轻量化改进方向其实有很多减少头数、降低d_model、共享某些层参数、用深度可分离卷积替换前馈网络中的线性层、引入蒸馏得到小模型。关键是知道你的瓶颈在哪里。如果显存占用高优先改注意力计算方式如果模型容量不够优先加层数或加宽度如果推理速度慢优先做剪枝和量化。6.3 从Transformer到任何模态一点关于通用模型的思考Transformer能火的根本原因在于它提供了一个极其通用的信息混合框架。不管输入是文本、图像还是时间序列只要你能把它变成一组向量并加入合适的位置信息Transformer就能对这些向量做全局的动态交互。这也是为什么今天很多多模态模型能实现图像描述生成视频问答这类任务——本质上都是把不同模态的数据切成Token然后丢给同一个Transformer处理。但通用性并不意味着没有代价。Transformer需要的数据量、算力和训练技巧都相当高。它不像CNN那样带着空间先验也不像RNN那样刻意建模时间顺序更多时候是先暴力拟合再从头学出规律。这也是为什么我在文末想提醒每一位刚开始接触Transformer的读者不要被各种眼花缭乱的变体带偏节奏先用手写一个极简的、小到能看清每个模块的Transformer把QKV、多头注意力、位置编码、Encoder和Decoder这些基本功打扎实之后再去看ViT也好、Swin也好、各种预测框架也好都会觉得它们只是在这个骨架上做了不同的加法和换法。在我自己动手实现之前一直觉得Transformer是某种高不可攀的神域。真正写过一遍之后才发现它其实就是一个由矩阵乘法、Softmax和LayerNorm这些基础操作堆叠出来的精妙结构。这份原来不过如此的踏实感比任何概念解释都更能支撑你走得更远。希望这篇初见文也能帮你把第一脚踩稳。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价