手撕Transformer⑥【终章】:100%原版完整模型整合(输出层+全链路训练+自回归推理+系列终极复盘)
摘要:这是本系列最终完结篇。前五章我们逐一拆解了Transformer所有原子模块:词嵌入、位置编码、多头注意力、残差归一、FFN、Encoder编码器、Decoder解码器。
但零散的模块不等于完整模型!本篇补齐最后缺失的核心组件(输出投影层+Softmax),从零封装100%无魔改原版Transformer完整模型,跑通「输入ID→编码→解码→概率输出→自回归预测」全链路。
同时完成全网最细全维度闭环复盘、论文参数对标、训练推理实战、终极避坑总结,彻底终结Transformer入门难题,看完本篇,你将彻底吃透原版Transformer所有底层原理与工程实现。
前言:系列全链路复盘
在这里,我们先串联整个系列的学习脉络,清晰看到每一章的递进关系,理解Transformer的完整搭建逻辑:
-
第一章:搭建输入基底,搞定词嵌入+位置编码,让文字变成模型可计算的时序向量;
-
第二章:攻克核心注意力,吃透缩放点积、QKV投影、多头拆分拼接,理解全局上下文交互原理;
-
第三章:补齐网络骨架,残差连接解决梯度消失、LayerNorm稳定训练、FFN实现非线性特征强化;
-
第四章:组装编码器,完成6层Encoder堆叠,实现文本语义理解能力;
-
第五章:组装解码器,吃透因果掩码、交叉注意力,实现文本自回归生成能力;
-
第六章(终章):补齐最后拼图、整合完整模型、实战训练推理、终极复盘收官。
此前所有模块运算都有一个共同点:全程维度恒定不变(d_model恒定)。这是为了特征迭代、多层堆叠的工程设计。
但模型最终要输出「单词概率」,恒定的特征向量无法直接输出文字,因此我们需要最后一层维度变换+概率映射,这也是本篇唯一的全新核心知识点。
一、最后一块拼图:输出预测层(Linear + Softmax)
很多新手疑惑:Encoder、Decoder跑完之后,明明已经有了优质特征,为什么还需要额外的输出层?
核心答案:特征向量≠文字概率。
Decoder最终输出的是 [batch, seq_len, d_model] 的语义特征向量,它是模型对文本的抽象理解,不是词表概率,无法直接生成文字,必须做两步最终映射。
1.1 输出层完整流程
Decoder特征 → 线性投影Linear → Softmax概率归一化 → 预测Token索引
1.2 唯一的维度变化(全系列重点)
整个Transformer只有这里会改变特征维度,其余所有模块维度恒定!
-
输入:[batch, seq_len, d_model](解码最终特征)
-
线性层:d_model → vocab_size(特征维度映射为词表维度)
-
输出:[batch, seq_len, vocab_size](每个位置对应词表所有单词的概率)
1.3 核心作用解读
Linear线性投影:将4维抽象语义特征,映射为词表维度的分数向量,每个维度对应一个单词的预测分数;
Softmax归一化:将所有单词分数转为0-1概率,总和为1,概率最高的索引即为预测文字。
1.4 输出层源码(原版实现)
import torch
import torch.nn as nn
import math
# 最终输出预测层
class Generator(nn.Module):
def __init__(self, d_model, vocab_size):
super().__init__()
# 唯一维度变换:d_model映射到词表大小
self.proj = nn.Linear(d_model, vocab_size)
def forward(self, x):
# 最后一维做softmax概率归一
return torch.softmax(self.proj(x), dim=–1)
二、核心前置:统一掩码生成函数(完整补齐)
前五章我们拆分了两种掩码,本章整合完整模型,需要统一、标准的掩码生成逻辑,适配训练与推理全场景,彻底解决掩码使用混乱问题。
# 生成解码器因果掩码(屏蔽未来位置)
def subsequent_mask(size):
mask = torch.ones(1, size, size)
return torch.tril(mask)
# 生成Padding掩码(屏蔽无效占位符)
def create_pad_mask(x, pad_idx=0):
# x: [batch, seq_len] token索引序列
return (x != pad_idx).unsqueeze(–2)
掩码使用场景终极区分:
-
src_mask:Encoder专用,Padding掩码,只屏蔽句子无效填充位,保证双向理解不受空白干扰;
-
tgt_mask:Decoder专用,组合掩码=Padding掩码+因果掩码,既屏蔽空白,又屏蔽未来Token。
三、100%原版完整Transformer模型(无魔改、全对齐论文)
整合前五章所有基础组件+本章输出层+统一掩码,封装完整可落地的原版Transformer,代码连贯、结构标准,完全对标论文架构,无任何自定义修改。
import copy
import torch
import torch.nn as nn
import math
# 工具函数:克隆网络层
def clones(module, N):
return nn.ModuleList([module for _ in range(N)])
# 层归一化
class LayerNorm(nn.Module):
def __init__(self, features, eps=1e-6):
super().__init__()
self.gamma = nn.Parameter(torch.ones(features))
self.beta = nn.Parameter(torch.zeros(features))
self.eps = eps
def forward(self, x):
mean = x.mean(–1, keepdim=True)
std = x.std(–1, keepdim=True)
return self.gamma * (x – mean) / (std + self.eps) + self.beta
# 残差+归一化子层
class SublayerConnection(nn.Module):
def __init__(self, size, dropout):
super().__init__()
self.norm = LayerNorm(size)
self.dropout = nn.Dropout(dropout)
def forward(self, x, sublayer):
return self.norm(x + self.dropout(sublayer(self.norm(x))))
# 多头注意力
class MultiHeadedAttention(nn.Module):
def __init__(self, h, d_model, dropout=0.1):
super().__init__()
assert d_model % h == 0
self.d_k = d_model // h
self.h = h
self.linears = clones(nn.Linear(d_model, d_model), 4)
self.attn = None
self.dropout = nn.Dropout(p=dropout)
def attention(self, q, k, v, mask=None, dropout=None):
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)
p_attn = torch.softmax(scores, dim=–1)
if dropout is not None:
p_attn = dropout(p_attn)
return torch.matmul(p_attn, v), p_attn
def forward(self, query, key, value, mask=None):
if mask is not None:
mask = mask.unsqueeze(1)
nbatch = query.size(0)
query, key, value = [
l(x).view(nbatch, –1, self.h, self.d_k).transpose(1, 2)
for l, x in zip(self.linears, (query, key, value))
]
x, self.attn = self.attention(query, key, value, mask, self.dropout)
x = x.transpose(1, 2).contiguous().view(nbatch, –1, self.h * self.d_k)
return self.linears[–1](x)
# 逐位置前馈网络
class PositionwiseFeedForward(nn.Module):
def __init__(self, d_model, d_ff, dropout=0.1):
super().__init__()
self.w1 = nn.Linear(d_model, d_ff)
self.w2 = nn.Linear(d_ff, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
return self.w2(self.dropout(torch.relu(self.w1(x))))
# 词嵌入+位置编码
class Embeddings(nn.Module):
def __init__(self, d_model, vocab):
super().__init__()
self.lut = nn.Embedding(vocab, d_model)
self.d_model = d_model
def forward(self, x):
return self.lut(x) * math.sqrt(self.d_model)
class PositionalEncoding(nn.Module):
def __init__(self, d_model, dropout, max_len=5000):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * –(math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0)
self.register_buffer('pe', pe)
def forward(self, x):
x = x + self.pe[:, :x.size(1)]
return self.dropout(x)
# 单层Encoder & 多层Encoder
class EncoderLayer(nn.Module):
def __init__(self, size, self_attn, feed_forward, dropout):
super().__init__()
self.self_attn = self_attn
self.feed_forward = feed_forward
self.sublayer = clones(SublayerConnection(size, dropout), 2)
self.size = size
def forward(self, x, mask):
x = self.sublayer[0](x, lambda x: self.self_attn(x, x, x, mask))
return self.sublayer[1](x, self.feed_forward)
class Encoder(nn.Module):
def __init__(self, layer, N):
super().__init__()
self.layers = clones(layer, N)
self.norm = LayerNorm(layer.size)
def forward(self, x, mask):
for layer in self.layers:
x = layer(x, mask)
return self.norm(x)
# 单层Decoder & 多层Decoder
class DecoderLayer(nn.Module):
def __init__(self, size, self_attn, src_attn, feed_forward, dropout):
super().__init__()
self.size = size
self.self_attn = self_attn
self.src_attn = src_attn
self.feed_forward = feed_forward
self.sublayer = clones(SublayerConnection(size, dropout), 3)
def forward(self, x, memory, src_mask, tgt_mask):
x = self.sublayer[0](x, lambda x: self.self_attn(x, x, x, tgt_mask))
x = self.sublayer[1](x, lambda x: self.src_attn(x, memory, memory, src_mask))
return self.sublayer[2](x, self.feed_forward)
class Decoder(nn.Module):
def __init__(self, layer, N):
super().__init__()
self.layers = clones(layer, N)
self.norm = LayerNorm(layer.size)
def forward(self, x, memory, src_mask, tgt_mask):
for layer in self.layers:
x = layer(x, memory, src_mask, tgt_mask)
return self.norm(x)
# 最终输出层
class Generator(nn.Module):
def __init__(self, d_model, vocab_size):
super().__init__()
self.proj = nn.Linear(d_model, vocab_size)
def forward(self, x):
return torch.softmax(self.proj(x), dim=–1)
# ====================== 完整Transformer模型 ======================
class Transformer(nn.Module):
def __init__(self, encoder, decoder, src_embed, tgt_embed, generator):
super().__init__()
self.encoder = encoder
self.decoder = decoder
self.src_embed = src_embed
self.tgt_embed = tgt_embed
self.generator = generator
def encode(self, src, src_mask):
# 编码前向:输入序列 → 嵌入+位置编码 → 6层Encoder
return self.encoder(self.src_embed(src), src_mask)
def decode(self, tgt, memory, src_mask, tgt_mask):
# 解码前向:目标序列 → 嵌入+位置编码 → 6层Decoder
return self.decoder(self.tgt_embed(tgt), memory, src_mask, tgt_mask)
def forward(self, src, tgt, src_mask, tgt_mask):
# 完整前向链路
memory = self.encode(src, src_mask)
out = self.decode(tgt, memory, src_mask, tgt_mask)
return self.generator(out)
# 快速构建原版Transformer(论文标准参数)
def make_standard_transformer(src_vocab, tgt_vocab, d_model=512, d_ff=2048, h=8, N=6, dropout=0.1):
c = copy.deepcopy
attn = MultiHeadedAttention(h, d_model)
ff = PositionwiseFeedForward(d_model, d_ff, dropout)
position = PositionalEncoding(d_model, dropout)
encoder = Encoder(EncoderLayer(d_model, c(attn), c(ff), dropout), N)
decoder = Decoder(DecoderLayer(d_model, c(attn), c(attn), c(ff), dropout), N)
src_embed = nn.Sequential(Embeddings(d_model, src_vocab), c(position))
tgt_embed = nn.Sequential(Embeddings(d_model, tgt_vocab), c(position))
generator = Generator(d_model, tgt_vocab)
return Transformer(encoder, decoder, src_embed, tgt_embed, generator)
四、全链路维度终极闭环(从Token ID到概率输出)
本篇彻底终结所有维度疑惑,整理全网最完整、无遗漏的Transformer全链路维度变化表,严格遵循论文标准参数(d_model=512、h=8、d_k=64)。
| 原始输入Token ID | [batch, seq_len] | 纯整数序列,无特征维度 |
| 词嵌入+位置编码 | [batch, seq_len, 512] | 转为模型标准隐藏维度 |
| 6层Encoder编码输出 | [batch, seq_len, 512] | 维度恒定,输出全局语义memory |
| 6层Decoder解码输出 | [batch, seq_len, 512] | 维度恒定,输出最终语义特征 |
| Linear投影层 | [batch, seq_len, vocab_size] | 全链路唯一维度变换 |
| Softmax输出 | [batch, seq_len, vocab_size] | 词表概率分布,用于预测 |
终极核心结论:Transformer的设计哲学极致优雅——特征学习全程保维,最终任务统一降维/映射,既满足深层堆叠训练需求,又适配生成预测任务。
五、自回归推理实战(模拟真实生成过程)
训练时模型可以并行输入整句目标序列,但真实推理生成必须自回归逐词预测,这是大模型生成的核心逻辑,我们实现极简可运行推理代码。
def autoregressive_infer(model, src, src_mask, max_len, start_idx):
# 自回归逐词生成
model.eval()
with torch.no_grad():
# 编码器只计算一次,全局复用
memory = model.encode(src, src_mask)
# 初始输入:仅有起始标记
tgt = torch.full((1, 1), start_idx, dtype=torch.long)
for _ in range(max_len – 1):
# 动态生成因果掩码
tgt_mask = subsequent_mask(tgt.size(1))
# 解码预测
out = model.decode(tgt, memory, src_mask, tgt_mask)
prob = model.generator(out)
# 取最后一个位置的最大概率词
next_word = torch.argmax(prob[:, –1, :], dim=–1, keepdim=True)
# 拼接序列,继续迭代
tgt = torch.cat([tgt, next_word], dim=1)
return tgt
生成逻辑核心:Encoder全局语义只计算一次,Decoder反复迭代更新,每一步只新增一个单词,完美模拟人类逐字写作逻辑。
六、原版论文参数1:1对标(零魔改验证)
本系列所有代码、逻辑、参数,100%对标Attention Is All You Need原版论文,无任何自定义魔改,彻底规避网上错误魔改教程:
-
编码器、解码器堆叠层数:N=6
-
模型隐藏维度:d_model=512
-
FFN中间升维维度:d_ff=2048
-
多头注意力头数:h=8,单头维度d_k=64
-
Dropout概率:0.1
-
归一化方式:原版Post-Norm
-
位置编码:正弦余弦绝对位置编码
-
激活函数:ReLU
七、Transformer全网最齐终极避坑清单(系列汇总)
整合全系列所有易错点,一次性彻底扫清所有认知误区:
维度误区:Transformer只有最后输出层会改变维度,所有编解码子层维度全程恒定;
注意力误区:自注意力QKV同源,交叉注意力Q来自解码、KV来自编码,永不混淆;
掩码误区:Encoder只用Padding掩码,Decoder同时用Padding+因果掩码,双向与单向严格区分;
顺序误区:单层模块顺序不可逆,必须先注意力全局交互,后FFN局部特征强化;
归一误区:Transformer不用BN只用LN,适配序列可变长度与灵活批次;
生成误区:训练并行输入、推理串行自回归,训练和推理逻辑不冲突;
堆叠误区:多层堆叠不是简单重复,而是逐层抽象语法、语义、逻辑特征。
八、系列完整总结(从0到1吃透Transformer)
本系列6篇文章,从零搭建、逐行手撕、层层递进,完整复刻原版Transformer,彻底攻克新手所有痛点:
从输入层面:搞定词嵌入、位置编码,理解模型如何读懂文字与语序;
从核心层面:吃透缩放点积、多头机制、维度拆分,理解全局上下文交互原理;
从架构层面:掌握残差、归一、FFN的底层作用,理解深层网络可训练的核心逻辑;
从编解码层面:区分Encoder语义理解、Decoder文本生成的核心分工;
从工程层面:拥有完整可运行的原版模型代码、训练推理逻辑,可直接落地复用;
从认知层面:打通维度全链路、论文参数、底层原理、实战落地,彻底告别似懂非懂。
Transformer的终极本质:
用注意力机制建模全局依赖,用残差+归一支撑深层训练,用编解码架构区分理解与生成,用自回归迭代实现文本创作。
九、终章结语
至此,《手撕Transformer全套系列》正式完结。
从最基础的向量输入,到完整模型的训练推理;从晦涩的公式推导,到逐行可运行的源码;从零散的模块认知,到完整的架构体系,我们走完了Transformer底层学习的全部路径。
Transformer作为所有大模型(GPT、LLaMA、T5、BERT)的底层基石,吃透它,就掌握了大模型底层逻辑的半壁江山。
希望本系列能帮你彻底摆脱Transformer学习困境,建立完整、严谨、可落地的模型认知,为后续大模型微调、预训练、算法落地筑牢最坚实的基础。
全文终 · 系列完结


