ARTICLE DETAIL

资讯详情

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

PyTorch从零实现Transformer:手写多头注意力与编解码器

PyTorch从零实现Transformer:手写多头注意力与编解码器 手撕 Transformer 这件事我一直觉得是深度学习中“高性价比”的动手项目。别管后面多少大模型套壳核心组件仍然是注意力机制 编解码器结构看懂原版 Transformer再去看 GPT、BERT、ViT、T5 都会顺畅很多。这次我们直接从 PyTorch 开始不借助torch.nn.Transformer的高级封装一步步把“输入嵌入 → 位置编码 → 多头注意力 → 前馈网络 → 编码器 → 解码器 → 输出层”全部敲出来。每一段代码都按行拆开解释不会出现“这里用了 nn.MultiheadAttention 就跳过”的情况。本篇文章会给出完整可运行的模型代码、训练循环示例和验证方式你拿到之后可以复制到本地跑通也可以继续改造成自己的实验框架。这篇文章适合这几类读者刚学完 PyTorch 基础想通过一个完整项目把张量操作真正用起来的。读论文时对多头注意力、因果掩码、层归一化等概念似懂非懂想通过代码确认细节的。准备面试或做课程项目需要手撸一个最小可运行的 Transformer。做 NLP 方向研究想从零搭建基线模型而不直接依赖 HuggingFace 的高级接口。硬件方面纯训练一个微型翻译模型用 CPU 也能跑通显存不是硬性要求。如果想跑稍微大一点的语料建议准备一块 4GB 以上显存的 NVIDIA 显卡配置好 CUDA 版 PyTorch。本文会用一个小型中英文平行语料做演示目的是把模型跑起来不追求刷 BLEU 指标。1. 核心能力速览先给一张总览表方便你快速判断这篇文章和代码能不能用在你的场景里能力项说明项目类型PyTorch 从零实现 Transformer 的教程代码不依赖torch.nn.Transformer封装核心功能文本序列到序列建模、多头注意力、位置编码、编码器解码器结构、训练与推理验证运行环境Python 3.8PyTorch 1.13Windows / Linux / macOS 均可CPU 支持支持小语料小模型下 CPU 可完成训练GPU 支持支持建议 4GB 以上显存通过 CUDA 加速是否需要下载大模型不需要完全从随机初始化开始训练演示模型数据集内置迷你平行语料也可以替换成自己的文本数据是否提供 API不涉及 Web API但代码中的模型类可被其他 Python 项目直接导入调用是否支持批量任务训练和推理均支持 batch 形式通过 DataLoader 实现适合场景教学演示、论文复现、基线实验、自定义 Encoder-Decoder 模型开发有一点要提前说明这个实现的目标是“把架构完整写出来并跑通”所以代码在效率上不会刻意追求极致优化。如果你想做大规模训练还是建议基于成熟的训练框架继续改但如果你要做的是理解结构和改造实验这版代码足够清晰。2. 适用场景与使用边界从实践角度讲手敲 Transformer 代码主要有三个使用场景学习与面试准备。面试中常见的“多头注意力计算过程”“为什么需要位置编码”“Decoder 的掩码到底怎么加”这些问题代码里都有明确对应。你不需要背答案把代码跑一遍观察每个张量的形状变化理解会比看书快很多。基线模型搭建。在做一些论文复现、课程设计或者小规模实验时从零实现一个可控的基线模型反而比使用大而全的框架更灵活。你可以随意修改注意力变体、位置编码方式、层数、头数而不被高级 API 的固定参数绑架。自定义结构改造。如果你后续要研究稀疏注意力、线性注意力、相对位置编码等方向从这份代码改起会非常顺手。因为所有组件都拆成了独立类替换某个模块只需要关注输入和输出形状。使用边界同样要清楚这不是一个拿来即用的生产级翻译模型。演示代码训练出的模型能力有限不要直接部署到生产环境。如果使用外部数据集、网上爬取的文本或他人语料务必确认版权和授权。公开论文使用的数据集一般可以用于学术研究但不代表可以任意商用。模型仅做技术学习和实验验证不要用它处理未脱敏的个人隐私数据。训练包含姓名、手机号、地址等信息的语料时需要先做匿名化处理。简单说用它学习、做实验、改结构都很合适但别把它当成一个“开箱即用”的工业级 NLP 工具。3. 环境准备与前置条件先把环境准备好。这个项目的依赖非常少主要就是 Python、PyTorch 和一个文本处理库其他都可以用标准库完成。3.1 安装 Python 和 PyTorch推荐使用 conda 或 venv 创建独立环境避免污染系统 Python。下面以 conda 示例# 创建 Python 3.10 环境 conda create -n transformer-demo python3.10 -y conda activate transformer-demo然后安装 PyTorch。CPU 版本可以直接用官方默认命令pip install torch如果你有 NVIDIA 显卡建议到 PyTorch 官网选择对应 CUDA 版本的安装命令。例如安装带 CUDA 11.8 或 CUDA 12.1 的版本需要按实际环境调整# 以 CUDA 12.1 为例实际命令以 torch 官网为例 pip install torch --index-url https://download.pytorch.org/whl/cu121验证安装是否成功import torch print(torch.__version__) print(torch.cuda.is_available())如果torch.cuda.is_available()返回True说明 GPU 可用。返回False也没关系后面训练部分用 CPU 也能跑通示例。3.2 安装辅助库本文只需要一个轻量的 BPE 或分词工具吗不需要。为了把注意力放在 Transformer 结构本身我们直接用字符级词表。也就是说把每个字符当作一个 token这样不需要额外安装 tokenizers、sentencepiece 之类的库代码也可以完整自洽。如果后续你想换成真正的分词器可以再安装pip install sentencepiece但本文代码不会强制依赖它。3.3 检查磁盘和内存整个项目代码和迷你语料加在一起不到几 MB。训练过程中主要内存消耗来自批次数据和优化器状态8GB 内存的机器完全够用。如果使用 CPU 训练建议把 batch size 调小一些避免一次加载过多数据。4. 数据准备与处理Transformer 是序列到序列模型所以需要一组“源文本 → 目标文本”的平行数据。这里我们构建一个迷你机器翻译数据英语短句 → 中文短句。4.1 构建迷你语料# 小规模平行语料英文 - 中文 pairs [ (hello world, 你好 世界), (i love you, 我 爱 你), (good morning, 早上 好), (how are you, 你 怎么 样), (see you tomorrow, 明天 见), (the cat sits on the mat, 猫 坐 在 垫子 上), (the dog runs fast, 狗 跑 得 很 快), (we study deep learning, 我们 学习 深度 学习), (the weather is nice today, 今天 天气 很 好), (i like to read books, 我 喜欢 读 书), ]这里做了一个简化中文分词直接用空格分开。实际项目中中文需要分词工具但为了专注 Transformer 结构演示数据手动分好即可。4.2 构建词表Word-level 和 Char-level 各有优缺点。这里用空格切分后的 word-level 词表更容易观察张量形状变化。def build_vocab(token_lists): vocab {pad: 0, bos: 1, eos: 2, unk: 3} idx len(vocab) for toks in token_lists: for t in toks: if t not in vocab: vocab[t] idx idx 1 return vocab src_token_lists [s.split() for s, _ in pairs] tgt_token_lists [t.split() for _, t in pairs] src_vocab build_vocab(src_token_lists) tgt_vocab build_vocab(tgt_token_lists) src_itos {i: w for w, i in src_vocab.items()} tgt_itos {i: w for w, i in tgt_vocab.items()} print(源语言词表大小, len(src_vocab)) print(目标语言词表大小, len(tgt_vocab))4.3 转换为 ID 序列需要把文本行转成 ID 序列并加上bos和eos标记。Decoder 的输入以bos开头预测目标以eos结尾。def encode(seq_list, vocab, add_bosFalse, add_eosFalse): encoded [] for toks in seq_list: ids [vocab.get(t, vocab[unk]) for t in toks] if add_bos: ids [vocab[bos]] ids if add_eos: ids ids [vocab[eos]] encoded.append(ids) return encoded src_ids encode(src_token_lists, src_vocab) tgt_ids encode(tgt_token_lists, tgt_vocab, add_bosFalse, add_eosTrue) print(源序列示例, src_ids[0]) print(目标序列示例, tgt_ids[0])4.4 按批次填充一个 batch 内的序列长度必须一致所以需要 pad 到当前 batch 的最大长度。这里使用pad_sequence或手写 collate 函数。import torch from torch.nn.utils.rnn import pad_sequence def collate_batch(batch): src_batch, tgt_batch [], [] for src, tgt in batch: src_batch.append(torch.tensor(src, dtypetorch.long)) tgt_batch.append(torch.tensor(tgt, dtypetorch.long)) src_padded pad_sequence(src_batch, batch_firstTrue, padding_value0) tgt_padded pad_sequence(tgt_batch, batch_firstTrue, padding_value0) return src_padded, tgt_padded dataset list(zip(src_ids, tgt_ids))训练时把dataset传入DataLoader设置collate_fncollate_batch即可。5. Transformer 核心模块代码实现现在进入正题。我们按模块顺序实现整个过程不引入任何现成的 Transformer 层。5.1 输入嵌入与位置编码Token 嵌入用nn.Embedding实现。位置编码采用原版 Transformer 的正余弦函数形式好处是不需要学习且能外推到更长的序列。import math import torch import torch.nn as nn import torch.nn.functional as F class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len512, dropout0.1): super().__init__() self.dropout nn.Dropout(pdropout) pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) 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) if d_model % 2 1: # 当 d_model 为奇数时保证列数不越界 pe[:, 1::2] torch.cos(position * div_term)[:, : d_model // 2] else: 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, seq_len, d_model) x x self.pe[:, : x.size(1)] return self.dropout(x) class TokenEmbedding(nn.Module): def __init__(self, vocab_size, d_model): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.d_model d_model def forward(self, x): return self.embed(x) * math.sqrt(self.d_model)关于位置编码原版 Transformer 使用正余弦函数优点是确定性强不增加额外参数如果你愿意也可以换成可学习位置编码把pe改为nn.Parameter效果在短序列上差别不大。5.2 多头自注意力这是整个模型的核心也是需要重点拆解的部分。先把计算过程拆成五步对输入做线性映射生成 Q、K、V。将 Q、K、V 按头数切分变成多头形状。缩放点积注意力。拼接多个头的结果。输出线性映射。代码实现如下class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_head, dropout0.1): super().__init__() assert d_model % n_head 0 self.d_model d_model self.n_head n_head self.head_dim d_model // n_head 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(pdropout) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 1. 线性映射后分头 Q self.w_q(query) # (batch, seq_len, d_model) K self.w_k(key) V self.w_v(value) # 2. 切分为多头 Q Q.view(batch_size, -1, self.n_head, self.head_dim).transpose(1, 2) K K.view(batch_size, -1, self.n_head, self.head_dim).transpose(1, 2) V V.view(batch_size, -1, self.n_head, self.head_dim).transpose(1, 2) # 3. 缩放点积注意力 scores Q K.transpose(-2, -1) / math.sqrt(self.head_dim) 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) context attn_weights V # (batch, n_head, seq_len, head_dim) # 4. 拼接多头 context context.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 5. 输出映射 output self.out_proj(context) return output这里要注意mask的形状。我们约定mask是(batch, 1, seq_len, seq_len)或(1, 1, seq_len, seq_len)的 0/1 张量其中需要屏蔽的位置为 0。在masked_fill时0 的位置被替换成负无穷softmax 之后权重为 0从而不参与注意力聚合。5.3 前馈网络与层归一化前馈网络就是两个线性层加一个 ReLU 激活。层归一化使用nn.LayerNorm残差连接通过x sublayer(x)实现。class FeedForward(nn.Module): def __init__(self, d_model, d_ff2048, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(pdropout) def forward(self, x): return self.linear2(self.dropout(F.relu(self.linear1(x))))5.4 编码器层编码器层做三件事自注意力、残差和层归一化、前馈网络和残差和层归一化。注意原版使用的是 Post-Norm也就是先做子层计算再加残差最后层归一化。class EncoderLayer(nn.Module): def __init__(self, d_model, n_head, d_ff2048, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_head, 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(pdropout) def forward(self, x, src_maskNone): # 自注意力子层 attn_out self.self_attn(x, x, x, masksrc_mask) x self.norm1(x self.dropout(attn_out)) # 前馈子层 ff_out self.feed_forward(x) x self.norm2(x self.dropout(ff_out)) return x如果你更熟悉 Pre-Norm即先 LayerNorm 再进子层也可以按自己的习惯调整。Pre-Norm 训练通常更稳定但原版论文用的是 Post-Norm这里保持一致。5.5 解码器层解码器层比编码器多一个“交叉注意力”子层同时还引入了因果掩码。class DecoderLayer(nn.Module): def __init__(self, d_model, n_head, d_ff2048, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, n_head, dropout) self.cross_attn MultiHeadAttention(d_model, n_head, dropout) self.feed_forward FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.norm3 nn.LayerNorm(d_model) self.dropout nn.Dropout(pdropout) def forward(self, x, encoder_output, src_maskNone, tgt_maskNone): # 1. 因果自注意力 attn_out self.self_attn(x, x, x, masktgt_mask) x self.norm1(x self.dropout(attn_out)) # 2. 交叉注意力Query 来自解码器Key/Value 来自编码器 cross_out self.cross_attn(x, encoder_output, encoder_output, masksrc_mask) x self.norm2(x self.dropout(cross_out)) # 3. 前馈网络 ff_out self.feed_forward(x) x self.norm3(x self.dropout(ff_out)) return x交叉注意力是 Encoder-Decoder 架构建模两个语言之间关系的关键。它让解码器在生成每个词时都能“查看”编码器输出的整个源序列。5.6 生成因果掩码Decoder 自注意力需要保证当前位置只能看到当前位置之前的信息。实现方法是通过一个上三角矩阵来屏蔽未来位置。def generate_square_subsequent_mask(sz): mask torch.ones(sz, sz) mask torch.triu(mask, diagonal1) return 1 - mask这个函数返回的矩阵中允许关注的位置为 1不允许的位置为 0。在解码器计算注意力分数时这个 mask 会被传入MultiHeadAttention.forward最终把未来位置的分数替换成负无穷。5.7 整合编码器和解码器class Encoder(nn.Module): def __init__(self, vocab_size, d_model, n_head, num_layers, d_ff, max_len, dropout): super().__init__() self.token_embed TokenEmbedding(vocab_size, d_model) self.pos_embed PositionalEncoding(d_model, max_len, dropout) self.layers nn.ModuleList( [EncoderLayer(d_model, n_head, d_ff, dropout) for _ in range(num_layers)] ) self.norm nn.LayerNorm(d_model) def forward(self, src, src_maskNone): x self.token_embed(src) x self.pos_embed(x) for layer in self.layers: x layer(x, src_mask) return self.norm(x) class Decoder(nn.Module): def __init__(self, vocab_size, d_model, n_head, num_layers, d_ff, max_len, dropout): super().__init__() self.token_embed TokenEmbedding(vocab_size, d_model) self.pos_embed PositionalEncoding(d_model, max_len, dropout) self.layers nn.ModuleList( [DecoderLayer(d_model, n_head, d_ff, dropout) for _ in range(num_layers)] ) self.norm nn.LayerNorm(d_model) def forward(self, tgt, encoder_output, src_maskNone, tgt_maskNone): x self.token_embed(tgt) x self.pos_embed(x) for layer in self.layers: x layer(x, encoder_output, src_mask, tgt_mask) return self.norm(x) class Transformer(nn.Module): def __init__(self, src_vocab_size, tgt_vocab_size, d_model512, n_head8, num_layers6, d_ff2048, max_len128, dropout0.1): super().__init__() self.encoder Encoder(src_vocab_size, d_model, n_head, num_layers, d_ff, max_len, dropout) self.decoder Decoder(tgt_vocab_size, d_model, n_head, num_layers, d_ff, max_len, dropout) self.output_proj nn.Linear(d_model, tgt_vocab_size) def forward(self, src, tgt, src_maskNone, tgt_maskNone): encoder_output self.encoder(src, src_mask) decoder_output self.decoder(tgt, encoder_output, src_mask, tgt_mask) logits self.output_proj(decoder_output) return logits到这里一个原版结构的 Transformer 已经完整写出来了。里面的src_vocab_size和tgt_vocab_size根据前面构建的词表传入即可。6. 训练与推理验证模型写完之后最关键的验证方式是把它跑起来训练。这里实现一个简单的训练循环。6.1 训练配置为了快速演示把模型规模调小一点例如d_model128、n_head4、num_layers2、d_ff512。这样 CPU 也能跑。device torch.device(cuda if torch.cuda.is_available() else cpu) model Transformer( src_vocab_sizelen(src_vocab), tgt_vocab_sizelen(tgt_vocab), d_model128, n_head4, num_layers2, d_ff512, max_len128, dropout0.1 ).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3, betas(0.9, 0.98), eps1e-9) criterion nn.CrossEntropyLoss(ignore_index0)6.2 训练循环def train_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss 0.0 for src_batch, tgt_batch in dataloader: src_batch src_batch.to(device) tgt_batch tgt_batch.to(device) # 解码器输入去掉目标序列最后一个 token tgt_input tgt_batch[:, :-1] # 标签去掉目标序列第一个 token因为 bos 不需要预测 tgt_label tgt_batch[:, 1:] tgt_mask generate_square_subsequent_mask(tgt_input.size(1)).to(device) logits model(src_batch, tgt_input, tgt_masktgt_mask) loss criterion( logits.reshape(-1, logits.size(-1)), tgt_label.reshape(-1) ) optimizer.zero_grad() loss.backward() # 梯度裁剪防止梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() total_loss loss.item() return total_loss / max(len(dataloader), 1)为什么解码器输入要去掉最后一个 token而标签要去掉第一个 token因为训练时模型在位置 i 的输入是目标序列的 token i要预测的则是 token i1。bos是第一个输入不需要预测最后一个位置的输入对应输出应该预测eos所以标签保留到末尾。6.3 简单批量训练封装迷你语料规模很小不需要写完整 epoch 循环就可以看出 loss 下降。为了演示 data loader 用法可以这样组合from torch.utils.data import DataLoader dataloader DataLoader( dataset, batch_size2, shuffleTrue, collate_fncollate_batch ) for epoch in range(50): loss train_epoch(model, dataloader, optimizer, criterion, device) if (epoch 1) % 10 0: print(fEpoch {epoch 1}, Loss {loss:.4f})运行这个循环你会看到 loss 逐步下降。由于语料只有 10 条模型很小几十个 epoch 之后 loss 就能降到一个比较低的水平说明前向传播和反向传播逻辑本身没有问题。6.4 贪心解码推理训练好模型后我们还需要验证它能不能生成结果。这里实现一个贪心解码函数每次输出一个 token然后把它拼到解码器输入末尾继续预测下一个 token直到遇到eos或达到最大长度。def greedy_decode(model, src_sentence, src_vocab, tgt_vocab, src_itos, tgt_itos, max_len64): model.eval() src_tokens src_sentence.split() src_ids [src_vocab.get(t, src_vocab[unk]) for t in src_tokens] src_tensor torch.tensor([src_ids], dtypetorch.long).to(device) encoder_output model.encoder(src_tensor) tgt_ids [tgt_vocab[bos]] with torch.no_grad(): for _ in range(max_len): tgt_tensor torch.tensor([tgt_ids], dtypetorch.long).to(device) tgt_mask generate_square_subsequent_mask(len(tgt_ids)).to(device) decoder_output model.decoder(tgt_tensor, encoder_output, tgt_masktgt_mask) logits model.output_proj(decoder_output) # (1, seq_len, tgt_vocab) next_token_logits logits[0, -1, :] # 只取最后一个位置 next_token int(next_token_logits.argmax(dim-1)) if next_token tgt_vocab[eos]: break tgt_ids.append(next_token) output_tokens [tgt_itos[i] for i in tgt_ids[1:]] # 去掉 bos return .join(output_tokens)调用示例print(greedy_decode(model, hello world, src_vocab, tgt_vocab, src_itos, tgt_itos))如果输出和“你好 世界”接近说明模型已经学会从训练集中拟合对应关系。6.5 验证维度这里给出几个验证模型是否正常的判断标准验证维度方法成功标志整体 loss 下降观察每个 epoch 的 loss 值loss 明显下降训练集拟合对训练集中的英文短句做贪心解码输出和中文目标句相似新句泛化将训练集句子做轻微改写如把hello world改成hello my world不一定完全正确但至少不崩溃形状正确打印中间张量形状各模块输出形状符合预期掩码生效打印注意力权重观察是否有未来位置的权重被 mask 的位置权重为 0其中“形状正确”这一点非常建议做。你可以给模型喂一个(batch2, seq_len5)的随机输入然后逐个模块打印输出形状dummy_src torch.randint(0, 20, (2, 5)).to(device) dummy_tgt torch.randint(0, 20, (2, 4)).to(device) mask generate_square_subsequent_mask(4).to(device) logits model(dummy_src, dummy_tgt, tgt_maskmask) print(logits.shape) # 期望 (2, 4, tgt_vocab_size)如果形状不对优先检查多头注意力里view和transpose之后的维度变化。7. 资源占用与性能观察手敲模型时性能观察的重点不是“显存够不够跑大模型”而是“我的实现是否正确、是否高效”。不过仍然值得关注资源占用情况这会影响实验速度和能处理的规模。7.1 CPU 与 GPU 差异CPU 训练小模型、小数据量下速度可以接受几十个 epoch 可能只要几十秒。适合快速验证逻辑正确性。GPU 训练如果数据量增大到几万条或者把d_model提升到 512、层数加到 6GPU 优势会明显体现。device切换非常方便只需要定义device时根据torch.cuda.is_available()自动选择即可前向传播代码不需要修改。7.2 显存占用观察训练时显存主要消耗在四部分模型参数、优化器状态、前向传播中间激活、反向传播梯度。# Linux / Windows 上观察 GPU 显存占用 watch -n 1 nvidia-smi在 N 卡上用torch.cuda.max_memory_allocated()可以打印训练过程中的峰值显存print(f当前显存分配峰值: {torch.cuda.max_memory_allocated() / 1024 ** 2:.2f} MB)不过要再次强调不同d_model、n_head、batch_size、max_len下的显存占用差异很大。这里不写死“多少 G 够用”更稳妥的做法是你自己跑一次观察峰值再调整。7.3 影响性能的关键因素因素影响方式优化手段batch_size越大越占显存/内存GPU 利用率更高显存不足时减小 batch_sized_model模型宽度影响参数量和 QKV 计算量实验阶段从 128 开始num_layers模型深度影响参数量和延迟先 2 层跑通再逐步增加n_head多头数量影响并行注意力维度必须能整除 d_modelmax_len序列长度影响注意力矩阵大小复杂度是 O(L^2)能缩短就缩短d_ff前馈网络中间维度一般采用 4 倍 d_model7.4 如何降低资源占用减小batch_size和max_len这是最直接的手段。使用梯度累积模拟更大 batch每accumulation_steps步做一次optimizer.step()但为了简洁本文未实现。使用自动混合精度AMP减少显存占用和加速训练scaler torch.cuda.amp.GradScaler() with torch.autocast(device_typecuda, dtypetorch.float16): logits model(src_batch, tgt_input, tgt_masktgt_mask) loss criterion(logits.reshape(-1, logits.size(-1)), tgt_label.reshape(-1))AMP 在一些老显卡上可能不稳定建议先跑通普通 FP32 版本再做精度优化。8. 常见问题与排查方法从经验看手写 Transformer 时大家遇到的问题比较集中这里整理成排查表。问题现象可能原因排查方式解决方案安装 PyTorch 后torch.cuda.is_available()为 False安装的是 CPU 版或 CUDA 驱动不匹配运行nvidia-smi查看驱动版本打印torch.version.cuda重新安装对应 CUDA 版本的 PyTorch模型前向传播报维度错误d_model不能被n_head整除或 mask 形状不对打印每个模块输入输出形状调整d_model和n_head关系统一 mask 形状注意力分数出现 NaN学习率过大或softmax输入中有inf减小学习率打印梯度使用学习率预热或梯度裁剪loss 不下降数据量太小、学习率设置不合理、词表对齐错误打印src_vocab和tgt_vocab的映射使用损失下降更好的优化器参数或增大 epoch解码输出全是unk词表未包含对应词或模型未充分训练检查src_vocab.get返回的 ID保证词表和推理时使用的 tokenizer 一致DataLoader 报长度不一致没有使用pad_sequence检查collate_fn是否正确将collate_batch传入DataLoader位置编码形状越界输入序列长度超过max_len打印x.size(1)和pe.size(1)增大max_len或在forward中动态截断梯度爆炸训练不稳定深层网络常见观察 loss 是否突然跳到极大值添加clip_grad_norm_降低学习率训练速度很慢使用 CPU 且d_model较大查看 CPU 占用减小模型尺寸或降低batch_size也可以尝试 GPU生成结果出现重复 token贪心解码的常见问题或模型容量不足打印生成序列改为 beam search或增加模型规模和数据量如果你用的是 Windows 系统训练时遇到多进程 DataLoader 报错通常把num_workers0就行避免 multiprocessing 在交互式环境中的各种坑。此外还有几个容易忽略的地方词表对齐src_vocab和tgt_vocab是两套独立词表不要混用。Decoder 的嵌入矩阵和输出投影矩阵按tgt_vocab_size定义。Padding 掩码本文示例因为所有句子都比较短没有额外传src_mask来处理 padding 位置。实际大规模数据中padding 位置不应该参与注意力计算需要构造文本长度 mask。位置编码溢出pe[:, 0::2]和pe[:, 1::2]的写法在d_model为偶数时没问题但奇数维度要小心。上面代码已经做了兼容但强烈建议d_model保持偶数省很多烦恼。训练集很小迷你语料由于样本太少过拟合非常严重。这不是模型 bug而是数据规模导致的预期现象。要真正训练一个翻译模型请换成较大的公开数据集。9. 最佳实践与使用建议从工程角度看从零实现 Transformer 之后如果想继续做实验有几点建议比较实用9.1 先固定一个小配置跑通不要一上来就复刻原论文d_model512, num_layers6, d_ff2048的大配置。先把模型缩到d_model64或128、层数 2 层、batch size 2确保前向反向、loss 下降、解码输出整个链路跑通再逐步放大。这样排查问题成本最低。9.2 建立一套最小可运行配置建议把以下内容固定下来方便每次启动实验# 训练演示 python train.py --d_model 128 --n_head 4 --num_layers 2 --d_ff 512 --batch_size 4 --epochs 50你可以把网络结构参数和训练超参数独立成配置文件或命令行参数避免每次改代码。9.3 目录结构管理推荐这样组织文件transformer-from-scratch/ ├── data.py # 数据加载与词表构建 ├── model.py # Transformer 模型定义 ├── train.py # 训练脚本 ├── decode.py # 推理脚本 ├── config.py # 超参数配置 └── outputs/ # 模型保存目录模型保存和加载使用 PyTorch 的state_dict# 保存 torch.save(model.state_dict(), outputs/transformer_demo.pt) # 加载 model Transformer(...) model.load_state_dict(torch.load(outputs/transformer_demo.pt, map_locationdevice))9.4 训练过程要记录日志数据规模一旦增大训练日志就很重要。建议记录每个 epoch 的 loss、学习率、GPU 峰值显存、当前模型在验证集上的表现。可以不引入复杂框架先用 Python 自带的logging模块写日志文件。9.5 调试技巧先用torch.autograd.set_detect_anomaly(True)找梯度出现 NaN 的位置但训练正常后要关闭因为它会拖慢速度。打印注意力权重把MultiHeadAttention里的attn_weights返回出来观察模型是否学到了有意义的对齐关系。这是理解 Transformer 内部行为最直观的方式。单步调试时把输入和 mask 都固定为最小形状例如(1, 3)逐层查看张量形状。9.6 合规与授权提醒如果你准备使用公开数据集训练并发布或商用务必确认数据集的许可证。常见数据集的授权条款各不相同有的仅限学术研究有的允许商用。训练语料如果涉及私人对话、人物信息、受版权保护的文本需要先脱敏并获得授权。在技术博客中贴训练结果时也不要展示未经授权的人名、地址等隐私内容。10. 总结与下一步这个项目最值得尝试的一点你完全不用依赖torch.nn.Transformer就能从零搭出一个可训练、可推理的完整 Transformer。位置编码、多头注意力、因果掩码、编码器解码器交互这些概念在代码里都有非常直观的对应。最先应该验证的功能是把训练循环跑通让 loss 降下去然后对训练集中的句子做一个贪心解码确认输入输出链路没有断。最容易踩的坑有三个mask 的形状不对导致注意力分数广播错误。解码器输入和标签的错位处理不对导致模型学不到正确的映射。词表不统一推理时某个词没有对应 ID导致输出全是unk。后续可以直接扩展的方向把贪心解码改成 beam search提升生成质量。加入学习率预热warmup和 Noam 衰减策略训练更稳定。加入 padding mask适配变长文本。把MultiHeadAttention替换成线性注意力、局部注意力等变体做对比实验。把模型接到翻译、摘要、文本复述等具体任务上替换成更大的公开数据集。到这里一个从输入到注意力再到编解码器的完整 Transformer 已经全部敲完。你可以直接复制代码到本地跑也可以把它作为后续研究的基础框架需要调整结构时改动起来会比较顺手。建议收藏备用后面做 Transformer 变体实验时可以直接回来对照。
返回列表