多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

训练循环全解析:attention-is-all-you-need-pytorch 中 train.py 从数据加载到 TensorBoard 监控

训练循环全解析:attention-is-all-you-need-pytorch 中 train.py 从数据加载到 TensorBoard 监控 训练循环全解析attention-is-all-you-need-pytorch 中 train.py 从数据加载到 TensorBoard 监控【免费下载链接】attention-is-all-you-need-pytorchA PyTorch implementation of the Transformer model in Attention is All You Need.项目地址: https://gitcode.com/gh_mirrors/at/attention-is-all-you-need-pytorch如果你正在学习 Transformerattention-is-all-you-need-pytorch 这个项目值得逐行精读——它用 PyTorch 完整实现了论文《Attention is All You Need》中的模型而 train.py 正是整个训练流程的总控台从加载 BPE 语料、构建数据迭代器到前向传播、反向传播、学习率 warmup再到 TensorBoard 曲线监控一条 367 行的脚本全部打通。本文将带你按数据流顺序拆解这个训练循环帮你快速看懂 PyTorch 训练一个 Transformer 的完整套路。️ 训练流程一览main() 的 4 个关键动作整个入口在 main()逻辑非常清晰解析参数train.py#L209-L240batch size、d_model、warmup 步数、是否开启 label smoothing 等加载数据train.py#L272-L277根据传入的是 BPE 文件还是预处理好的 pkl走不同的数据加载分支构建模型与优化器train.py#L281-L300实例化Transformer并包装带学习率调度的优化器启动训练循环train.py#L302调用train()进入逐 epoch 迭代 一个小细节如果指定了-seed脚本会固定torch、numpy、random的随机种子并关闭 cudnn benchmarktrain.py#L248-L253保证实验可复现。 数据加载prepare_dataloaders 如何喂数据训练数据有两条加载路径pkl 路径→ prepare_dataloaders()读取preprocess.py预处理好的 pickle 文件含词表、训练集、验证集直接用Dataset构建BPE 文件路径→ prepare_dataloaders_from_bpe_files()用TranslationDataset从.src/.trg编码文件中加载两条路径最后都做了同一件事创建BucketIterator分桶迭代器train.py#L360-L361。BucketIterator(train, batch_sizebatch_size, devicedevice, trainTrue)为什么叫分桶它会把长度相近的句子分到同一个 batch减少 padding 浪费这是机器翻译训练的经典技巧。同时脚本还会从 pickle 中回填一批关键配置词表大小、PAD 索引定义在 transformer/Constants.py 的blank、最大序列长度——这些正是后面实例化模型所需的参数。⚙️ 模型与优化器两个隐藏彩蛋模型实例化在 train.py#L281-L296Transformer类定义在 transformer/Models.py支持两个经典技巧参数作用-embs_share_weight源/目标语言共享同一个词嵌入矩阵-proj_share_weight词嵌入与输出层线性投影共享权重论文 3.4 节做法优化器部分train.py#L298-L300用Adam(betas(0.9, 0.98))并包了一层 ScheduledOptim。它的核心是论文中的学习率公式transformer/Optim.py#L26-L29lr lr_mul × d_model^(-0.5) × min(step^(-0.5), step × warmup^(-1.5))前半段warmup线性升温之后按步数的平方根倒数衰减——这就是常说的Noam 调度。脚本还会贴心地提醒batch size 小于 2048 而 warmup 不足 4000 时warmup 阶段可能还没训热就结束了train.py#L262-L266。 单个 Epochtrain_epoch 的 5 步循环核心函数 train_epoch() 中每个 batch 经历标准五步数据整形patch_src/patch_trgtrain.py#L61-L69把序列转置为 [seq, batch] 布局并做teacher forcing偏移——trg序列左移一位作为输入右移一位作为标签前向传播pred model(src_seq, trg_seq)计算损失调用cal_performance()反向传播loss.backward()后执行optimizer.step_and_update_lr()——注意这里是更新学习率 参数更新一步完成记账累计总损失与词级正确数每个 epoch 结束返回平均词损失与词准确率ppl困惑度则由exp(loss)换算train.py#L167。 损失函数里的 label smoothingcal_loss() 提供了两种模式普通交叉熵ignore_indexpad_idx直接跳过填充位标签平滑-label_smoothing开启时把正确答案的置信度从 1 降到 0.9剩余 0.1 均分给其他词train.py#L45-L55标签平滑能抑制模型过度自信是 Transformer 翻译任务提升 BLEU 的常用手段官方训练命令中默认开启。 验证与保存eval_epoch 和 checkpoint 策略eval_epoch() 与训练循环几乎同构但有两个关键区别model.eval()关闭 Dropout且整个循环包在torch.no_grad()中省显存又提速验证时不做label smoothing得到更真实的损失估计每个 epoch 结束后train() 根据-save_mode决定保存策略train.py#L181-L188all每个 epoch 都存一个带准确率的 checkpointbest仅当验证损失刷新历史最低时覆盖model.chkpt保存的 checkpoint 包含epoch、全部超参settings和model.state_dict()可直接被 translate.py 加载做推理。 TensorBoard 监控三条曲线看健康度加上-use_tb参数后脚本会在output_dir/tensorboard下写入事件文件train.py#L138-L141每个 epoch 记录三组指标train.py#L198-L201指标说明ppl训练/验证困惑度应随 epoch 稳步下降accuracy词级准确率train 与 val 曲线差距过大是过拟合信号learning_rate学习率直观验证 warmup 调度是否符合预期同时每个 epoch 的 loss、ppl、accuracy 还会以 CSV 格式追加写入train.log和valid.logtrain.py#L149-L151即使不用 TensorBoard 也能用任何图表工具复现曲线。 快速上手一条命令跑通以 Multi30k 德英翻译为例官方示例脚本 train_multi30k_de_en.sh 展示了推荐配置python train.py \ -data_pkl m30k_deen_shr.pkl \ -embs_share_weight -proj_share_weight -label_smoothing \ -b 256 -warmup 4000 -epoch 200 \ -output_dir output -use_tb数据预处理则交给 preprocess.py 提前完成先下载语料、构建 BPE 词表、dump 成 pkl。训练结束后用 translate.py 加载 checkpoint 即可翻译。 相关文件速查文件职责train.py训练入口与训练循环transformer/Models.pyTransformer 主模型、位置编码transformer/Layers.py编码器/解码器层transformer/SubLayers.py多头注意力、前馈网络transformer/Optim.pyNoam 学习率调度包装器transformer/Constants.pyPAD/BOS/EOS 特殊符号preprocess.py语料下载、BPE 编码、pkl 打包translate.py加载 checkpoint 做推理总结train.py 用一个不到 400 行的脚本展示了 Transformer 训练的标准范式——分桶数据加载、teacher forcing、warmup 学习率、标签平滑、best checkpoint 策略与 TensorBoard 监控。读懂它你就掌握了绝大多数 PyTorch 序列模型训练循环的骨架接下来只需替换模型和数据即可迁移到自己的项目上。【免费下载链接】attention-is-all-you-need-pytorchA PyTorch implementation of the Transformer model in Attention is All You Need.项目地址: https://gitcode.com/gh_mirrors/at/attention-is-all-you-need-pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表