手撕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 输出层源码原版实现importtorchimporttorch.nnasnnimportmath# 最终输出预测层classGenerator(nn.Module):def__init__(self,d_model,vocab_size):super().__init__()# 唯一维度变换d_model映射到词表大小self.projnn.Linear(d_model,vocab_size)defforward(self,x):# 最后一维做softmax概率归一returntorch.softmax(self.proj(x),dim-1)二、核心前置统一掩码生成函数完整补齐前五章我们拆分了两种掩码本章整合完整模型需要统一、标准的掩码生成逻辑适配训练与推理全场景彻底解决掩码使用混乱问题。# 生成解码器因果掩码屏蔽未来位置defsubsequent_mask(size):masktorch.ones(1,size,size)returntorch.tril(mask)# 生成Padding掩码屏蔽无效占位符defcreate_pad_mask(x,pad_idx0):# x: [batch, seq_len] token索引序列return(x!pad_idx).unsqueeze(-2)掩码使用场景终极区分src_maskEncoder专用Padding掩码只屏蔽句子无效填充位保证双向理解不受空白干扰tgt_maskDecoder专用组合掩码Padding掩码因果掩码既屏蔽空白又屏蔽未来Token。三、100%原版完整Transformer模型无魔改、全对齐论文整合前五章所有基础组件本章输出层统一掩码封装完整可落地的原版Transformer代码连贯、结构标准完全对标论文架构无任何自定义修改。importcopyimporttorchimporttorch.nnasnnimportmath# 工具函数克隆网络层defclones(module,N):returnnn.ModuleList([modulefor_inrange(N)])# 层归一化classLayerNorm(nn.Module):def__init__(self,features,eps1e-6):super().__init__()self.gammann.Parameter(torch.ones(features))self.betann.Parameter(torch.zeros(features))self.epsepsdefforward(self,x):meanx.mean(-1,keepdimTrue)stdx.std(-1,keepdimTrue)returnself.gamma*(x-mean)/(stdself.eps)self.beta# 残差归一化子层classSublayerConnection(nn.Module):def__init__(self,size,dropout):super().__init__()self.normLayerNorm(size)self.dropoutnn.Dropout(dropout)defforward(self,x,sublayer):returnself.norm(xself.dropout(sublayer(self.norm(x))))# 多头注意力classMultiHeadedAttention(nn.Module):def__init__(self,h,d_model,dropout0.1):super().__init__()assertd_model%h0self.d_kd_model//h self.hh self.linearsclones(nn.Linear(d_model,d_model),4)self.attnNoneself.dropoutnn.Dropout(pdropout)defattention(self,q,k,v,maskNone,dropoutNone):scorestorch.matmul(q,k.transpose(-2,-1))/math.sqrt(self.d_k)ifmaskisnotNone:scoresscores.masked_fill(mask0,-1e9)p_attntorch.softmax(scores,dim-1)ifdropoutisnotNone:p_attndropout(p_attn)returntorch.matmul(p_attn,v),p_attndefforward(self,query,key,value,maskNone):ifmaskisnotNone:maskmask.unsqueeze(1)nbatchquery.size(0)query,key,value[l(x).view(nbatch,-1,self.h,self.d_k).transpose(1,2)forl,xinzip(self.linears,(query,key,value))]x,self.attnself.attention(query,key,value,mask,self.dropout)xx.transpose(1,2).contiguous().view(nbatch,-1,self.h*self.d_k)returnself.linears[-1](x)# 逐位置前馈网络classPositionwiseFeedForward(nn.Module):def__init__(self,d_model,d_ff,dropout0.1):super().__init__()self.w1nn.Linear(d_model,d_ff)self.w2nn.Linear(d_ff,d_model)self.dropoutnn.Dropout(dropout)defforward(self,x):returnself.w2(self.dropout(torch.relu(self.w1(x))))# 词嵌入位置编码classEmbeddings(nn.Module):def__init__(self,d_model,vocab):super().__init__()self.lutnn.Embedding(vocab,d_model)self.d_modeld_modeldefforward(self,x):returnself.lut(x)*math.sqrt(self.d_model)classPositionalEncoding(nn.Module):def__init__(self,d_model,dropout,max_len5000):super().__init__()self.dropoutnn.Dropout(pdropout)petorch.zeros(max_len,d_model)positiontorch.arange(0,max_len).unsqueeze(1)div_termtorch.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)pepe.unsqueeze(0)self.register_buffer(pe,pe)defforward(self,x):xxself.pe[:,:x.size(1)]returnself.dropout(x)# 单层Encoder 多层EncoderclassEncoderLayer(nn.Module):def__init__(self,size,self_attn,feed_forward,dropout):super().__init__()self.self_attnself_attn self.feed_forwardfeed_forward self.sublayerclones(SublayerConnection(size,dropout),2)self.sizesizedefforward(self,x,mask):xself.sublayer[0](x,lambdax:self.self_attn(x,x,x,mask))returnself.sublayer[1](x,self.feed_forward)classEncoder(nn.Module):def__init__(self,layer,N):super().__init__()self.layersclones(layer,N)self.normLayerNorm(layer.size)defforward(self,x,mask):forlayerinself.layers:xlayer(x,mask)returnself.norm(x)# 单层Decoder 多层DecoderclassDecoderLayer(nn.Module):def__init__(self,size,self_attn,src_attn,feed_forward,dropout):super().__init__()self.sizesize self.self_attnself_attn self.src_attnsrc_attn self.feed_forwardfeed_forward self.sublayerclones(SublayerConnection(size,dropout),3)defforward(self,x,memory,src_mask,tgt_mask):xself.sublayer[0](x,lambdax:self.self_attn(x,x,x,tgt_mask))xself.sublayer[1](x,lambdax:self.src_attn(x,memory,memory,src_mask))returnself.sublayer[2](x,self.feed_forward)classDecoder(nn.Module):def__init__(self,layer,N):super().__init__()self.layersclones(layer,N)self.normLayerNorm(layer.size)defforward(self,x,memory,src_mask,tgt_mask):forlayerinself.layers:xlayer(x,memory,src_mask,tgt_mask)returnself.norm(x)# 最终输出层classGenerator(nn.Module):def__init__(self,d_model,vocab_size):super().__init__()self.projnn.Linear(d_model,vocab_size)defforward(self,x):returntorch.softmax(self.proj(x),dim-1)# 完整Transformer模型 classTransformer(nn.Module):def__init__(self,encoder,decoder,src_embed,tgt_embed,generator):super().__init__()self.encoderencoder self.decoderdecoder self.src_embedsrc_embed self.tgt_embedtgt_embed self.generatorgeneratordefencode(self,src,src_mask):# 编码前向输入序列 → 嵌入位置编码 → 6层Encoderreturnself.encoder(self.src_embed(src),src_mask)defdecode(self,tgt,memory,src_mask,tgt_mask):# 解码前向目标序列 → 嵌入位置编码 → 6层Decoderreturnself.decoder(self.tgt_embed(tgt),memory,src_mask,tgt_mask)defforward(self,src,tgt,src_mask,tgt_mask):# 完整前向链路memoryself.encode(src,src_mask)outself.decode(tgt,memory,src_mask,tgt_mask)returnself.generator(out)# 快速构建原版Transformer论文标准参数defmake_standard_transformer(src_vocab,tgt_vocab,d_model512,d_ff2048,h8,N6,dropout0.1):ccopy.deepcopy attnMultiHeadedAttention(h,d_model)ffPositionwiseFeedForward(d_model,d_ff,dropout)positionPositionalEncoding(d_model,dropout)encoderEncoder(EncoderLayer(d_model,c(attn),c(ff),dropout),N)decoderDecoder(DecoderLayer(d_model,c(attn),c(attn),c(ff),dropout),N)src_embednn.Sequential(Embeddings(d_model,src_vocab),c(position))tgt_embednn.Sequential(Embeddings(d_model,tgt_vocab),c(position))generatorGenerator(d_model,tgt_vocab)returnTransformer(encoder,decoder,src_embed,tgt_embed,generator)四、全链路维度终极闭环从Token ID到概率输出本篇彻底终结所有维度疑惑整理全网最完整、无遗漏的Transformer全链路维度变化表严格遵循论文标准参数d_model512、h8、d_k64。网络环节输出维度 Shape维度变化说明原始输入Token ID[batch, seq_len]纯整数序列无特征维度词嵌入位置编码[batch, seq_len, 512]转为模型标准隐藏维度6层Encoder编码输出[batch, seq_len, 512]维度恒定输出全局语义memory6层Decoder解码输出[batch, seq_len, 512]维度恒定输出最终语义特征Linear投影层[batch, seq_len, vocab_size]全链路唯一维度变换Softmax输出[batch, seq_len, vocab_size]词表概率分布用于预测终极核心结论Transformer的设计哲学极致优雅——特征学习全程保维最终任务统一降维/映射既满足深层堆叠训练需求又适配生成预测任务。五、自回归推理实战模拟真实生成过程训练时模型可以并行输入整句目标序列但真实推理生成必须自回归逐词预测这是大模型生成的核心逻辑我们实现极简可运行推理代码。defautoregressive_infer(model,src,src_mask,max_len,start_idx):# 自回归逐词生成model.eval()withtorch.no_grad():# 编码器只计算一次全局复用memorymodel.encode(src,src_mask)# 初始输入仅有起始标记tgttorch.full((1,1),start_idx,dtypetorch.long)for_inrange(max_len-1):# 动态生成因果掩码tgt_masksubsequent_mask(tgt.size(1))# 解码预测outmodel.decode(tgt,memory,src_mask,tgt_mask)probmodel.generator(out)# 取最后一个位置的最大概率词next_wordtorch.argmax(prob[:,-1,:],dim-1,keepdimTrue)# 拼接序列继续迭代tgttorch.cat([tgt,next_word],dim1)returntgt生成逻辑核心Encoder全局语义只计算一次Decoder反复迭代更新每一步只新增一个单词完美模拟人类逐字写作逻辑。六、原版论文参数1:1对标零魔改验证本系列所有代码、逻辑、参数100%对标Attention Is All You Need原版论文无任何自定义魔改彻底规避网上错误魔改教程编码器、解码器堆叠层数N6模型隐藏维度d_model512FFN中间升维维度d_ff2048多头注意力头数h8单头维度d_k64Dropout概率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学习困境建立完整、严谨、可落地的模型认知为后续大模型微调、预训练、算法落地筑牢最坚实的基础。全文终 · 系列完结