Seq2Seq Transformer 中英翻译实战:从 GRU 到 Transformer,BLEU 0.225 → 0.307
项目地址(Gitee):https://gitee.com/Touari/seq2seq_zh_en_translation.git
本文记录
Seq2Seq_Transformer项目的完整构建过程:从Seq2Seq_Attention(GRU + Luong 注意力)复制出来 → 改造为 Transformer 架构 → 重新训练 → 实测对比。所有数据均为实测(RTX 3060 6GB)。
1. 项目概况
| 项 | 内容 |
|---|---|
| 前身 | Seq2Seq_Attention(GRU + Luong 注意力,完整复制保留 git 历史) |
| 任务 | 中译英(中文字符级 → 英文 Treebank 分词) |
| 模型 | Transformer (nn.Transformer,4 头注意力,2+2 层) |
| 数据 | Tatoeba cmn-eng,29,155 句对 |
| 词表 | 中文 2,747 / 英文 7,412(pad=0,sos=2,eos=3) |
| 参数量 | 4,761,716(约 476 万) |
| 关键结果 | 测试集 BLEU 0.1952 → 0.2252 → 0.3065(相对原版 +57.0%) |
2. 改造动机:从 RNN 走向 Transformer
前两个版本(GRU 原版、GRU+注意力版)的核心问题是串行:
#mermaid-svg-KH8KvXpdSmmTecfA{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-KH8KvXpdSmmTecfA .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-KH8KvXpdSmmTecfA .error-icon{fill:#552222;}#mermaid-svg-KH8KvXpdSmmTecfA .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-KH8KvXpdSmmTecfA .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-KH8KvXpdSmmTecfA .marker{fill:#333333;stroke:#333333;}#mermaid-svg-KH8KvXpdSmmTecfA .marker.cross{stroke:#333333;}#mermaid-svg-KH8KvXpdSmmTecfA svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-KH8KvXpdSmmTecfA p{margin:0;}#mermaid-svg-KH8KvXpdSmmTecfA .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-KH8KvXpdSmmTecfA .cluster-label text{fill:#333;}#mermaid-svg-KH8KvXpdSmmTecfA .cluster-label span{color:#333;}#mermaid-svg-KH8KvXpdSmmTecfA .cluster-label span p{background-color:transparent;}#mermaid-svg-KH8KvXpdSmmTecfA .label text,#mermaid-svg-KH8KvXpdSmmTecfA span{fill:#333;color:#333;}#mermaid-svg-KH8KvXpdSmmTecfA .node rect,#mermaid-svg-KH8KvXpdSmmTecfA .node circle,#mermaid-svg-KH8KvXpdSmmTecfA .node ellipse,#mermaid-svg-KH8KvXpdSmmTecfA .node polygon,#mermaid-svg-KH8KvXpdSmmTecfA .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-KH8KvXpdSmmTecfA .rough-node .label text,#mermaid-svg-KH8KvXpdSmmTecfA .node .label text,#mermaid-svg-KH8KvXpdSmmTecfA .image-shape .label,#mermaid-svg-KH8KvXpdSmmTecfA .icon-shape .label{text-anchor:middle;}#mermaid-svg-KH8KvXpdSmmTecfA .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-KH8KvXpdSmmTecfA .rough-node .label,#mermaid-svg-KH8KvXpdSmmTecfA .node .label,#mermaid-svg-KH8KvXpdSmmTecfA .image-shape .label,#mermaid-svg-KH8KvXpdSmmTecfA .icon-shape .label{text-align:center;}#mermaid-svg-KH8KvXpdSmmTecfA .node.clickable{cursor:pointer;}#mermaid-svg-KH8KvXpdSmmTecfA .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-KH8KvXpdSmmTecfA .arrowheadPath{fill:#333333;}#mermaid-svg-KH8KvXpdSmmTecfA .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-KH8KvXpdSmmTecfA .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-KH8KvXpdSmmTecfA .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-KH8KvXpdSmmTecfA .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-KH8KvXpdSmmTecfA .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-KH8KvXpdSmmTecfA .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-KH8KvXpdSmmTecfA .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-KH8KvXpdSmmTecfA .cluster text{fill:#333;}#mermaid-svg-KH8KvXpdSmmTecfA .cluster span{color:#333;}#mermaid-svg-KH8KvXpdSmmTecfA div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-KH8KvXpdSmmTecfA .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-KH8KvXpdSmmTecfA rect.text{fill:none;stroke-width:0;}#mermaid-svg-KH8KvXpdSmmTecfA .icon-shape,#mermaid-svg-KH8KvXpdSmmTecfA .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-KH8KvXpdSmmTecfA .icon-shape p,#mermaid-svg-KH8KvXpdSmmTecfA .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-KH8KvXpdSmmTecfA .icon-shape .label rect,#mermaid-svg-KH8KvXpdSmmTecfA .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-KH8KvXpdSmmTecfA .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-KH8KvXpdSmmTecfA .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-KH8KvXpdSmmTecfA :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} GRU Decoder(串行生成)
GRU Encoder(串行读取)
中文句子
逐步读取
每个词依赖前一个隐状态
output 序列
逐词生成
每步依赖上一步隐状态
- 训练慢:句子有多长,就要循环多少次,GPU 无法并行
- 长距离依赖弱:信息沿时间轴逐级传递,早期信息容易衰减(即使加了注意力,也只能在解码端"补救",编码端依然是串行的)
Transformer 的答案是完全抛弃循环,用自注意力一次性并行处理整句:
#mermaid-svg-xtg6VNYBTF5hc03C{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-xtg6VNYBTF5hc03C .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-xtg6VNYBTF5hc03C .error-icon{fill:#552222;}#mermaid-svg-xtg6VNYBTF5hc03C .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-xtg6VNYBTF5hc03C .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-xtg6VNYBTF5hc03C .marker{fill:#333333;stroke:#333333;}#mermaid-svg-xtg6VNYBTF5hc03C .marker.cross{stroke:#333333;}#mermaid-svg-xtg6VNYBTF5hc03C svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-xtg6VNYBTF5hc03C p{margin:0;}#mermaid-svg-xtg6VNYBTF5hc03C .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-xtg6VNYBTF5hc03C .cluster-label text{fill:#333;}#mermaid-svg-xtg6VNYBTF5hc03C .cluster-label span{color:#333;}#mermaid-svg-xtg6VNYBTF5hc03C .cluster-label span p{background-color:transparent;}#mermaid-svg-xtg6VNYBTF5hc03C .label text,#mermaid-svg-xtg6VNYBTF5hc03C span{fill:#333;color:#333;}#mermaid-svg-xtg6VNYBTF5hc03C .node rect,#mermaid-svg-xtg6VNYBTF5hc03C .node circle,#mermaid-svg-xtg6VNYBTF5hc03C .node ellipse,#mermaid-svg-xtg6VNYBTF5hc03C .node polygon,#mermaid-svg-xtg6VNYBTF5hc03C .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-xtg6VNYBTF5hc03C .rough-node .label text,#mermaid-svg-xtg6VNYBTF5hc03C .node .label text,#mermaid-svg-xtg6VNYBTF5hc03C .image-shape .label,#mermaid-svg-xtg6VNYBTF5hc03C .icon-shape .label{text-anchor:middle;}#mermaid-svg-xtg6VNYBTF5hc03C .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-xtg6VNYBTF5hc03C .rough-node .label,#mermaid-svg-xtg6VNYBTF5hc03C .node .label,#mermaid-svg-xtg6VNYBTF5hc03C .image-shape .label,#mermaid-svg-xtg6VNYBTF5hc03C .icon-shape .label{text-align:center;}#mermaid-svg-xtg6VNYBTF5hc03C .node.clickable{cursor:pointer;}#mermaid-svg-xtg6VNYBTF5hc03C .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-xtg6VNYBTF5hc03C .arrowheadPath{fill:#333333;}#mermaid-svg-xtg6VNYBTF5hc03C .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-xtg6VNYBTF5hc03C .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-xtg6VNYBTF5hc03C .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-xtg6VNYBTF5hc03C .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-xtg6VNYBTF5hc03C .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-xtg6VNYBTF5hc03C .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-xtg6VNYBTF5hc03C .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-xtg6VNYBTF5hc03C .cluster text{fill:#333;}#mermaid-svg-xtg6VNYBTF5hc03C .cluster span{color:#333;}#mermaid-svg-xtg6VNYBTF5hc03C div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-xtg6VNYBTF5hc03C .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-xtg6VNYBTF5hc03C rect.text{fill:none;stroke-width:0;}#mermaid-svg-xtg6VNYBTF5hc03C .icon-shape,#mermaid-svg-xtg6VNYBTF5hc03C .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-xtg6VNYBTF5hc03C .icon-shape p,#mermaid-svg-xtg6VNYBTF5hc03C .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-xtg6VNYBTF5hc03C .icon-shape .label rect,#mermaid-svg-xtg6VNYBTF5hc03C .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-xtg6VNYBTF5hc03C .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-xtg6VNYBTF5hc03C .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-xtg6VNYBTF5hc03C :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} Transformer Decoder(并行训练)
Transformer Encoder(并行)
整句同时输入
self-attention 一次看全
memory
目标句同时输入
masked self-attention + 交叉注意力
直观类比:GRU 像"顺着念完再倒着写",Transformer 像"整页摊开,任意两个词之间直连"。
2.1 构建步骤与 git 提交历史
整个改造按"模型层 → 训练层 → 推理层"三批提交,与代码依赖顺序一致:
#mermaid-svg-6oj205CB7yUh68hL{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-6oj205CB7yUh68hL .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-6oj205CB7yUh68hL .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-6oj205CB7yUh68hL .error-icon{fill:#552222;}#mermaid-svg-6oj205CB7yUh68hL .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-6oj205CB7yUh68hL .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-6oj205CB7yUh68hL .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-6oj205CB7yUh68hL .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-6oj205CB7yUh68hL .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-6oj205CB7yUh68hL .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-6oj205CB7yUh68hL .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-6oj205CB7yUh68hL .marker{fill:#333333;stroke:#333333;}#mermaid-svg-6oj205CB7yUh68hL .marker.cross{stroke:#333333;}#mermaid-svg-6oj205CB7yUh68hL svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-6oj205CB7yUh68hL p{margin:0;}#mermaid-svg-6oj205CB7yUh68hL .commit-id,#mermaid-svg-6oj205CB7yUh68hL .commit-msg,#mermaid-svg-6oj205CB7yUh68hL .branch-label{fill:lightgrey;color:lightgrey;font-family:'trebuchet ms',verdana,arial,sans-serif;font-family:var(--mermaid-font-family);}#mermaid-svg-6oj205CB7yUh68hL .branch-label0{fill:#ffffff;}#mermaid-svg-6oj205CB7yUh68hL .commit0{stroke:hsl(240, 100%, 46.2745098039%);fill:hsl(240, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight0{stroke:hsl(60, 100%, 3.7254901961%);fill:hsl(60, 100%, 3.7254901961%);}#mermaid-svg-6oj205CB7yUh68hL .label0{fill:hsl(240, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow0{stroke:hsl(240, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label1{fill:black;}#mermaid-svg-6oj205CB7yUh68hL .commit1{stroke:hsl(60, 100%, 43.5294117647%);fill:hsl(60, 100%, 43.5294117647%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight1{stroke:rgb(0, 0, 160.5);fill:rgb(0, 0, 160.5);}#mermaid-svg-6oj205CB7yUh68hL .label1{fill:hsl(60, 100%, 43.5294117647%);}#mermaid-svg-6oj205CB7yUh68hL .arrow1{stroke:hsl(60, 100%, 43.5294117647%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label2{fill:black;}#mermaid-svg-6oj205CB7yUh68hL .commit2{stroke:hsl(80, 100%, 46.2745098039%);fill:hsl(80, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight2{stroke:rgb(48.8333333334, 0, 146.5000000001);fill:rgb(48.8333333334, 0, 146.5000000001);}#mermaid-svg-6oj205CB7yUh68hL .label2{fill:hsl(80, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow2{stroke:hsl(80, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label3{fill:#ffffff;}#mermaid-svg-6oj205CB7yUh68hL .commit3{stroke:hsl(210, 100%, 46.2745098039%);fill:hsl(210, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight3{stroke:rgb(146.5000000001, 73.2500000001, 0);fill:rgb(146.5000000001, 73.2500000001, 0);}#mermaid-svg-6oj205CB7yUh68hL .label3{fill:hsl(210, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow3{stroke:hsl(210, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label4{fill:black;}#mermaid-svg-6oj205CB7yUh68hL .commit4{stroke:hsl(180, 100%, 46.2745098039%);fill:hsl(180, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight4{stroke:rgb(146.5000000001, 0, 0);fill:rgb(146.5000000001, 0, 0);}#mermaid-svg-6oj205CB7yUh68hL .label4{fill:hsl(180, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow4{stroke:hsl(180, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label5{fill:black;}#mermaid-svg-6oj205CB7yUh68hL .commit5{stroke:hsl(150, 100%, 46.2745098039%);fill:hsl(150, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight5{stroke:rgb(146.5000000001, 0, 73.2500000001);fill:rgb(146.5000000001, 0, 73.2500000001);}#mermaid-svg-6oj205CB7yUh68hL .label5{fill:hsl(150, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow5{stroke:hsl(150, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label6{fill:black;}#mermaid-svg-6oj205CB7yUh68hL .commit6{stroke:hsl(300, 100%, 46.2745098039%);fill:hsl(300, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight6{stroke:rgb(0, 146.5000000001, 0);fill:rgb(0, 146.5000000001, 0);}#mermaid-svg-6oj205CB7yUh68hL .label6{fill:hsl(300, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow6{stroke:hsl(300, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch-label7{fill:black;}#mermaid-svg-6oj205CB7yUh68hL .commit7{stroke:hsl(0, 100%, 46.2745098039%);fill:hsl(0, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight7{stroke:rgb(0, 146.5000000001, 146.5000000001);fill:rgb(0, 146.5000000001, 146.5000000001);}#mermaid-svg-6oj205CB7yUh68hL .label7{fill:hsl(0, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .arrow7{stroke:hsl(0, 100%, 46.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .branch{stroke-width:1;stroke:#333333;stroke-dasharray:2;}#mermaid-svg-6oj205CB7yUh68hL .commit-label{font-size:10px;fill:#000021;}#mermaid-svg-6oj205CB7yUh68hL .commit-label-bkg{font-size:10px;fill:#ffffde;opacity:0.5;}#mermaid-svg-6oj205CB7yUh68hL .tag-label{font-size:10px;fill:#131300;}#mermaid-svg-6oj205CB7yUh68hL .tag-label-bkg{fill:#ECECFF;stroke:hsl(240, 60%, 86.2745098039%);}#mermaid-svg-6oj205CB7yUh68hL .tag-hole{fill:#333;}#mermaid-svg-6oj205CB7yUh68hL .commit-merge{stroke:#ECECFF;fill:#ECECFF;}#mermaid-svg-6oj205CB7yUh68hL .commit-reverse{stroke:#ECECFF;fill:#ECECFF;stroke-width:3;}#mermaid-svg-6oj205CB7yUh68hL .commit-highlight-inner{stroke:#ECECFF;fill:#ECECFF;}#mermaid-svg-6oj205CB7yUh68hL .arrow{stroke-width:8;stroke-linecap:round;fill:none;}#mermaid-svg-6oj205CB7yUh68hL .gitTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-6oj205CB7yUh68hL :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} main transformer 4665558 完成 BLEU 评估 3c67444 实现注意力机制 3b6f0e6 重构模型架构 df56a78 训练并行化 7cc90e1 推理适配
说明:前两个 commit 是复制自
Seq2Seq_Attention的历史(保留 git 与远程),transformer分支后的三个 commit 为本次改造。其中7cc90e1的消息误用了训练循环的标题,不影响代码。
步骤 1:模型层(src/model.py + src/config.py)
| 改动 | 内容 |
|---|---|
| 删除 | Attention、TranslationEncoder(GRU)、TranslationDecoder(GRU) |
| 新增 | PositionEncoding:sin/cos 位置编码,register_buffer 注册 |
| 新增 | TranslationModel:中英双 Embedding + nn.Transformer + 线性输出层 |
| 配置 | EMBEDDING_DIM/HIDDEN_DIM → DIM_MODEL=128 / NUM_HEADS=4 / 2+2 层 |
步骤 2:训练层(src/train.py)
删除"编码 + for 循环逐步解码",替换为一次性并行前向:
python
src_pad_mask = (encoder_inputs == model.zh_embedding.padding_idx)
tgt_mask = model.transformer.generate_square_subsequent_mask(decoder_inputs.shape[1])
decoder_outputs = model(encoder_inputs, decoder_inputs, src_pad_mask, tgt_mask)
步骤 3:推理层(src/predict.py)
- 编码只做一次:
memory = model.encode(inputs, src_pad_mask)缓存复用 - 自回归循环:每步重建
tgt_mask(.to(device)),decode后取[:, -1, :] - 生成序列逐步拼接,eos 终止判定保留
evaluate.py无需修改 :predict_batch保留了device参数,签名兼容。
3. 模型架构总览
#mermaid-svg-dxYLon9QuOgmZm0Z{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-dxYLon9QuOgmZm0Z .error-icon{fill:#552222;}#mermaid-svg-dxYLon9QuOgmZm0Z .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-dxYLon9QuOgmZm0Z .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-dxYLon9QuOgmZm0Z .marker{fill:#333333;stroke:#333333;}#mermaid-svg-dxYLon9QuOgmZm0Z .marker.cross{stroke:#333333;}#mermaid-svg-dxYLon9QuOgmZm0Z svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-dxYLon9QuOgmZm0Z p{margin:0;}#mermaid-svg-dxYLon9QuOgmZm0Z .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z .cluster-label text{fill:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z .cluster-label span{color:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z .cluster-label span p{background-color:transparent;}#mermaid-svg-dxYLon9QuOgmZm0Z .label text,#mermaid-svg-dxYLon9QuOgmZm0Z span{fill:#333;color:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z .node rect,#mermaid-svg-dxYLon9QuOgmZm0Z .node circle,#mermaid-svg-dxYLon9QuOgmZm0Z .node ellipse,#mermaid-svg-dxYLon9QuOgmZm0Z .node polygon,#mermaid-svg-dxYLon9QuOgmZm0Z .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-dxYLon9QuOgmZm0Z .rough-node .label text,#mermaid-svg-dxYLon9QuOgmZm0Z .node .label text,#mermaid-svg-dxYLon9QuOgmZm0Z .image-shape .label,#mermaid-svg-dxYLon9QuOgmZm0Z .icon-shape .label{text-anchor:middle;}#mermaid-svg-dxYLon9QuOgmZm0Z .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-dxYLon9QuOgmZm0Z .rough-node .label,#mermaid-svg-dxYLon9QuOgmZm0Z .node .label,#mermaid-svg-dxYLon9QuOgmZm0Z .image-shape .label,#mermaid-svg-dxYLon9QuOgmZm0Z .icon-shape .label{text-align:center;}#mermaid-svg-dxYLon9QuOgmZm0Z .node.clickable{cursor:pointer;}#mermaid-svg-dxYLon9QuOgmZm0Z .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-dxYLon9QuOgmZm0Z .arrowheadPath{fill:#333333;}#mermaid-svg-dxYLon9QuOgmZm0Z .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-dxYLon9QuOgmZm0Z .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-dxYLon9QuOgmZm0Z .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-dxYLon9QuOgmZm0Z .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-dxYLon9QuOgmZm0Z .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-dxYLon9QuOgmZm0Z .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-dxYLon9QuOgmZm0Z .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-dxYLon9QuOgmZm0Z .cluster text{fill:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z .cluster span{color:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-dxYLon9QuOgmZm0Z .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-dxYLon9QuOgmZm0Z rect.text{fill:none;stroke-width:0;}#mermaid-svg-dxYLon9QuOgmZm0Z .icon-shape,#mermaid-svg-dxYLon9QuOgmZm0Z .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-dxYLon9QuOgmZm0Z .icon-shape p,#mermaid-svg-dxYLon9QuOgmZm0Z .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-dxYLon9QuOgmZm0Z .icon-shape .label rect,#mermaid-svg-dxYLon9QuOgmZm0Z .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-dxYLon9QuOgmZm0Z .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-dxYLon9QuOgmZm0Z .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-dxYLon9QuOgmZm0Z :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 目标语言侧(英文)
源语言侧(中文)
memory
(batch, src_len, 128)
src (batch, src_len)
zh_embedding
Embedding(2747→128)
- 位置编码
PositionEncoding
Transformer Encoder
2 层 × 4 头
tgt (batch, tgt_len)
en_embedding
Embedding(7412→128)
- 位置编码
Transformer Decoder
2 层 × 4 头
masked self-attn + cross-attn
Linear(128→7412)
逐位置输出词表 logits
关键设计:模型拆分为 encode / decode / forward 三个方法:
| 方法 | 用途 | 调用场景 |
|---|---|---|
encode(src, src_pad_mask) |
编码,返回 memory | 训练(forward 内部)/ 推理(只调用一次,缓存 memory) |
decode(tgt, memory, tgt_mask, memory_pad_mask) |
解码,返回词表 logits | 推理(每步自回归调用) |
forward(src, tgt, src_pad_mask, tgt_mask) |
组合完整前向 | 训练(一次并行) |
4. 核心机制详解
4.1 注意力计算的数学本质(QKV)
Transformer 的一切都建立在缩放点积注意力之上:
Attention(Q,K,V)=softmax(QK⊤dk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)VAttention(Q,K,V)=softmax(dk QK⊤)V
#mermaid-svg-JkeynmLpF7NzvKYA{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-JkeynmLpF7NzvKYA .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-JkeynmLpF7NzvKYA .error-icon{fill:#552222;}#mermaid-svg-JkeynmLpF7NzvKYA .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-JkeynmLpF7NzvKYA .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-JkeynmLpF7NzvKYA .marker{fill:#333333;stroke:#333333;}#mermaid-svg-JkeynmLpF7NzvKYA .marker.cross{stroke:#333333;}#mermaid-svg-JkeynmLpF7NzvKYA svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-JkeynmLpF7NzvKYA p{margin:0;}#mermaid-svg-JkeynmLpF7NzvKYA .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-JkeynmLpF7NzvKYA .cluster-label text{fill:#333;}#mermaid-svg-JkeynmLpF7NzvKYA .cluster-label span{color:#333;}#mermaid-svg-JkeynmLpF7NzvKYA .cluster-label span p{background-color:transparent;}#mermaid-svg-JkeynmLpF7NzvKYA .label text,#mermaid-svg-JkeynmLpF7NzvKYA span{fill:#333;color:#333;}#mermaid-svg-JkeynmLpF7NzvKYA .node rect,#mermaid-svg-JkeynmLpF7NzvKYA .node circle,#mermaid-svg-JkeynmLpF7NzvKYA .node ellipse,#mermaid-svg-JkeynmLpF7NzvKYA .node polygon,#mermaid-svg-JkeynmLpF7NzvKYA .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-JkeynmLpF7NzvKYA .rough-node .label text,#mermaid-svg-JkeynmLpF7NzvKYA .node .label text,#mermaid-svg-JkeynmLpF7NzvKYA .image-shape .label,#mermaid-svg-JkeynmLpF7NzvKYA .icon-shape .label{text-anchor:middle;}#mermaid-svg-JkeynmLpF7NzvKYA .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-JkeynmLpF7NzvKYA .rough-node .label,#mermaid-svg-JkeynmLpF7NzvKYA .node .label,#mermaid-svg-JkeynmLpF7NzvKYA .image-shape .label,#mermaid-svg-JkeynmLpF7NzvKYA .icon-shape .label{text-align:center;}#mermaid-svg-JkeynmLpF7NzvKYA .node.clickable{cursor:pointer;}#mermaid-svg-JkeynmLpF7NzvKYA .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-JkeynmLpF7NzvKYA .arrowheadPath{fill:#333333;}#mermaid-svg-JkeynmLpF7NzvKYA .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-JkeynmLpF7NzvKYA .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-JkeynmLpF7NzvKYA .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-JkeynmLpF7NzvKYA .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-JkeynmLpF7NzvKYA .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-JkeynmLpF7NzvKYA .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-JkeynmLpF7NzvKYA .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-JkeynmLpF7NzvKYA .cluster text{fill:#333;}#mermaid-svg-JkeynmLpF7NzvKYA .cluster span{color:#333;}#mermaid-svg-JkeynmLpF7NzvKYA div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-JkeynmLpF7NzvKYA .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-JkeynmLpF7NzvKYA rect.text{fill:none;stroke-width:0;}#mermaid-svg-JkeynmLpF7NzvKYA .icon-shape,#mermaid-svg-JkeynmLpF7NzvKYA .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-JkeynmLpF7NzvKYA .icon-shape p,#mermaid-svg-JkeynmLpF7NzvKYA .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-JkeynmLpF7NzvKYA .icon-shape .label rect,#mermaid-svg-JkeynmLpF7NzvKYA .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-JkeynmLpF7NzvKYA .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-JkeynmLpF7NzvKYA .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-JkeynmLpF7NzvKYA :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} Q (query) 查询
我要找什么
QKᵀ 打分
每个位置对彼此的关联强度
K (key) 键
我能提供什么
÷ √d_k 缩放
防止点积过大
softmax 梯度消失
softmax 归一化
注意力权重
V (value) 值
实际携带的信息
加权求和
按权重聚合信息
输出
- 缩放 √d_k 的意义:点积随维度增大而增大,softmax 会进入梯度饱和区;除以 √d_k 把方差拉回 1
- 多头注意力:把 128 维切成 4 个头(每头 32 维),各自独立计算注意力再拼接------让模型同时关注不同的"关系类型"(语法关系、指代关系、语义相似性)
- 本项目中的三类注意力:
| 类型 | Q / K / V 来源 | 作用 |
|---|---|---|
| Encoder 自注意力 | Q=K=V=中文序列 | 中文词之间互相建立联系 |
| Decoder 自注意力 | Q=K=V=英文序列(masked) | 英文词之间建立联系,且只看已生成的词 |
| 交叉注意力 | Q=英文,K=V=memory | 英文词"查阅"中文原文对应位置 |
4.2 位置编码(PositionEncoding)
Transformer 没有循环结构,序列顺序信息必须显式注入,否则"我爱你"和"你爱我"的表示完全相同。采用原论文的固定正弦余弦编码:
PE(pos,2i)=sin(pos100002i/dmodel),PE(pos,2i+1)=cos(pos100002i/dmodel)PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right), \quad PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)PE(pos,2i)=sin(100002i/dmodelpos),PE(pos,2i+1)=cos(100002i/dmodelpos)
python
for pos in range(max_len): # 位置维度
for _2i in range(0, dim_model, 2): # 特征维度(步长 2)
pe[pos, _2i] = math.sin(pos / (10000 ** (_2i / dim_model)))
pe[pos, _2i + 1] = math.cos(pos / (10000 ** (_2i / dim_model)))
要点:
- 偶维 sin / 奇维 cos:不同维度拥有不同波长,使模型能区分绝对位置与相对位置
register_buffer('pe', pe)注册:buffer 随.to(device)迁移、不参与梯度forward中按x.size(1)截取前 seq_len 行,与 embedding 直接广播相加
4.3 三种 Mask(最容易混淆的部分)
| Mask | 形状 | 作用 | 构造方式 |
|---|---|---|---|
tgt_mask(因果/下三角) |
(tgt_len, tgt_len) |
训练并行时屏蔽未来位置,第 i 个位置只能看 0...i-1 | generate_square_subsequent_mask(tgt_len) |
src_key_padding_mask |
(batch, src_len) |
屏蔽中文句的 pad 位置参与注意力 | (src == padding_idx) |
memory_key_padding_mask |
(batch, src_len) |
交叉注意力中屏蔽 pad | 同 src 的 mask(复用 src_pad_mask) |
tgt_mask 的可视化 (4 个位置,上三角为 -inf):
位置0 位置1 位置2 位置3
位置0 0 -inf -inf -inf
位置1 0 0 -inf -inf
位置2 0 0 0 -inf
位置3 0 0 0 0
第 i 行的含义:位置 i 生成时能"看"哪些位置 ------只能看 0...i(含自身),后面的全是 -inf,softmax 后权重为 0,等效屏蔽。这就是"训练并行、行为等价于串行"的保障。
⚠️ 注意
tgt_mask与 padding mask 语义差异:tgt_mask是"位置对位置"的方阵(-inf屏蔽未来),padding mask 是"样本对位置"的 bool 矩阵(True屏蔽 pad),形状不同、传参位置不同(tgt_maskvstgt_key_padding_mask),不要混淆。
5. 训练机制:一次性并行前向
GRU 版训练是 for 循环逐步喂词(teacher forcing 逐时间步),Transformer 版整个目标序列一次前向:
python
src_pad_mask = (encoder_inputs == model.zh_embedding.padding_idx) # (batch, src_len)
tgt_mask = model.transformer.generate_square_subsequent_mask(decoder_inputs.shape[1])
decoder_outputs = model(encoder_inputs, decoder_inputs, src_pad_mask, tgt_mask)
loss = loss_fn(decoder_outputs.reshape(-1, decoder_outputs.shape[-1]),
decoder_targets.reshape(-1))
decoder_inputs = targets[:, :-1](右移一位,去掉句尾 eos)decoder_targets = targets[:, 1:](去掉句首 sos,作为监督标签)- 因果 mask 保证第 i 个位置的预测只依赖前 i-1 个真实词------并行下的 teacher forcing 与串行等价
- 损失用
CrossEntropyLoss(ignore_index=pad),pad 位置不参与
实测收敛曲线(50 epoch,TensorBoard 记录):
| Epoch | Loss | Epoch | Loss |
|---|---|---|---|
| 2 | 4.1853 | 31 | 0.3705 |
| 11 | 0.8471 | 41 | 0.3168 |
| 21 | 0.4835 | 50 | 0.2820(最优) |
全程单调下降、无震荡,最终 loss 0.2820。
6. 推理机制:自回归生成
推理时没有参考答案,必须逐个词生成,且每步把已生成的序列重新喂入模型:
python
src_pad_mask = (inputs == model.zh_embedding.padding_idx)
memory = model.encode(inputs, src_pad_mask) # 编码只做一次,缓存复用
decoder_input = torch.full([batch_size, 1], en_tokenizer.sos_token_index, device=device)
for i in range(config.MAX_SEQ_LEN):
tgt_mask = model.transformer.generate_square_subsequent_mask(decoder_input.size(1)).to(device)
decoder_output = model.decode(decoder_input, memory, tgt_mask, src_pad_mask)
next_token = torch.argmax(decoder_output[:, -1, :], dim=-1, keepdim=True) # 只取最后位置
decoder_input = torch.cat([decoder_input, next_token], dim=-1) # 拼接增长
if (next_token.squeeze(1) == en_tokenizer.eos_token_index).all():
break
三个关键理解:
- 为什么取
[:, -1, :]:decode对整段已生成序列重算,输出(batch, 当前长度, vocab),新词永远在最后一位------与 GRU 版"每步只输出一个词"不同 - 为什么每步重建
tgt_mask:序列长度在增长,方阵尺寸要跟着长 - 为什么 memory 只算一次 :中文句子不变,编码结果不变;
encode/decode拆分正是为推理的"一次编码、多次解码"设计的(KV cache 是进一步的优化,本项目未采用)
自回归生成时序:
翻译过程 模型 翻译过程 模型 #mermaid-svg-IpRtk2OMF2Tv33hR{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-IpRtk2OMF2Tv33hR .error-icon{fill:#552222;}#mermaid-svg-IpRtk2OMF2Tv33hR .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-IpRtk2OMF2Tv33hR .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-IpRtk2OMF2Tv33hR .marker{fill:#333333;stroke:#333333;}#mermaid-svg-IpRtk2OMF2Tv33hR .marker.cross{stroke:#333333;}#mermaid-svg-IpRtk2OMF2Tv33hR svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-IpRtk2OMF2Tv33hR p{margin:0;}#mermaid-svg-IpRtk2OMF2Tv33hR .actor{stroke:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);fill:#ECECFF;}#mermaid-svg-IpRtk2OMF2Tv33hR text.actor>tspan{fill:black;stroke:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .actor-line{stroke:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);}#mermaid-svg-IpRtk2OMF2Tv33hR .innerArc{stroke-width:1.5;stroke-dasharray:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .messageLine0{stroke-width:1.5;stroke-dasharray:none;stroke:#333;}#mermaid-svg-IpRtk2OMF2Tv33hR .messageLine1{stroke-width:1.5;stroke-dasharray:2,2;stroke:#333;}#mermaid-svg-IpRtk2OMF2Tv33hR #arrowhead path{fill:#333;stroke:#333;}#mermaid-svg-IpRtk2OMF2Tv33hR .sequenceNumber{fill:white;}#mermaid-svg-IpRtk2OMF2Tv33hR #sequencenumber{fill:#333;}#mermaid-svg-IpRtk2OMF2Tv33hR #crosshead path{fill:#333;stroke:#333;}#mermaid-svg-IpRtk2OMF2Tv33hR .messageText{fill:#333;stroke:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .labelBox{stroke:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);fill:#ECECFF;}#mermaid-svg-IpRtk2OMF2Tv33hR .labelText,#mermaid-svg-IpRtk2OMF2Tv33hR .labelText>tspan{fill:black;stroke:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .loopText,#mermaid-svg-IpRtk2OMF2Tv33hR .loopText>tspan{fill:black;stroke:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .loopLine{stroke-width:2px;stroke-dasharray:2,2;stroke:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);fill:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);}#mermaid-svg-IpRtk2OMF2Tv33hR .note{stroke:#aaaa33;fill:#fff5ad;}#mermaid-svg-IpRtk2OMF2Tv33hR .noteText,#mermaid-svg-IpRtk2OMF2Tv33hR .noteText>tspan{fill:black;stroke:none;}#mermaid-svg-IpRtk2OMF2Tv33hR .activation0{fill:#f4f4f4;stroke:#666;}#mermaid-svg-IpRtk2OMF2Tv33hR .activation1{fill:#f4f4f4;stroke:#666;}#mermaid-svg-IpRtk2OMF2Tv33hR .activation2{fill:#f4f4f4;stroke:#666;}#mermaid-svg-IpRtk2OMF2Tv33hR .actorPopupMenu{position:absolute;}#mermaid-svg-IpRtk2OMF2Tv33hR .actorPopupMenuPanel{position:absolute;fill:#ECECFF;box-shadow:0px 8px 16px 0px rgba(0,0,0,0.2);filter:drop-shadow(3px 5px 2px rgb(0 0 0 / 0.4));}#mermaid-svg-IpRtk2OMF2Tv33hR .actor-man line{stroke:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);fill:#ECECFF;}#mermaid-svg-IpRtk2OMF2Tv33hR .actor-man circle,#mermaid-svg-IpRtk2OMF2Tv33hR line{stroke:hsl(259.6261682243, 59.7765363128%, 87.9019607843%);fill:#ECECFF;stroke-width:2px;}#mermaid-svg-IpRtk2OMF2Tv33hR :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 编码阶段(只做一次) 解码阶段(逐步自回归) 直到预测出 eos 或达到 MAX_SEQ_LEN encode(中文, src_pad_mask)memory(中文语义缓存)decode(sos, memory, mask₁)预测词₁ (取最后位置)decode(sos, 词₁, memory, mask₂)预测词₂decode(sos, 词₁, 词₂, memory, mask₃)预测词₃
7. 改造中的踩坑记录(实测排雷)
| # | 位置 | 现象 | 原因 | 解法 |
|---|---|---|---|---|
| 1 | PositionEncoding |
GPU 训练时 x + part_pe 设备不匹配 |
self.pe 是普通 tensor,.to(device) 不会迁移普通属性 |
self.register_buffer('pe', pe) |
| 2 | PositionEncoding.forward |
TypeError: 'builtin_function_or_method' object is not subscriptable |
x.size 是方法不是属性,x.size[1] 应为 x.size(1) |
x.size(1) 或 x.shape[1] |
| 3 | nn.Transformer 实例化 |
RuntimeError: the batch number of src and tgt must be equal |
默认 batch_first=False,输入被解释成 (seq, batch, dim),把 10 和 8 当成了 batch 数 |
实例化传 batch_first=True |
| 4 | tgt_mask 设备 |
预警过 device mismatch | generate_square_subsequent_mask 默认生成在 CPU |
实测未触发;predict.py 中主动 .to(device) 双保险 |
踩坑 3 的报错信息解读 (很有迷惑性):不传 batch_first=True 时,src=(2,10,128) 被读成"batch=10"、tgt=(2,8,128) 被读成"batch=8",10≠8 所以报"batch 数必须相等"------不是真的 batch 不相等,是维度解释错了。
8. 结果对比
| 版本 | 架构 | 参数量 | BLEU | 相对原版 |
|---|---|---|---|---|
| Seq2Seq_Translation | GRU | 5,695,604 | 0.1952 | --- |
| Seq2Seq_Attention | GRU + Luong | 5,695,604 | 0.2252 | +15.4% |
| Seq2Seq_Transformer | Transformer | 4,761,716 | 0.3065 | +57.0% |
结论:Transformer 用更少的参数(476 万 vs 569 万,-16.4%)取得了更高的 BLEU(+36% vs 注意力版),验证了 self-attention 架构在机器翻译上的优势。
与 GRU 系列架构对比
| 维度 | GRU + Luong | Transformer(本项目) |
|---|---|---|
| 参数量 | 5,695,604 | 4,761,716(-16.4%) |
| BLEU | 0.2252 | 0.3065(+36%) |
| 训练 | 逐步循环(串行) | 一次性并行 |
| 长距离依赖 | 沿时间轴衰减,注意力补救 | 任意两位置直接关联(O(n²) 注意力) |
| 位置信息 | 隐状态天然含顺序 | 需显式位置编码 |
| 复杂度 | 时间 O(n),无法并行 | 时间 O(n²),但可并行(GPU 友好) |
本质差异:GRU 是"压缩式记忆"(整句压成隐状态逐步传递),Transformer 是"直接寻址式记忆"(任意词之间直接建立联系)。RNN 的 O(n) 串行在长序列上反而慢于 O(n²) 的并行注意力。
翻译示例(实测)
| 中文 | Transformer 输出 | 点评 |
|---|---|---|
| 今天天气很好 | It's fine today. | ✅ 准确 |
| 他正在图书馆看书 | He is reading in the library. | ✅ 准确 |
| 我想去北京旅游 | I want to go to Tokyo... | ⚠️ 专有名词按数据频率猜错 |
| 这个苹果多少钱 | This apple juice cost of pineapples. | ⚠️ 抓对结构但出现词语幻觉 |
9. 已知局限与下一步
- 专有名词易错(实测"北京"→"Tokyo"):训练数据中 Tokyo 出现频率更高,模型按先验"猜词"------数据分布问题,非架构缺陷
- 小模型 + 小数据下存在词语幻觉(如 "This apple juice cost of pineapples")
- 下一步可选方向:
- 学习率 warmup + 衰减(Transformer 对 LR 敏感,当前固定 lr=1e-3 已收敛良好,可尝试进一步压榨)
- 增加层数/维度、beam search 解码、检查 tgt_mask 用因果 mask + padding mask 的组合
- 集成到 Agentic AI 项目的 AI 服务容器化栈(与 LSTM 情感 API、Qwen-VL 并列)
小结
从 GRU 到 GRU+Attention 再到 Transformer,同一个中译英任务的三次迭代清晰地展示了序列建模的进化主线:串行循环 → 循环+注意力补救 → 完全并行注意力。Transformer 版以 -16.4% 的参数换来 +57% 的 BLEU,也印证了"Attention Is All You Need"------而本项目 4 头注意力、2+2 层的小配置,正是理解现代大语言模型架构的最小可行样本。
完整代码与数据见项目仓库:https://gitee.com/Touari/seq2seq_zh_en_translation.git