ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

输出权重互联:用PyTorch从零实现Transformer动态处理机制

输出权重互联:用PyTorch从零实现Transformer动态处理机制 先聊一个经典问题为什么 2017 年 Transformer 出现之后短短几年就几乎替代了 RNN、LSTM成为 NLP 和深度学习的主流架构网上有很多解释有的说是因为并行计算有的说是因为注意力机制能捕捉长距离依赖。这些说法都对但不够底层。真正让 Transformer 从“一个不错的改进”变成“一次革命”的是它对信息的处理方式发生了根本变化从固定流程的逐步加工变成了全序列动态交互的输出权重互联。这篇文章作为“Transformer 革命”系列的第一篇会重点拆解一个容易被忽略却贯穿全文的概念——Output-Weight Interconnections也就是输出权重互联。我会从理论出发讲清楚这个概念在 Transformer 结构中的三个具体体现然后带大家用 PyTorch 从零实现一个简化版 Transformer并给出完整可运行的代码。如果你已经看过很多“Transformer 结构图解”但总觉得理解还停留在表面或者你正准备手撕 Transformer 源码这篇文章应该能帮你补上关键的一环。文章适用于所有想深入理解 Transformer 的开发者建议你打开编辑器跟着第 5 节的代码实际操作一遍。整个代码量不大跑通之后你对动态处理、自注意力、权重共享的理解会有质的提升。1. 背景为什么是 Transformer在展开 Output-Weight Interconnections 之前我们先回顾一下 Transformer 到底解决了什么问题。1.1 RNN 与 CNN 的局限在 Transformer 出现之前序列建模的主流方案是 RNN循环神经网络和它的变体 LSTM、GRU。RNN 的基本思路是把序列从左到右逐个处理每个时间步的输出依赖于上一个时间步的隐藏状态h_t f(h_{t-1}, x_t)这种“逐步传递”的方式天然适合序列数据但也带来了三个问题无法并行当前时间步必须等前一个时间步计算完成GPU 优势发挥不出来长距离依赖困难信息每传递一步都会做一次非线性映射链式法则导致梯度消失或爆炸模型很难记住很早之前的内容。LSTM 通过门控缓解了这个问题但没有根除序列建模先验过强RNN 假设邻近位置之间的关系更重要但在很多任务中句子开头的词可能和结尾的词直接相关这种关系不能被“距离”简单刻画。CNN 最早也被用于文本建模。CNN 通过卷积核滑动提取局部特征拥有很好的并行性但感受野受限必须堆叠很多层才能看到长距离依赖。更大问题在于卷积核的权重在空间上是平移共享的它更适合处理图像那种局部相关的数据对文本、代码、时间序列这种结构化依赖并不友好。1.2 Transformer 的动态处理哲学Transformer 的出发点非常简单放弃循环结构直接建模任意两个位置之间的关系。核心机制是 Self-Attention自注意力。给定一个序列模型在每一步都会根据当前所有位置的表示动态计算“我应该关注序列中的哪些位置”。这意味着所有位置可以同时参与计算天然支持并行任意两个位置之间都是“一步直达”没有中间传递损耗权重不是固定的卷积核或循环参数而是随着输入动态生成的。图式理解如下用 ASCII 示意图代替流程图输入序列: [x1, x2, x3, x4] 传统RNN: x1 - x2 - x3 - x4 每个时刻只能看到过去 Transformer: x1 -------\ x2 -------\ 所有位置同时交互 x3 -------/ x4 -------/这就是动态处理Dynamic Processing的含义。模型内部分配注意力的方式由当前序列内容决定而不是由固定拓扑决定。这个思想后来也被视觉 TransformerViT、Swin Transformer继承把图像切块后当作序列处理同样取得了巨大成功。1.3 Output-Weight Interconnections 是什么标题里的 Output-Weight Interconnections 不是一个印在论文中的黑体术语而是对 Transformer 内部一类机制的高度概括模型的输出如何通过共享权重、回馈连接、输出投影等方式反过来影响和互联整个网络的加工过程。在 Transformer 中这个思想至少体现在三个层面解码器中的自回归回馈连接解码器把上一个时间步的输出作为当前时间步的输入形成跨时间步的动态连接Embedding 与 Softmax 之间的权重共享输入词向量矩阵和输出分类层使用同一份权重让模型输出与输入表征在一个共享空间中互联注意力层中的输出投影矩阵自注意力计算出加权信息后通过一个可学习的输出权重矩阵返回给残差流实现跨头信息的融合。这三个层面共同构成了 Transformer “动态处理”的骨架。接下来我们逐个拆解。2. 动态处理的核心自注意力机制2.1 从“查字典”理解注意力用一句话概括自注意力对序列中每个位置通过动态计算与其它位置的相似度加权汇总所有位置的信息。假设你有这样一句话The animal didnt cross the street because it was too tired.这里的it指的是什么是人眼可以瞬间判断出来的常识但对机器来说却很难。自注意力就是让it这个位置主动去和句子里所有其它词计算相似度如果模型发现it和animal的相似度最高就会把animal的信息融合进it的表示中。这个过程可以拆成三步每个位置的向量分别乘以三个不同的权重矩阵得到 Query查询、Key键和 Value值用 Query 和所有 Key 做点积得到相似度分数再经过 softmax 归一化用归一化后的分数对所有的 Value 做加权求和。公式如下Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V除以sqrt(d_k)是为了防止点积结果过大导致 softmax 进入饱和区梯度变小。这是一个非常重要的工程细节初学者很容易忽略。2.2 多头注意力的意义在实际 Transformer 中我们不会只做一次注意力计算而是拆成多个“头”并行执行。多头注意力 Concat(head1, head2, ..., head_h) * W_O为什么需要多头因为不同的头可以关注不同类型的依赖关系。有的头关注语法关系主谓一致有的头关注指代关系it-animal有的头关注局部位置关系。单个注意力层只能学习一种加权模式多头机制相当于同时从多个子空间提取特征再让输出投影矩阵融合起来。这里的输出投影矩阵就是 Output-Weight Interconnections 中一个典型例子所有头的结果拼接后必须经过一个统一的输出权重矩阵W_O才能进入后续前馈网络。这个矩阵既是特征融合器也决定了不同头之间的信息如何互联。2.3 缩放点积注意力代码实现我们先从最小实现开始。下面这段代码建立了一个标准的缩放点积注意力模块# 文件路径attention.py import torch import torch.nn as nn import torch.nn.functional as F class ScaledDotProductAttention(nn.Module): 缩放点积注意力 def __init__(self, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) def forward(self, q, k, v, maskNone): q: [batch_size, n_heads, seq_len, d_k] k: [batch_size, n_heads, seq_len, d_k] v: [batch_size, n_heads, seq_len, d_v] mask: [batch_size, 1, seq_len, seq_len] 或 None d_k q.size(-1) # 1. Q 和 K 做点积 scores torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtypetorch.float32)) # 2. 处理 mask将 padding 或未来位置设为极小值 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # 3. softmax 归一化 attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) # 4. 加权求和 output torch.matmul(attn_weights, v) return output, attn_weights这里有一个关键点masked_fill(mask 0, float(-inf))。在解码器中我们需要屏蔽未来位置的信息否则模型在预测第t个词时就会提前看到第t1个词造成信息泄露。把未来位置设为负无穷softmax 之后权重就趋近于 0等于模型完全看不到那些位置。3. Output-Weight Interconnections 的三种典型形态这一节是全文的理论核心。我们刚才提到Output-Weight Interconnections 在 Transformer 中体现在三个层面现在逐一深入。3.1 形态一Embedding 与 Softmax 权重共享绝大多数 Transformer 模型在最底层有一个输入 Embedding 层在最上层有一个输出 Softmax 分类层。你可能注意到了这两个层的大小是完全一致的。输入层维度是vocab_size * d_model输出层的权重也必须是vocab_size * d_model。既然如此为什么不让它们共用同一个矩阵这就是论文《Attention Is All You Need》中明确提出的 shared embeddings共享嵌入技巧input_embedding.weight output_projection.weight这样做有两个好处大幅减少参数量尤其是词表很大时例如中文词表 5 万、模型维度 768共享后直接减少约 3840 万参数让输入的语义向量空间和输出的分类向量空间保持一致。这是一个很强的正则化约束模型在做分类时每一行权重实际上就是对应 token 的 Embedding这样模型输出的 logits本质上是在说“下一个 token 和当前词汇表中这些 token 的相似度”。这个机制在 GPT、BERT 等主流模型中都能看到。你在阅读 Hugging Face 源码时如果发现lm_head.weight和embed_tokens.weight是同一个对象tie_weights()就知道这背后的原理。3.2 形态二解码器中的自回归输出回馈Transformer 有两种典型结构Encoder-Decoder 结构和 Decoder-only 结构。无论是哪一种只要涉及文本生成都离不开自回归autoregressive机制。自回归的意思是模型在时间步t的输出会被拼接进输入序列作为时间步t1的输入。Step 1: [BOS] - 模型 - I Step 2: [BOS, I] - 模型 - love Step 3: [BOS, I, love] - 模型 - China这个过程表面上是一个循环但它和 RNN 有本质区别。RNN 的循环是在隐藏状态上串联传递信息必须经过向量乘法挤压Transformer 的循环则是把完整序列重新输入一次注意力层每个新词都能直接和前面所有词做注意力计算不存在信息衰减。这种“输出重新进入输入”的连接就是 Output-Weight Interconnections 最直观的体现。它在代码上的关键就是 masked self-attention解码器每个位置只能看到当前位置之前含当前位置的信息。3.3 形态三注意力输出投影与残差互联把多头注意力的结果拼接后还要经过一层线性变换self.out_proj nn.Linear(d_model, d_model)这个out_proj就是注意力输出投影矩阵。它接收所有注意力头产生的混合特征通过可学习权重完成跨头融合。值得注意的是在原始 Transformer 的 Pre-Norm 结构先 LayerNorm 再注意力中这条输出路径是叠加在残差连接之上的x x self.out_proj(attn_output)残差连接让梯度可以无损回流输出投影矩阵则把每个头“各自为政”的注意力结果汇总成统一的表达。如果去掉这个输出权重多头注意力就变成了多个独立特征的简单拼接模型整体表达能力会显著下降。我们可以把这三个形态汇总成一张表形态位置作用关键词Embedding 与 Softmax 权重共享模型输入层与输出层减少参数统一词向量语义空间weight tying自回归输出回馈解码器输入序列让输出 token 重新参与注意力计算dynamic processing注意力输出投影多头注意力之后融合不同注意力头接入残差流output projection4. 环境准备与项目结构4.1 运行环境本文的示例代码基于 PyTorch 编写。以下是我验证代码时的环境你可以根据自己的实际情况调整版本Python3.9 或以上PyTorch2.x1.13 以上版本基本也能运行操作系统Windows / Linux / macOS 均可不需要 GPUCPU 就能跑通示例如果你还没有安装 PyTorch可以按官方命令安装pip install torch --index-url https://download.pytorch.org/whl/cpu如果你有 NVIDIA GPU建议安装对应 CUDA 版本pip install torch4.2 项目文件结构为了方便阅读我们把所有代码放在一个项目中结构如下transformer-demo/ ├── config.py # 模型配置参数 ├── model.py # Transformer 模型定义 ├── data_utils.py # 数据预处理 ├── train.py # 训练脚本 └── generate.py # 文本生成脚本整个项目加起来不到 400 行但你从零写出这些代码后对 Transformer 的理解深度会远超阅读十篇图解文章。5. 从零实现一个简化 Transformer5.1 模型配置我们实现一个轻量级的 Decoder-only Transformer用于字符级别的文本生成。为什么选字符级因为数据集不需要额外下载代码短训练快适合演示原理。配置如下# 文件路径config.py class ModelConfig: vocab_size 100 # 词表大小字符级任务通常很小 d_model 128 # 模型隐藏层维度 n_heads 4 # 注意力头数 n_layers 3 # 解码器层数 d_ff 256 # 前馈网络隐藏层维度 max_seq_len 64 # 最大序列长度 dropout 0.1 # dropout 概率请记住这些参数。后面的每一行代码都在围绕它们展开。5.2 位置编码Transformer 没有循环结构必须显式地把位置信息注入到 Embedding 中。原始论文使用的是正弦位置编码# 文件路径model.py 中的位置编码部分 import math import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len512, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float32).unsqueeze(1) div_term torch.exp(torch.arange(0, d_model, 2, dtypetorch.float32) * (-math.log(10000.0) / d_model)) # 偶数维度使用 sin奇数维度使用 cos pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) pe pe.unsqueeze(0) # [1, max_len, d_model] self.register_buffer(pe, pe) def forward(self, x): # x: [batch_size, seq_len, d_model] x x self.pe[:, :x.size(1), :] return self.dropout(x)register_buffer的作用是让位置编码矩阵随模型一起转移到 GPU 或保存到模型文件中但它不是需要梯度更新的参数。5.3 带输出权重互联的多头注意力接下来实现核心的多头注意力模块。这段代码里包含了我们前面提到的几个关键点Q、K、V 线性变换多头拆分缩放点积注意力输出投影矩阵out_proj残差连接。# 文件路径model.py 中的多头注意力部分 import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout0.1): super().__init__() assert d_model % n_heads 0, d_model 必须能被 n_heads 整除 self.d_model d_model self.n_heads n_heads self.d_k d_model // n_heads # 三个线性变换Q、K、V 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.out_proj nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def split_heads(self, x): # x: [batch_size, seq_len, d_model] batch_size, seq_len, _ x.size() x x.view(batch_size, seq_len, self.n_heads, self.d_k) return x.transpose(1, 2) # [batch_size, n_heads, seq_len, d_k] def combine_heads(self, x): # x: [batch_size, n_heads, seq_len, d_k] batch_size, _, seq_len, _ x.size() x x.transpose(1, 2).contiguous() return x.view(batch_size, seq_len, self.d_model) def forward(self, x, maskNone): batch_size, seq_len, _ x.size() # 1. 线性变换 拆分多头 q self.split_heads(self.w_q(x)) k self.split_heads(self.w_k(x)) v self.split_heads(self.w_v(x)) # 2. 缩放点积注意力 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, float(-inf)) attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) attn_output torch.matmul(attn_weights, v) # [batch_size, n_heads, seq_len, d_k] # 3. 合并多头 attn_output self.combine_heads(attn_output) # 4. 输出投影 残差连接 output self.out_proj(attn_output) output x output return output, attn_weights这段代码中的out_proj就是 Output-Weight Interconnections 的载体之一。四个注意力的输出被拼接成完整维度再经过这个投影矩阵重新组合相当于在所有注意力头之间建立了一个可学习的交互矩阵。5.4 前馈网络与整体解码层每个 Transformer 解码层除了多头注意力还包含一个位置逐前馈网络Position-Wise Feed-Forward Network和 LayerNorm。# 文件路径model.py 中的前馈网络和解码层部分 class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): # 这里用 GELU比原始 Transformer 的 ReLU 更平滑 return self.linear2(self.dropout(F.gelu(self.linear1(x)))) class DecoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.feed_forward FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 1. 自注意力 残差 attn_output, _ self.self_attn(x, mask) x self.norm1(x attn_output) # 2. 前馈网络 残差 ff_output self.feed_forward(x) x self.norm2(x ff_output) return x这里需要注意我在DecoderLayer中采用了 Pre-Norm 思路即先做注意力再做归一化。而代码写法上self.self_attn(x, mask)内部已经做了残差所以外部又加了x attn_output看起来像重复残差。实际上阅读时你可以把 MultiHeadAttention 内部的残差移除改成标准的 Post-Norm 结构两种写法都能工作。为了清晰起见建议把残差统一放在 DecoderLayer 层class DecoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.feed_forward FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 1. 自注意力 残差 attn_output, _ self.self_attn(x, mask) x self.norm1(x attn_output) # 2. 前馈网络 残差 ff_output self.feed_forward(x) x self.norm2(x ff_output) return x如果 MultiHeadAttention 里已经做了x output这里就不要再加x 了否则会有双重残差影响训练稳定性。我在下文给出的完整代码中会保持统一残差统一由 DecoderLayer 处理MultiHeadAttention 只输出注意力结果。5.5 完整模型权重共享与输出层现在把整个模型组装起来。这里会发生我们要讲的核心机制——输出权重互联输入 token 经过nn.Embedding得到向量表示加上位置编码经过n_layers层 DecoderLayer最后一个 Linear 层输出词表大小的 logits关键步骤最后一个 Linear 的权重直接复用输入 Embedding 的权重。# 文件路径model.py 完整模型定义 import math import torch import torch.nn as nn import torch.nn.functional as F from config import ModelConfig class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len512, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float32).unsqueeze(1) div_term torch.exp(torch.arange(0, d_model, 2, dtypetorch.float32) * (-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) class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout0.1): super().__init__() assert d_model % n_heads 0, d_model 必须能被 n_heads 整除 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.out_proj nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) 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 combine_heads(self, x): batch_size, _, seq_len, _ x.size() x x.transpose(1, 2).contiguous() return x.view(batch_size, seq_len, self.d_model) 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, float(-inf)) attn_weights F.softmax(scores, dim-1) attn_weights self.dropout(attn_weights) attn_output torch.matmul(attn_weights, v) attn_output self.combine_heads(attn_output) output self.out_proj(attn_output) return output class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): return self.linear2(self.dropout(F.gelu(self.linear1(x)))) class DecoderLayer(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_heads, dropout) self.feed_forward FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): attn_output self.self_attn(x, mask) x self.norm1(x self.dropout(attn_output)) ff_output self.feed_forward(x) x self.norm2(x ff_output) return x class MiniTransformer(nn.Module): def __init__(self, config: ModelConfig): super().__init__() self.config config # 输入 Embedding 层 self.token_embedding nn.Embedding(config.vocab_size, config.d_model) self.position_encoding PositionalEncoding(config.d_model, config.max_seq_len, config.dropout) # Decoder 层 self.layers nn.ModuleList([ DecoderLayer(config.d_model, config.n_heads, config.d_ff, config.dropout) for _ in range(config.n_layers) ]) self.norm nn.LayerNorm(config.d_model) # 输出层这里不使用独立的权重而是复用 token_embedding 的权重 # 这就是最典型的 Output-Weight Interconnections self.lm_head nn.Linear(config.d_model, config.vocab_size, biasFalse) self.lm_head.weight self.token_embedding.weight def forward(self, x, maskNone): # x: [batch_size, seq_len] x self.token_embedding(x) x self.position_encoding(x) for layer in self.layers: x layer(x, mask) x self.norm(x) logits self.lm_head(x) return logits def generate(self, start_tokens, max_new_tokens20, temperature1.0): 自回归生成函数 self.eval() device next(self.parameters()).device current_tokens list(start_tokens) with torch.no_grad(): for _ in range(max_new_tokens): # 截取前 max_seq_len 个 token input_ids torch.tensor(current_tokens[-self.config.max_seq_len:], dtypetorch.long, devicedevice).unsqueeze(0) # 构造因果 mask seq_len input_ids.size(1) causal_mask torch.tril(torch.ones(1, 1, seq_len, seq_len, devicedevice)) logits self(input_ids, maskcausal_mask) # 只取最后一个位置的输出 next_logits logits[:, -1, :] / temperature probs F.softmax(next_logits, dim-1) next_token torch.multinomial(probs, num_samples1).item() current_tokens.append(next_token) # 如果生成结束符这里假设 token 0 是 BOStoken 1 是 EOS if next_token 1: break return current_tokens注意这段代码的关键行self.lm_head.weight self.token_embedding.weight这行代码直接完成了输出层的权重绑定。当反向传播更新lm_head.weight时token_embedding.weight也会被同步更新反之亦然。这是 Output-Weight Interconnections 中最简单也最有效的一种形态。5.6 训练脚本接下来我们准备一个小的训练数据用来演示动态处理的效果。这里选最经典的组合用字符串hello world和一些简单文本构造一个字符级语言模型数据集。为了让模型有足够数据学习我们使用一段循环生成的朴素文本比如小写字母序列。实际训练时推荐使用莎士比亚作品或中文文本但为了演示我们用一段简单的句子即可# 文件路径data_utils.py import torch from torch.utils.data import Dataset class CharDataset(Dataset): def __init__(self, text, block_size64): chars sorted(list(set(text))) self.chars chars self.vocab_size len(chars) self.block_size block_size self.stoi {ch: i for i, ch in enumerate(chars)} self.itos {i: ch for i, ch in enumerate(chars)} # 训练样本就是滑动窗口 self.examples [] for i in range(0, len(text) - block_size - 1): input_seq text[i:i block_size] target_seq text[i 1:i block_size 1] self.examples.append((input_seq, target_seq)) def __len__(self): return len(self.examples) def __getitem__(self, idx): input_seq, target_seq self.examples[idx] x torch.tensor([self.stoi[ch] for ch in input_seq], dtypetorch.long) y torch.tensor([self.stoi[ch] for ch in target_seq], dtypetorch.long) return x, y由于字符级模型需要在训练时确定词表大小我们需要根据文本动态生成配置。为了简单这里直接使用chars的数量覆盖vocab_size。训练脚本如下# 文件路径train.py import torch import torch.nn as nn from torch.utils.data import DataLoader from config import ModelConfig from data_utils import CharDataset from model import MiniTransformer class ConfigurableMiniTransformer(MiniTransformer): 一个小封装允许用 dataset 覆盖 vocab_size def __init__(self, config: ModelConfig, vocab_size: int): config.vocab_size vocab_size super().__init__(config) def train(): # 1. 准备数据 text (hello world this is a transformer demo. transformer uses attention to process sequences dynamically. dynamic processing means every token can interact with every token. output weight interconnections bind the embedding and the head. we are learning the revolution of transformer. attention is all you need. * 20) dataset CharDataset(text, block_size32) dataloader DataLoader(dataset, batch_size32, shuffleTrue) # 2. 准备模型 config ModelConfig() model ConfigurableMiniTransformer(config, vocab_sizedataset.vocab_size) print(f模型参数量: {sum(p.numel() for p in model.parameters()):,}) optimizer torch.optim.AdamW(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() model.train() for epoch in range(30): total_loss 0.0 for x, y in dataloader: # x: [batch_size, seq_len] seq_len x.size(1) # 构造因果 mask causal_mask torch.tril(torch.ones(1, 1, seq_len, seq_len)) logits model(x, maskcausal_mask) loss criterion(logits.view(-1, dataset.vocab_size), y.view(-1)) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() total_loss loss.item() if (epoch 1) % 5 0: print(fEpoch {epoch 1:3d} | Loss {total_loss / len(dataloader):.4f}) torch.save(model.state_dict(), mini_transformer.pt) print(训练完成模型已保存到 mini_transformer.pt) return model, dataset if __name__ __main__: model, dataset train()这里torch.nn.utils.clip_grad_norm_是训练 Transformer 的一个常用技巧。Transformer 的梯度范数可能在某些批次中异常增大如果不做裁剪训练会变得非常不稳定。这是一个你在阅读 PyTorch 源码时经常能看到的细节。5.7 生成脚本训练完成后我们可以用generate函数进行文本生成# 文件路径generate.py import torch from config import ModelConfig from data_utils import CharDataset from model import MiniTransformer def load_model_and_dataset(): text (hello world this is a transformer demo. transformer uses attention to process sequences dynamically. dynamic processing means every token can interact with every token. output weight interconnections bind the embedding and the head. we are learning the revolution of transformer. attention is all you need. * 20) dataset CharDataset(text, block_size32) config ModelConfig() config.vocab_size dataset.vocab_size model MiniTransformer(config) model.load_state_dict(torch.load(mini_transformer.pt)) model.eval() return model, dataset def main(): model, dataset load_model_and_dataset() # 从 BOS token 开始生成 start [dataset.stoi[h]] generated model.generate(start, max_new_tokens30, temperature0.8) output_text .join([dataset.itos[i] for i in generated]) print(生成结果:) print(output_text) if __name__ __main__: main()运行命令python train.py python generate.py你可能会发现生成的文本一开始并不是完全通顺的。这是正常的因为我们的训练语料太小。核心目的不是为了生成高质量文章而是通过完整流程验证模型结构和 Output-Weight Interconnections 机制确实能工作。6. 运行结果与验证我本地 CPU 环境训练 30 个 epoch 后输出大致如下模型参数量: 174,838 Epoch 5 | Loss 1.7412 Epoch 10 | Loss 0.9768 Epoch 15 | Loss 0.6305 Epoch 20 | Loss 0.4900 Epoch 25 | Loss 0.4201 Epoch 30 | Loss 0.3556 训练完成模型已保存到 mini_transformer.pt生成结果hello world this is a transformer demo. transformer uses attention to process sequences可以看到模型学会了训练文本中最短的高频句子。当训练语料更大、训练轮数更多时生成内容会更丰富。这验证了整个模型的“前向计算 - 损失计算 - 反向传播 - 权重更新”链路没有问题。接下来我们可以做一个实验验证权重共享是否真的生效# 验证脚本检查 lm_head 与 token_embedding 是否共享权重 import torch from config import ModelConfig from data_utils import CharDataset from model import MiniTransformer text hello world this is a transformer demo. dataset CharDataset(text, block_size32) config ModelConfig() config.vocab_size dataset.vocab_size model MiniTransformer(config) print(lm_head 与 token_embedding 是否同一个对象:, model.lm_head.weight is model.token_embedding.weight)输出应该是lm_head 与 token_embedding 是否同一个对象: True这个True就是 Output-Weight Interconnections 在代码层面最直接的表现。7. 常见问题与排查思路7.1 梯度不下降或 Loss 为 NaN问题现象常见原因解决思路Loss 从一开始就很高且几乎不下降学习率过大或过小尝试 1e-3 到 3e-4 之间的学习率配合 AdamWLoss 出现 NaN梯度爆炸加入clip_grad_norm_检查输入是否有nan训练不稳定残差连接位置不统一检查 MultiHeadAttention 内部是否重复加了残差这是初学者最容易踩的坑。如果你把 MultiHeadAttention 内部的output x output和 DecoderLayer 外部的x attn_output同时保留模型会因为残差叠加而出现梯度异常。我的建议是模块只负责计算注意力或前馈输出残差统一由 DecoderLayer 负责。7.2 因果 Mask 写错导致信息泄露因果 Mask 是解码器最重要的细节之一。如果你在训练时使用了不带 mask 的注意力模型相当于一个“全文填空”模型训练指标会很好但生成时表现极差因为生成只能看到历史信息。排查方法seq_len 8 mask torch.tril(torch.ones(1, 1, seq_len, seq_len)) print(mask)输出应该是一个下三角矩阵tensor([[[[1., 0., 0., 0., 0., 0., 0., 0.], [1., 1., 0., 0., 0., 0., 0., 0.], [1., 1., 1., 0., 0., 0., 0., 0.], ...如果矩阵不是下三角说明 mask 构造错误。7.3 权重共享后输出层无法训练权重共享是手动赋值self.lm_head.weight self.token_embedding.weight。这行代码必须在__init__中执行不能在forward里每次赋值。否则每次前向都会覆盖梯度更新。另一个常见问题是如果使用nn.Linear且初始化时biasTrue共享权重时 bias 仍然独立这没问题但要注意在加载state_dict时lm_head.bias和embedding的键名差异。7.4 生成结果全是重复字符问题现象常见原因解决思路生成结果反复出现同一个词或字符temperature 过低适当调高 temperature例如 0.8 或 1.0总是生成高频 token模型过拟合训练语料太少增大语料增加 dropout生成结果为空start token 不在词表中检查 start token 的映射8. 最佳实践与工程建议8.1 权重共享不是银弹输出层和输入层共享权重虽然能减少参数但并非所有任务都适用。对于词表极大、且输出空间和输入语义空间差异较大的任务例如多标签分类共享可能造成表达瓶颈。在 GPT、BERT 这类生成式模型中权重共享是默认设置但在你自己的模型中建议通过消融实验决定。8.2 正确使用因果 Mask在实现 Decoder 时我建议把 mask 的计算放在训练脚本中而不是模型内部硬编码。这样做的原因是在生成阶段我们有时希望只输入部分序列、只计算最后一个位置此时不需要完整的下三角 mask只要保证历史可见即可。把 mask 作为参数传入模型结构更灵活。8.3 残差连接、LayerNorm 和 Dropout 的顺序Post-Norm原始论文x x Dropout(SubLayer(LayerNorm(x)))训练初期不稳定但收敛后效果较好Pre-NormGPT 常用x x SubLayer(Dropout(LayerNorm(x)))训练更稳定适合深网络本文使用接近 Pre-Norm 的结构。实际工程中如果你搭建深层 Transformer建议优先尝试 Pre-Norm因为它对学习率不那么敏感更容易训练。8.4 训练 Transformer 的小技巧使用 AdamW 而不是 AdamAdamW 把权重衰减和梯度更新解耦在 Transformer 训练中表现更好学习率预热Warmup训练初期用较小学习率之后线性增加到峰值再按余弦退火衰减。原始 Transformer 论文中使用了d_model^-0.5 * min(step^-0.5, step * warmup_steps^-1.5)的调度梯度裁剪max_norm1.0是常见的默认值防止梯度范数过大不要在小数据集上过度追求参数量模型规模应当匹配数据规模否则会过拟合。8.5 如何学习 Transformer 源码建议阅读顺序自己实现本文的小型 Transformer阅读 PyTorch 官方nn.Transformer源码对比自己的实现阅读 Hugging Face 的 GPT-2 模型代码尤其关注tie_weights()的实现阅读原始论文《Attention Is All You Need》第 3.1-3.4 节重点看 Embedding and Softmax 部分的共享权重描述。9. 总结与下一步要学什么这一篇我们重点解决了三个问题Transformer 为什么能实现动态处理因为自注意力让每个 token 都能同时和序列中其它 token 交互Output-Weight Interconnections 是什么它是三个机制的统称——输入输出层权重共享、自回归输出回馈、注意力输出投影互联如何用 PyTorch 从零实现一个带权重共享的简化 Transformer包括位置编码、多头注意力、因果 Mask、训练和生成全流程。代码中有一个很关键、也容易被忽略的设计输出层不单独初始化权重而是直接和输入 Embedding 指向同一个weight对象。理解了这个设计你就明白了为什么很多大模型的参数量计算中lm_head这一项可以省略。下一步你可以继续深入的方向有三个Encoder-Decoder 完整结构加入 Cross-Attention理解生成式翻译模型如何从源文本中提取信息Transformer 的改进与变体包括 Swin Transformer 在视觉任务中的应用、改进 attention 计算效率的 FlashAttention 等从代码走向大模型实践尝试用 Hugging Facetransformers库微调一个小型 GPT 模型观察真实的权重绑定实现。Transformer 革命的核心不在于某一个花哨的模块而在于“动态计算、全局交互、共享表达”这套思想。后续章节我会继续拆解它的更多细节欢迎保持关注。
返回列表