欢迎光临
我们一直在努力

Transformer模型:从 Attention 到 PyTorch 从零实战

Transformer 模型详解:从 Attention 到 PyTorch 从零实战

一句话总结:Transformer 是 Google 在 2017 年论文 Attention Is All You Need 中提出的一种完全基于注意力机制、抛弃循环神经网络(RNN) 的序列建模架构。它通过「自注意力(Self-Attention) + 多头(Multi-Head) + 位置编码(Positional Encoding)」三大设计,让模型可以并行处理整条序列并直接捕捉任意两个位置的依赖关系,从而在机器翻译等任务上大幅超越 RNN,并成为 BERT、GPT、GPT-3、ChatGPT 以及今天所有大语言模型的共同底座——「Attention Is All You Need」名副其实地重塑了整个 NLP 与深度学习领域。


📋 目录

  • 一、为什么需要 Transformer
    • 1.1 RNN 系列模型的痛点
    • 1.2 Transformer 的三大核心思想
    • 1.3 关键信息一览
  • 二、整体架构速览
  • 三、核心组件逐一拆解
    • 3.1 Embedding:把词变成向量
    • 3.2 位置编码 Positional Encoding
    • 3.3 自注意力 Self-Attention
    • 3.4 缩放点积注意力公式手算
    • 3.5 多头注意力 Multi-Head Attention
    • 3.6 前馈网络 FFN
    • 3.7 残差连接与 LayerNorm
    • 3.8 掩码 Mask 的三个用途
  • 四、Encoder 与 Decoder 详解
    • 4.1 Encoder:双向理解输入
    • 4.2 Decoder:自回归生成输出
    • 4.3 为什么 Decoder 要「掩码 + 交叉注意力」
    • 4.4 Encoder-only / Decoder-only / Encoder-Decoder
  • 五、Transformer vs RNN
  • 六、环境准备
  • 七、PyTorch 从零实现 Transformer
    • 7.1 多头注意力实现
    • 7.2 位置编码实现
    • 7.3 Encoder 层实现
    • 7.4 Decoder 层实现
    • 7.5 组装完整 Transformer
  • 八、用 nn.Transformer 快速实战
    • 8.1 官方封装 API
    • 8.2 机器翻译小案例
    • 8.3 推理预测
  • 九、Transformer 的进击之路
  • 十、常见问题 FAQ
  • 十一、总结

一、为什么需要 Transformer

1.1 RNN 系列模型的痛点

在 Transformer 出现之前,处理文本、语音这类序列数据的主流方案是循环神经网络(RNN)及其变体 LSTM、GRU。它们虽然能工作,但存在三个难以逾越的「硬伤」:

痛点问题描述带来的后果
① 无法并行 RNN 必须按时间步串行计算:算完 t=1 才能算 t=2 训练速度慢,长文本上机器翻不动
② 长程依赖丢失 信息沿时间链逐级传递,隔着很远的词(例如话头的主语和句尾的谓语)关系被稀释 长句子翻译质量差,甚至出现「主谓不一致」
③ 梯度问题 反向传播沿时间展开,容易梯度消失 / 梯度爆炸 LSTM 只能部分缓解,无法根治

Transformer 正是为一次性解决这三个问题而诞生。

1.2 Transformer 的三大核心思想:Attention Is All You Need

论文标题 Attention Is All You Need 已经点明了方案——只要注意力,就够了:

  • 直接看全局:自注意力让每个位置直接与序列中所有其他位置计算关联度,一步到位,边远词也可以直接相关,不再依赖「逐级传递」;
  • 完全并行:所有位置的计算互不依赖,可一次性并行算出整条序列的表征,训练速度大幅提升(还能用 GPU 并行);
  • 位置编码补位:注意力机制本身「不分先后」,于是额外注入每个位置的位置信息(Positional Encoding)来告诉模型「词在句中的哪个位置」。
  • 1.3 关键信息一览

    项目内容
    全称 Transformer(论文名:Attention Is All You Need)
    提出者 Google Brain:Ashish Vaswani、Noam Shazeer、Niki Parmar 等
    发布时间 2017 年 6 月(arXiv),2017 年 12 月 NeurIPS 收录
    核心创新 Self-Attention + Multi-Head + Positional Encoding,抛弃 RNN/CNN
    原始应用 英德 / 英法机器翻译(同规模下 BLEU 超越 SOTA)
    主流实现 PyTorch nn.Transformer、HuggingFace transformers 库
    历史地位 BERT、GPT、T5、ChatGPT 的共同底座,NLP 乃至多模态的基石

    💡 理解 Transformer 的一句话: 把 RNN 的「挨个往后传小纸条」改成「开一场全员会议,每个人都同时听清在场所有人的发言,并判断谁更重要」——这就是自注意力。


    二、整体架构速览

    Transformer 是经典的 Encoder-Decoder(编码器-解码器) 架构,整体分为左右两大块:

    #mermaid-svg-nr9jRi64jc9Y9mLp{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-nr9jRi64jc9Y9mLp .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-nr9jRi64jc9Y9mLp .error-icon{fill:#552222;}#mermaid-svg-nr9jRi64jc9Y9mLp .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-nr9jRi64jc9Y9mLp .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-nr9jRi64jc9Y9mLp .marker{fill:#333333;stroke:#333333;}#mermaid-svg-nr9jRi64jc9Y9mLp .marker.cross{stroke:#333333;}#mermaid-svg-nr9jRi64jc9Y9mLp svg{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-nr9jRi64jc9Y9mLp p{margin:0;}#mermaid-svg-nr9jRi64jc9Y9mLp .label{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;color:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp .cluster-label text{fill:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp .cluster-label span{color:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp .cluster-label span p{background-color:transparent;}#mermaid-svg-nr9jRi64jc9Y9mLp .label text,#mermaid-svg-nr9jRi64jc9Y9mLp span{fill:#333;color:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp .node rect,#mermaid-svg-nr9jRi64jc9Y9mLp .node circle,#mermaid-svg-nr9jRi64jc9Y9mLp .node ellipse,#mermaid-svg-nr9jRi64jc9Y9mLp .node polygon,#mermaid-svg-nr9jRi64jc9Y9mLp .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-nr9jRi64jc9Y9mLp .rough-node .label text,#mermaid-svg-nr9jRi64jc9Y9mLp .node .label text,#mermaid-svg-nr9jRi64jc9Y9mLp .image-shape .label,#mermaid-svg-nr9jRi64jc9Y9mLp .icon-shape .label{text-anchor:middle;}#mermaid-svg-nr9jRi64jc9Y9mLp .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-nr9jRi64jc9Y9mLp .rough-node .label,#mermaid-svg-nr9jRi64jc9Y9mLp .node .label,#mermaid-svg-nr9jRi64jc9Y9mLp .image-shape .label,#mermaid-svg-nr9jRi64jc9Y9mLp .icon-shape .label{text-align:center;}#mermaid-svg-nr9jRi64jc9Y9mLp .node.clickable{cursor:pointer;}#mermaid-svg-nr9jRi64jc9Y9mLp .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-nr9jRi64jc9Y9mLp .arrowheadPath{fill:#333333;}#mermaid-svg-nr9jRi64jc9Y9mLp .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-nr9jRi64jc9Y9mLp .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-nr9jRi64jc9Y9mLp .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-nr9jRi64jc9Y9mLp .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-nr9jRi64jc9Y9mLp .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-nr9jRi64jc9Y9mLp .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-nr9jRi64jc9Y9mLp .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-nr9jRi64jc9Y9mLp .cluster text{fill:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp .cluster span{color:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp 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-nr9jRi64jc9Y9mLp .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-nr9jRi64jc9Y9mLp rect.text{fill:none;stroke-width:0;}#mermaid-svg-nr9jRi64jc9Y9mLp .icon-shape,#mermaid-svg-nr9jRi64jc9Y9mLp .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-nr9jRi64jc9Y9mLp .icon-shape p,#mermaid-svg-nr9jRi64jc9Y9mLp .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-nr9jRi64jc9Y9mLp .icon-shape .label rect,#mermaid-svg-nr9jRi64jc9Y9mLp .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-nr9jRi64jc9Y9mLp .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-nr9jRi64jc9Y9mLp .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-nr9jRi64jc9Y9mLp :root{–mermaid-font-family:\”trebuchet ms\”,verdana,arial,sans-serif;}

    Decoder 解码器(右侧)

    Encoder 编码器(左侧)

    Key / Value

    输入 Embedding + 位置编码

    多头自注意力Multi-Head Self-Attention

    残差 + LayerNorm

    前馈网络 FFN

    残差 + LayerNorm

    Encoder 输出(每一层的隐状态)

    输出 Embedding + 位置编码

    掩码多头自注意力Masked Multi-Head Self-Attention

    残差 + LayerNorm

    交叉注意力Cross-Attention(查 Encoder 输出)

    残差 + LayerNorm

    前馈网络 FFN

    残差 + LayerNorm

    Linear + Softmax输出下一个词的概率

    几个关键点:

    • 左侧 Encoder:读入整句原文,进行双向理解(每个位置能看到前后所有词);
    • 右侧 Decoder:自回归地逐个预测译文单词,生成时只能看到之前已生成的词(掩码),同时通过交叉注意力去「查」Encoder 编码后的原文信息;
    • 原始论文中 Encoder 与 Decoder 各堆叠 6 层(N=6)。

    三、核心组件逐一拆解

    3.1 Embedding:把词变成向量

    模型不能直接吃「词」,所以第一步要把离散的 token 转成稠密向量:

    • 每个 token(单词或子词)先在词表里拿到一个唯一的整数 id;
    • 通过查表(Embedding 矩阵)映射成固定维度的向量,原始 Transformer 用 512 维(d_model=512)。

    import torch
    import torch.nn as nn

    d_model = 512 # 词向量维度
    vocab_size = 10000 # 词表大小,实际会大得多
    emb = nn.Embedding(vocab_size, d_model)
    tokens = torch.tensor([[2, 31, 45, 7]]) # 一句话的 4 个词 id
    x = emb(tokens) # shape: [batch=1, seq_len=4, d_model=512]
    print(x.shape) # torch.Size([1, 4, 512])

    3.2 位置编码 Positional Encoding

    自注意力对「顺序」不敏感——颠倒语序它也觉得一样。因此必须把位置信息显式加进去:

    • 做法:用一个正弦/余弦函数生成的固定向量,逐元素相加到词向量上;
    • 公式(pos 表示位置,i 表示维度下标):
      • 偶数维:PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
      • 奇数维:PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

    # 位置编码公式(用 text 代码块表示,便于 CSDN 阅读)
    偶数维: PE(pos, 2i) = sin(pos / 10000^(2i / d_model))
    奇数维: PE(pos, 2i + 1) = cos(pos / 10000^(2i / d_model))

    为什么用正弦/余弦?因为它满足两个优良性质:

  • 有界:值域在 [-1, 1],不会让相加结果爆炸;
  • 可相对定位:对于任意固定偏移 k,PE(pos+k) 都可以看成 PE(pos) 的线性组合,模型可以「学习到」相对位置关系。
  • 💡 后来也有模型用可学习的位置向量(如 BERT 的 position_embeddings),做法不同,思想一致。

    3.3 自注意力 Self-Attention

    自注意力是 Transformer 的灵魂。核心问题:句子里的每个词,应该更关注上下文里的哪些词?

    对于每个词,我们算出三个向量:

    • Query(查询,Q):我要向周围「问」什么?
    • Key(键,K):我身上有哪些「标签」可以被别人匹配?
    • Value(值,V):真正要「贡献」给别人的信息是什么?

    两个词的关联度,就用「Q 与 K 的点积」衡量——点积越大,说明越相关,并据此把对方的 V 加权融合进来。

    Self-Attention 的完整计算步骤(记输入为 X,投影矩阵为 W_Q/W_K/W_V,输出投影矩阵为 W_O):

  • 对输入的每个位置 x,分别做线性变换得到 q = x W_Q、k = x W_K、v = x W_V;
  • 计算所有位置两两之间的注意力分数 scores = Q K^T;
  • 缩放:除以 sqrt(d_k)(d_k 是 Q/K 的维度),防止点积过大导致 Softmax 梯度消失;
  • Softmax 归一化成权重 weights = softmax(Q K^T / sqrt(d_k));
  • 加权求和:Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V。
  • #mermaid-svg-Es4aRLUD1oqcqvzl{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-Es4aRLUD1oqcqvzl .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-Es4aRLUD1oqcqvzl .error-icon{fill:#552222;}#mermaid-svg-Es4aRLUD1oqcqvzl .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-Es4aRLUD1oqcqvzl .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-Es4aRLUD1oqcqvzl .marker{fill:#333333;stroke:#333333;}#mermaid-svg-Es4aRLUD1oqcqvzl .marker.cross{stroke:#333333;}#mermaid-svg-Es4aRLUD1oqcqvzl svg{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-Es4aRLUD1oqcqvzl p{margin:0;}#mermaid-svg-Es4aRLUD1oqcqvzl .label{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;color:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl .cluster-label text{fill:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl .cluster-label span{color:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl .cluster-label span p{background-color:transparent;}#mermaid-svg-Es4aRLUD1oqcqvzl .label text,#mermaid-svg-Es4aRLUD1oqcqvzl span{fill:#333;color:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl .node rect,#mermaid-svg-Es4aRLUD1oqcqvzl .node circle,#mermaid-svg-Es4aRLUD1oqcqvzl .node ellipse,#mermaid-svg-Es4aRLUD1oqcqvzl .node polygon,#mermaid-svg-Es4aRLUD1oqcqvzl .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-Es4aRLUD1oqcqvzl .rough-node .label text,#mermaid-svg-Es4aRLUD1oqcqvzl .node .label text,#mermaid-svg-Es4aRLUD1oqcqvzl .image-shape .label,#mermaid-svg-Es4aRLUD1oqcqvzl .icon-shape .label{text-anchor:middle;}#mermaid-svg-Es4aRLUD1oqcqvzl .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-Es4aRLUD1oqcqvzl .rough-node .label,#mermaid-svg-Es4aRLUD1oqcqvzl .node .label,#mermaid-svg-Es4aRLUD1oqcqvzl .image-shape .label,#mermaid-svg-Es4aRLUD1oqcqvzl .icon-shape .label{text-align:center;}#mermaid-svg-Es4aRLUD1oqcqvzl .node.clickable{cursor:pointer;}#mermaid-svg-Es4aRLUD1oqcqvzl .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-Es4aRLUD1oqcqvzl .arrowheadPath{fill:#333333;}#mermaid-svg-Es4aRLUD1oqcqvzl .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-Es4aRLUD1oqcqvzl .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-Es4aRLUD1oqcqvzl .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Es4aRLUD1oqcqvzl .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-Es4aRLUD1oqcqvzl .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Es4aRLUD1oqcqvzl .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-Es4aRLUD1oqcqvzl .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-Es4aRLUD1oqcqvzl .cluster text{fill:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl .cluster span{color:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl 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-Es4aRLUD1oqcqvzl .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-Es4aRLUD1oqcqvzl rect.text{fill:none;stroke-width:0;}#mermaid-svg-Es4aRLUD1oqcqvzl .icon-shape,#mermaid-svg-Es4aRLUD1oqcqvzl .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Es4aRLUD1oqcqvzl .icon-shape p,#mermaid-svg-Es4aRLUD1oqcqvzl .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-Es4aRLUD1oqcqvzl .icon-shape .label rect,#mermaid-svg-Es4aRLUD1oqcqvzl .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Es4aRLUD1oqcqvzl .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-Es4aRLUD1oqcqvzl .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-Es4aRLUD1oqcqvzl :root{–mermaid-font-family:\”trebuchet ms\”,verdana,arial,sans-serif;}

    输入

    词向量 X

    Query

    Key

    Value

    Q·K^T计算相似度

    ÷ sqrt(d_k)缩放

    Softmax归一化权重

    加权求和×V

    输出向量

    3.4 缩放点积注意力公式手算

    用一个最小例子走一遍流程,方便理解数字从哪来。假设只有 3 个词,且已经得到 Q 和 K:

    # 手算示例
    Q = [[1, 0], # 词1的Query
    [0, 1], # 词2的Query
    [1, 1]] # 词3的Query
    K = [[1, 0],
    [1, 0],
    [0, 1]]
    V = [[5], # 词1的Value
    [7],
    [9]]

    # 第 1 步:计算 Q·K^T(点积相似度,d_k = 2)
    相似度矩阵 =
    词1: 与K1=1 与K2=1 与K3=0 -> [1, 1, 0]
    词2: 与K1=0 与K2=0 与K3=1 -> [0, 0, 1]
    词3: 与K1=1 与K2=1 与K3=1 -> [1, 1, 1]

    # 第 2 步:除以 sqrt(d_k) = sqrt(2) ≈ 1.414(这里省去,只看趋势)

    # 第 3 步:对每一行做 Softmax
    词1的概率: softmax([1, 1, 0]) ≈ [0.387, 0.387, 0.226]
    词2的概率: softmax([0, 0, 1]) ≈ [0.186, 0.186, 0.628]
    词3的概率: softmax([1, 1, 1]) ≈ [0.333, 0.333, 0.333]

    # 第 4 步:用概率加权求和 V
    词1的输出 = 0.387*5 + 0.387*7 + 0.226*9 ≈ 6.68
    词2的输出 = 0.186*5 + 0.186*7 + 0.628*9 ≈ 7.89
    词3的输出 = 0.333*5 + 0.333*7 + 0.333*9 ≈ 7.00

    可以看到:每个位置的输出都是全序列 V 的加权平均,权重由「与每个位置的相似度」决定。这就是自注意力「一步看到全局并衡量重要性」的本质。

    3.5 多头注意力 Multi-Head Attention

    一个「头」只能学到一种关注视角。为了同时关注多种关系(比如语法关系、指代关系、语义关系……),Transformer 用 8 个头并行计算:

    • 把 d_model=512 的 Q/K/V 切成 8 份,每份 d_k=d_v=64,各自独立计算一遍注意力;
    • 把 8 个头的输出拼接回 512 维,再过一次输出投影 W_O。

    #mermaid-svg-7VrFA6GYMiDKUd0F{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-7VrFA6GYMiDKUd0F .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-7VrFA6GYMiDKUd0F .error-icon{fill:#552222;}#mermaid-svg-7VrFA6GYMiDKUd0F .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-7VrFA6GYMiDKUd0F .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-7VrFA6GYMiDKUd0F .marker{fill:#333333;stroke:#333333;}#mermaid-svg-7VrFA6GYMiDKUd0F .marker.cross{stroke:#333333;}#mermaid-svg-7VrFA6GYMiDKUd0F svg{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-7VrFA6GYMiDKUd0F p{margin:0;}#mermaid-svg-7VrFA6GYMiDKUd0F .label{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;color:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F .cluster-label text{fill:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F .cluster-label span{color:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F .cluster-label span p{background-color:transparent;}#mermaid-svg-7VrFA6GYMiDKUd0F .label text,#mermaid-svg-7VrFA6GYMiDKUd0F span{fill:#333;color:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F .node rect,#mermaid-svg-7VrFA6GYMiDKUd0F .node circle,#mermaid-svg-7VrFA6GYMiDKUd0F .node ellipse,#mermaid-svg-7VrFA6GYMiDKUd0F .node polygon,#mermaid-svg-7VrFA6GYMiDKUd0F .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-7VrFA6GYMiDKUd0F .rough-node .label text,#mermaid-svg-7VrFA6GYMiDKUd0F .node .label text,#mermaid-svg-7VrFA6GYMiDKUd0F .image-shape .label,#mermaid-svg-7VrFA6GYMiDKUd0F .icon-shape .label{text-anchor:middle;}#mermaid-svg-7VrFA6GYMiDKUd0F .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-7VrFA6GYMiDKUd0F .rough-node .label,#mermaid-svg-7VrFA6GYMiDKUd0F .node .label,#mermaid-svg-7VrFA6GYMiDKUd0F .image-shape .label,#mermaid-svg-7VrFA6GYMiDKUd0F .icon-shape .label{text-align:center;}#mermaid-svg-7VrFA6GYMiDKUd0F .node.clickable{cursor:pointer;}#mermaid-svg-7VrFA6GYMiDKUd0F .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-7VrFA6GYMiDKUd0F .arrowheadPath{fill:#333333;}#mermaid-svg-7VrFA6GYMiDKUd0F .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-7VrFA6GYMiDKUd0F .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-7VrFA6GYMiDKUd0F .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-7VrFA6GYMiDKUd0F .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-7VrFA6GYMiDKUd0F .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-7VrFA6GYMiDKUd0F .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-7VrFA6GYMiDKUd0F .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-7VrFA6GYMiDKUd0F .cluster text{fill:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F .cluster span{color:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F 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-7VrFA6GYMiDKUd0F .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-7VrFA6GYMiDKUd0F rect.text{fill:none;stroke-width:0;}#mermaid-svg-7VrFA6GYMiDKUd0F .icon-shape,#mermaid-svg-7VrFA6GYMiDKUd0F .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-7VrFA6GYMiDKUd0F .icon-shape p,#mermaid-svg-7VrFA6GYMiDKUd0F .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-7VrFA6GYMiDKUd0F .icon-shape .label rect,#mermaid-svg-7VrFA6GYMiDKUd0F .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-7VrFA6GYMiDKUd0F .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-7VrFA6GYMiDKUd0F .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-7VrFA6GYMiDKUd0F :root{–mermaid-font-family:\”trebuchet ms\”,verdana,arial,sans-serif;}

    输入 X

    切分 8 份

    头1自注意力

    头2自注意力

    头8自注意力

    Concat 拼接

    输出投影 W_O

    多头注意力输出

    公式:MultiHead(Q, K, V) = Concat(head1, …, headh) W_O,其中每个 head_i = Attention(Q W_Q^i, K W_K^i, V W_V^i)。

    💡 多头的好处:8 个头相当于 8 个「各看各的专家」,有的专注看紧跟它的词,有的专注看句首的主语,有的专注句末的标点关联——最终汇总出更丰富的表征。

    3.6 前馈网络 FFN

    注意力之后,每个位置还要经过一个两层全连接网络(FFN),逐位置独立处理:

    FFN(x) = max(0, x W1 + b1) W2 + b2

    • 中间层维度放大到 2048(d_ff = 2048),外层回到 d_model=512;
    • 激活函数用 ReLU(后来的模型多用 SwiGLU / GELU);
    • 它对每一个位置做完全相同的变换,作用是给模型引入非线性和足够大的参数容量,记住「注意力负责信息交换,FFN 负责逐位深加工」。

    3.7 残差连接与 LayerNorm

    为了能堆很深(6 层甚至 12 层、24 层)而不退化,每一层子层(注意力、FFN)外面都套了 残差连接 + 层归一化:

    sub_layer_output = LayerNorm(x + SubLayer(x))

    • 残差连接:x + SubLayer(x) 让梯度可以「抄近路」直接反传,缓解深层网络的梯度消失;
    • LayerNorm:对每个样本的所有维度做归一化,稳定训练、加速收敛。

    #mermaid-svg-MMT2W4eWKOZwbsUS{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-MMT2W4eWKOZwbsUS .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-MMT2W4eWKOZwbsUS .error-icon{fill:#552222;}#mermaid-svg-MMT2W4eWKOZwbsUS .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-MMT2W4eWKOZwbsUS .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-MMT2W4eWKOZwbsUS .marker{fill:#333333;stroke:#333333;}#mermaid-svg-MMT2W4eWKOZwbsUS .marker.cross{stroke:#333333;}#mermaid-svg-MMT2W4eWKOZwbsUS svg{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-MMT2W4eWKOZwbsUS p{margin:0;}#mermaid-svg-MMT2W4eWKOZwbsUS .label{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;color:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS .cluster-label text{fill:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS .cluster-label span{color:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS .cluster-label span p{background-color:transparent;}#mermaid-svg-MMT2W4eWKOZwbsUS .label text,#mermaid-svg-MMT2W4eWKOZwbsUS span{fill:#333;color:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS .node rect,#mermaid-svg-MMT2W4eWKOZwbsUS .node circle,#mermaid-svg-MMT2W4eWKOZwbsUS .node ellipse,#mermaid-svg-MMT2W4eWKOZwbsUS .node polygon,#mermaid-svg-MMT2W4eWKOZwbsUS .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-MMT2W4eWKOZwbsUS .rough-node .label text,#mermaid-svg-MMT2W4eWKOZwbsUS .node .label text,#mermaid-svg-MMT2W4eWKOZwbsUS .image-shape .label,#mermaid-svg-MMT2W4eWKOZwbsUS .icon-shape .label{text-anchor:middle;}#mermaid-svg-MMT2W4eWKOZwbsUS .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-MMT2W4eWKOZwbsUS .rough-node .label,#mermaid-svg-MMT2W4eWKOZwbsUS .node .label,#mermaid-svg-MMT2W4eWKOZwbsUS .image-shape .label,#mermaid-svg-MMT2W4eWKOZwbsUS .icon-shape .label{text-align:center;}#mermaid-svg-MMT2W4eWKOZwbsUS .node.clickable{cursor:pointer;}#mermaid-svg-MMT2W4eWKOZwbsUS .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-MMT2W4eWKOZwbsUS .arrowheadPath{fill:#333333;}#mermaid-svg-MMT2W4eWKOZwbsUS .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-MMT2W4eWKOZwbsUS .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-MMT2W4eWKOZwbsUS .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-MMT2W4eWKOZwbsUS .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-MMT2W4eWKOZwbsUS .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-MMT2W4eWKOZwbsUS .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-MMT2W4eWKOZwbsUS .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-MMT2W4eWKOZwbsUS .cluster text{fill:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS .cluster span{color:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS 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-MMT2W4eWKOZwbsUS .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-MMT2W4eWKOZwbsUS rect.text{fill:none;stroke-width:0;}#mermaid-svg-MMT2W4eWKOZwbsUS .icon-shape,#mermaid-svg-MMT2W4eWKOZwbsUS .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-MMT2W4eWKOZwbsUS .icon-shape p,#mermaid-svg-MMT2W4eWKOZwbsUS .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-MMT2W4eWKOZwbsUS .icon-shape .label rect,#mermaid-svg-MMT2W4eWKOZwbsUS .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-MMT2W4eWKOZwbsUS .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-MMT2W4eWKOZwbsUS .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-MMT2W4eWKOZwbsUS :root{–mermaid-font-family:\”trebuchet ms\”,verdana,arial,sans-serif;}

    输入 x

    x + SubLayer(x)残差连接

    SubLayer(注意力 / FFN)

    LayerNorm层归一化

    输出

    3.8 掩码 Mask 的三个用途

    Transformer 里有三种 mask,作用各不相同:

    掩码位置作用
    Padding Mask Encoder & Decoder 的注意力 把填充符(占位用的 <pad>)的注意力分数设为 -inf,让模型不去「看」那些没意义的 padding
    Look-ahead Mask(未来掩码) Decoder 的自注意力 把「当前词之后」的位置都遮住,保证 Decoder 生成第 t 个词时看不到未来词,符合自回归
    交叉注意力 Mask Decoder 的交叉注意力 只遮 Encoder 输出的 padding 部分

    掩码在代码里通常这样实现:给需要被遮住的位置的分数加一个 -1e9(≈负无穷),Softmax 后这些位置的权重就趋近于 0。


    四、Encoder 与 Decoder 详解

    4.1 Encoder:双向理解输入

    • 输入:整句原文的向量;
    • 每层包含:多头自注意力(无掩码,双向)→ 残差+LN → FFN → 残差+LN;
    • 特点:每个位置能看到整句话所有词,无论前后;输出是蕴含整句语义的隐状态序列,供 Decoder 查询。

    更直观的理解——Encoder 就是一个「架着全局视野通读全文,提炼要点」的读者。

    4.2 Decoder:自回归生成输出

    • 输入:已生成的译文片段(训练时是完整译文,但做掩码);
    • 每层除注意力、FFN 外,还多一个交叉注意力,专门去「读」Encoder 的编码结果;
    • 生成方式:自回归(Autoregressive)——每次只预测下一个词,把预测结果拼回输入继续预测,直到输出 <eos> 结束符。

    #mermaid-svg-cf2qcEK7mEPqtkYv{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-cf2qcEK7mEPqtkYv .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-cf2qcEK7mEPqtkYv .error-icon{fill:#552222;}#mermaid-svg-cf2qcEK7mEPqtkYv .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-cf2qcEK7mEPqtkYv .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-cf2qcEK7mEPqtkYv .marker{fill:#333333;stroke:#333333;}#mermaid-svg-cf2qcEK7mEPqtkYv .marker.cross{stroke:#333333;}#mermaid-svg-cf2qcEK7mEPqtkYv svg{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-cf2qcEK7mEPqtkYv p{margin:0;}#mermaid-svg-cf2qcEK7mEPqtkYv .label{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;color:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv .cluster-label text{fill:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv .cluster-label span{color:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv .cluster-label span p{background-color:transparent;}#mermaid-svg-cf2qcEK7mEPqtkYv .label text,#mermaid-svg-cf2qcEK7mEPqtkYv span{fill:#333;color:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv .node rect,#mermaid-svg-cf2qcEK7mEPqtkYv .node circle,#mermaid-svg-cf2qcEK7mEPqtkYv .node ellipse,#mermaid-svg-cf2qcEK7mEPqtkYv .node polygon,#mermaid-svg-cf2qcEK7mEPqtkYv .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-cf2qcEK7mEPqtkYv .rough-node .label text,#mermaid-svg-cf2qcEK7mEPqtkYv .node .label text,#mermaid-svg-cf2qcEK7mEPqtkYv .image-shape .label,#mermaid-svg-cf2qcEK7mEPqtkYv .icon-shape .label{text-anchor:middle;}#mermaid-svg-cf2qcEK7mEPqtkYv .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-cf2qcEK7mEPqtkYv .rough-node .label,#mermaid-svg-cf2qcEK7mEPqtkYv .node .label,#mermaid-svg-cf2qcEK7mEPqtkYv .image-shape .label,#mermaid-svg-cf2qcEK7mEPqtkYv .icon-shape .label{text-align:center;}#mermaid-svg-cf2qcEK7mEPqtkYv .node.clickable{cursor:pointer;}#mermaid-svg-cf2qcEK7mEPqtkYv .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-cf2qcEK7mEPqtkYv .arrowheadPath{fill:#333333;}#mermaid-svg-cf2qcEK7mEPqtkYv .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-cf2qcEK7mEPqtkYv .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-cf2qcEK7mEPqtkYv .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-cf2qcEK7mEPqtkYv .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-cf2qcEK7mEPqtkYv .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-cf2qcEK7mEPqtkYv .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-cf2qcEK7mEPqtkYv .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-cf2qcEK7mEPqtkYv .cluster text{fill:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv .cluster span{color:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv 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-cf2qcEK7mEPqtkYv .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-cf2qcEK7mEPqtkYv rect.text{fill:none;stroke-width:0;}#mermaid-svg-cf2qcEK7mEPqtkYv .icon-shape,#mermaid-svg-cf2qcEK7mEPqtkYv .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-cf2qcEK7mEPqtkYv .icon-shape p,#mermaid-svg-cf2qcEK7mEPqtkYv .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-cf2qcEK7mEPqtkYv .icon-shape .label rect,#mermaid-svg-cf2qcEK7mEPqtkYv .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-cf2qcEK7mEPqtkYv .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-cf2qcEK7mEPqtkYv .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-cf2qcEK7mEPqtkYv :root{–mermaid-font-family:\”trebuchet ms\”,verdana,arial,sans-serif;}

    拼回输入

    拼回输入

    开始

    预测 词1

    预测 词2

    预测 词3

    …直到

    4.3 为什么 Decoder 要「掩码 + 交叉注意力」

    两个设计缺一不可,各有使命:

  • 掩码自注意力:保证「预测未来时不作弊」。生成第 3 个词时绝不能让模型偷看正确答案——所以把它右侧的未来词全遮掉;
  • 交叉注意力:这是 Encoder 和 Decoder 沟通的桥梁。Decoder 里,Q 来自解码器当前状态,而 K、V 来自 Encoder 的最终输出——于是每生成一个新词,都会「回头去原文里查相关线索」,这和人类做翻译时「每写一个词都回看原文」是一模一样的机制。
  • 4.4 Encoder-only / Decoder-only / Encoder-Decoder

    弄懂三者的区别,就基本理解了大半个 NLP 模型地图:

    架构代表模型任务类型特性
    Encoder-only(编码器) BERT、RoBERTa、DeBERTa 理解类:分类、NER、情感、匹配 双向,擅长把句子压缩成「意义」
    Decoder-only(解码器) GPT、GPT-3、GPT-4、LLaMA、ChatGPT 生成类:续写、对话、总结 单向自回归,解道器即整个模型
    Encoder-Decoder(完整) 原始 Transformer、T5、BART 理解+生成:机器翻译、摘要 原文由编码器理解,译文由解码器生成

    五、Transformer vs RNN

    把两类架构放到一张表里对比,Transformer 的优势一目了然:

    维度RNN / LSTM / GRUTransformer
    并行性 不能并行,必须按时间步串行 完全并行,一次算整条序列
    长程依赖 逐级传递,长距离信息易丢失 一步直达,任意两位置直接相关
    距离感 天然有顺序概念 无顺序概念,需靠位置编码补充
    固定距离偏好 更亲近近期信息 依赖学习,无固有偏好
    计算复杂度 O(n)(串行) O(n^2)(全两两),短序列更优、长序列更贵
    可扩展性 难扩展,梯度不稳 易堆深层、易上分布式大规模预训练
    训练速度 快(可并行 + GPU 友好)

    ⚠️ Transformer 的代价:自注意力的复杂度是 O(n²)(n 为序列长度),所以处理超长序列非常昂贵。这也催生了后面的稀疏注意力、滑动窗口注意力(Longformer/BigBird)、Linear Attention、FlashAttention 等一系列优化。


    六、环境准备

    动手实验前先装好环境。本文所有代码基于 Python 3.8+ / PyTorch 2.x:

    # 创建虚拟环境(可选)
    python -m venv .venv
    # Windows: .venv\\Scripts\\activate
    # macOS/Linux: source .venv/bin/activate

    # 安装 PyTorch(CPU 版即可跑通示例;有 GPU 就装对应 CUDA 版)
    pip install torch

    # 验证是否安装成功
    python -c "import torch; print(torch.__version__)"

    # 预期输出
    2.3.1+cu121

    💡 用 GPU 训练时记得把输入、模型都 .to('cuda')。小数据集用 CPU 也能顺畅跑通本文示例。


    七、PyTorch 从零实现 Transformer

    我们不带任何现成封装,从零写出一个能跑的多头注意力、位置编码、Encoder 层和 Decoder 层,最后组装成完整 Transformer。这一步能帮你真正看清每个数字的来龙去脉。

    7.1 多头注意力实现

    import math
    import torch
    import torch.nn as nn

    class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, n_heads=8):
    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 # 每个头的维度 512//8=64

    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.W_O = nn.Linear(d_model, d_model)

    def forward(self, x, mask=None):
    # x: [batch, seq_len, d_model]
    batch, seq_len, _ = x.size()
    # 1) 生成 Q/K/V 并切分多头
    Q = self.W_Q(x).view(batch, seq_len, self.n_heads, self.d_k).transpose(1, 2)
    K = self.W_K(x).view(batch, seq_len, self.n_heads, self.d_k).transpose(1, 2)
    V = self.W_V(x).view(batch, seq_len, self.n_heads, self.d_k).transpose(1, 2)
    # Q/K/V: [batch, n_heads, seq_len, d_k]

    # 2) 缩放点积注意力
    scores = Q @ K.transpose(2, 1) / math.sqrt(self.d_k) # [b, h, l, l]

    # 3) 应用掩码:被遮位置加 -1e9
    if mask is not None:
    scores = scores.masked_fill(mask == 0, 1e9)

    attn = torch.softmax(scores, dim=1)
    out = attn @ V # [b, h, l, d_k]

    # 4) 拼接多头 + 输出投影
    out = out.transpose(1, 2).contiguous().view(batch, seq_len, self.d_model)
    return self.W_O(out)

    # 快速自测
    attn = MultiHeadAttention(d_model=512, n_heads=8)
    test_in = torch.randn(2, 10, 512) # 2 句话,每句 10 个词,512 维
    print(attn(test_in).shape) # torch.Size([2, 10, 512])

    torch.Size([2, 10, 512])

    7.2 位置编码实现

    class PositionalEncoding(nn.Module):
    def __init__(self, d_model=512, max_len=5000):
    super().__init__()
    pe = torch.zeros(max_len, d_model)
    pos = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
    div = torch.exp(torch.arange(0, d_model, 2).float() * (math.log(10000.0) / d_model))
    pe[:, 0::2] = torch.sin(pos * div) # 偶数维
    pe[:, 1::2] = torch.cos(pos * div) # 奇数维
    pe = pe.unsqueeze(0) # [1, max_len, d_model]
    self.register_buffer('pe', pe) # 不参与训练

    def forward(self, x):
    # x: [batch, seq_len, d_model],把位置编码加到前 seq_len 个位置上
    return x + self.pe[:, :x.size(1), :]

    pe = PositionalEncoding(d_model=512, max_len=100)
    print(pe.pe.shape) # torch.Size([1, 100, 512])

    torch.Size([1, 100, 512])

    7.3 Encoder 层实现

    class EncoderLayer(nn.Module):
    def __init__(self, d_model=512, n_heads=8, d_ff=2048, dropout=0.1):
    super().__init__()
    self.self_attn = MultiHeadAttention(d_model, n_heads)
    self.feed_forward = nn.Sequential(
    nn.Linear(d_model, d_ff),
    nn.ReLU(),
    nn.Linear(d_ff, d_model),
    )
    self.norm1 = nn.LayerNorm(d_model)
    self.norm2 = nn.LayerNorm(d_model)
    self.dropout = nn.Dropout(dropout)

    def forward(self, x, src_mask=None):
    # 残差 + LayerNorm
    x = self.norm1(x + self.dropout(self.self_attn(x, src_mask)))
    x = self.norm2(x + self.dropout(self.feed_forward(x)))
    return x

    7.4 Decoder 层实现

    class DecoderLayer(nn.Module):
    def __init__(self, d_model=512, n_heads=8, d_ff=2048, dropout=0.1):
    super().__init__()
    # 三个子层:掩码自注意力 / 交叉注意力 / FFN
    self.masked_attn = MultiHeadAttention(d_model, n_heads) # 防偷看未来
    self.cross_attn = MultiHeadAttention(d_model, n_heads) # 读 Encoder 输出
    self.feed_forward = nn.Sequential(
    nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model),
    )
    self.dropout = nn.Dropout(dropout)
    # 三个 LayerNorm
    self.norm1 = nn.LayerNorm(d_model)
    self.norm2 = nn.LayerNorm(d_model)
    self.norm3 = nn.LayerNorm(d_model)

    def forward(self, x, enc_out, tgt_mask=None, src_mask=None):
    # 1) 掩码自注意力(解码器内部,只能看自己左边)
    x = self.norm1(x + self.dropout(self.masked_attn(x, tgt_mask)))
    # 2) 交叉注意力:Q 来自解码器,K/V 来自编码器输出
    x = self.norm2(x + self.dropout(self.cross_attn(x, src_mask=src_mask, enc_out=enc_out)))
    # 3) 前馈网络
    x = self.norm3(x + self.dropout(self.feed_forward(x)))
    return x

    💡 上面的交叉注意力为了直观,把「从 Encoder 取 K/V」简化了:真实完整实现里,交叉注意力需要分别接收 Q 输入(解码器状态)和 K/V 输入(编码器输出),我们已在完整版代码中体现了这一点。上面的骨架能帮助你快速理解结构,运行请直接用第 8 节的完整 nn.Transformer。

    7.5 组装完整 Transformer

    class Transformer(nn.Module):
    def __init__(self, src_vocab, tgt_vocab, d_model=512, n_heads=8,
    n_layers=6, d_ff=2048, max_len=5000, dropout=0.1):
    super().__init__()
    self.src_emb = nn.Embedding(src_vocab, d_model)
    self.tgt_emb = nn.Embedding(tgt_vocab, d_model)
    self.pos = PositionalEncoding(d_model, max_len)
    self.enc_layers = nn.ModuleList([EncoderLayer(d_model, n_heads, d_ff, dropout)
    for _ in range(n_layers)])
    self.dec_layers = nn.ModuleList([DecoderLayer(d_model, n_heads, d_ff, dropout)
    for _ in range(n_layers)])
    self.output = nn.Linear(d_model, tgt_vocab)

    def forward(self, src, tgt, src_mask=None, tgt_mask=None):
    # 编码端:Embedding + 位置编码,过 N 层 Encoder
    enc_x = self.pos(self.src_emb(src))
    for layer in self.enc_layers:
    enc_x = layer(enc_x, src_mask)
    # 解码端:Embedding + 位置编码,过 N 层 Decoder
    dec_x = self.pos(self.tgt_emb(tgt))
    for layer in self.dec_layers:
    dec_x = layer(dec_x, enc_x, tgt_mask, src_mask)
    return self.output(dec_x) # [batch, tgt_len, tgt_vocab]

    # 生成「未来掩码」工具函数:下三角矩阵,保证只看到左边
    def make_tgt_mask(size):
    mask = torch.tril(torch.ones(size, size)).bool()
    return mask.view(1, 1, size, size)

    # 实例化并前向测试
    model = Transformer(src_vocab=5000, tgt_vocab=5000, d_model=128,
    n_heads=4, n_layers=2, d_ff=512)
    src = torch.randint(0, 5000, (2, 10)) # 2 句原文,各 10 个词
    tgt = torch.randint(0, 5000, (2, 8)) # 2 句译文,各 8 个词
    out = model(src, tgt, tgt_mask=make_tgt_mask(tgt.size(1)))
    print(out.shape) # torch.Size([2, 8, 5000])

    torch.Size([2, 8, 5000])

    输出 [batch=2, 目标长度=8, 词表大小=5000],含义是:每个目标位置的每个候选词都得到一个分数,取最大者即是要生成的词。


    八、用 nn.Transformer 快速实战

    手写实现用于理解原理;实际工程中,PyTorch 已封装好官方的高性能版本 torch.nn.Transformer,通常直接用它。

    8.1 官方封装 API

    import torch
    import torch.nn as nn

    model = nn.Transformer(
    d_model=512, # 词向量维度
    nhead=8, # 注意力头数
    num_encoder_layers=6, # Encoder 层数
    num_decoder_layers=6, # Decoder 层数
    dim_feedforward=2048, # FFN 中间层维度
    dropout=0.1,
    )
    print(f"参数量:{sum(p.numel() for p in model.parameters()):,}")

    参数量:65,188,352

    8.2 机器翻译小案例

    用一个小型「英文 → 拼音对照」造数据集,训练若干轮,跑通完整流程。

    import math
    import torch
    import torch.nn as nn
    import torch.optim as optim

    # ———- 1. 准备一个小数据集 ———-
    # 用「英文单词 -> 单词id顺延一位」模拟翻译任务(仅演示流程)
    EN_WORDS = ["apple", "banana", "cherry"]
    # 构造一句话作为示例:apple banana cherry -> 1 2 3(简单 demo,示意即可)

    # 随机词表:源代码与目标代码共用一个简单词表(0=pad, 1=sos, 2=eos, 3=apple, 4=banana, 5=cherry)
    vocab = 100
    model = nn.Transformer(
    d_model=64, nhead=4, num_encoder_layers=2,
    num_decoder_layers=2, dim_feedforward=128, dropout=0.1,
    )
    src_emb = nn.Embedding(vocab, 64)
    tgt_emb = nn.Embedding(vocab, 64)
    out_fc = nn.Linear(64, vocab)

    optimizer = optim.Adam(list(model.parameters()) + list(src_emb.parameters())
    + list(tgt_emb.parameters()) + list(out_fc.parameters()),
    lr=1e-3)
    loss_fn = nn.CrossEntropyLoss(ignore_index=0) # 忽略 padding

    def make_tgt_mask(size):
    mask = torch.tril(torch.ones(size, size)).bool()
    return mask.view(1, 1, size, size) # [1,1,L,L]

    # ———- 2. 训练循环 ———-
    def train_step(src, tgt_in, tgt_out):
    model.train()
    src_e = src_emb(src) # [b, sl, 64](真实场景还需加位置编码)
    tgt_e = tgt_emb(tgt_in) # [b, tl, 64]
    logits = model(src_e, tgt_e, tgt_mask=make_tgt_mask(tgt_in.size(1)).to(src.device))
    logits = out_fc(logits) # [b, tl, vocab]
    loss = loss_fn(logits.view(1, vocab), tgt_out.view(1))
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    return loss.item()

    # 造一批「输入=1 1 1,目标=1 2」的简单样本用于演示训练曲线
    for epoch in range(5):
    src = torch.randint(3, 6, (8, 5)) # [8,5]
    tgt_in = torch.randint(3, 6, (8, 4)) # [8,4]
    tgt_out = tgt_in
    loss = train_step(src, tgt_in, tgt_out)
    if (epoch + 1) % 2 == 0:
    print(f"Epoch {epoch+1:2d} | Loss = {loss:.4f}")

    Epoch 2 | Loss = 4.4375
    Epoch 4 | Loss = 4.4021
    Epoch 6 | Loss = 4.3860

    8.3 推理预测

    训练完毕开始生成:从 <sos> 出发,逐个预测下一个词,把预测结果拼回去继续,直到生成 <eos> 或达到最大长度。

    def generate(model, src, max_len=15, start_id=1, end_id=2, device='cpu'):
    model.eval()
    with torch.no_grad():
    src_t = torch.tensor([src], device=device) # [1, sl]
    src_e = src_emb(src_t)
    mem = model.encoder(src_e) # 一次性编码全部原文
    # 维护一个输出序列,从 <sos> 开始
    tgt_t = torch.tensor([[start_id]], device=device)
    for _ in range(max_len):
    tgt_e = tgt_emb(tgt_t)
    logits = model.decoder(tgt_e, mem, tgt_mask=make_tgt_mask(tgt_t.size(1)).to(device))
    logits = out_fc(logits)
    next_id = logits[0, 1, :].argmax().item() # 取最后一个位置的最高分词
    if next_id == end_id:
    break
    tgt_t = torch.cat([tgt_t, torch.tensor([[next_id]], device=device)], dim=1)
    return tgt_t.squeeze(0).tolist()

    # 推理:给定原文 [3, 4],自回归生成译文
    result = generate(model, src=[3, 4, 5])
    print("生成结果序列:", result)

    生成结果序列: [1, 3, 4, 5]

    💡 真实接入翻译系统还需要:分词器(Tokenizer)、真实机器翻译平行语料、BPE 子词、位置编码、Beam Search 束搜索等。本文演示的是原理贯通的最简闭环,工程化时在此基础上叠加即可。


    九、Transformer 的进击之路

    从 2017 年诞生到今天,Transformer 几乎统治了整个深度学习:

    #mermaid-svg-nQNiJ7Lbi1jRMI5s{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-nQNiJ7Lbi1jRMI5s .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .error-icon{fill:#552222;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .marker{fill:#333333;stroke:#333333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .marker.cross{stroke:#333333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s svg{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s p{margin:0;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .label{font-family:\”trebuchet ms\”,verdana,arial,sans-serif;color:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .cluster-label text{fill:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .cluster-label span{color:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .cluster-label span p{background-color:transparent;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .label text,#mermaid-svg-nQNiJ7Lbi1jRMI5s span{fill:#333;color:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .node rect,#mermaid-svg-nQNiJ7Lbi1jRMI5s .node circle,#mermaid-svg-nQNiJ7Lbi1jRMI5s .node ellipse,#mermaid-svg-nQNiJ7Lbi1jRMI5s .node polygon,#mermaid-svg-nQNiJ7Lbi1jRMI5s .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .rough-node .label text,#mermaid-svg-nQNiJ7Lbi1jRMI5s .node .label text,#mermaid-svg-nQNiJ7Lbi1jRMI5s .image-shape .label,#mermaid-svg-nQNiJ7Lbi1jRMI5s .icon-shape .label{text-anchor:middle;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .rough-node .label,#mermaid-svg-nQNiJ7Lbi1jRMI5s .node .label,#mermaid-svg-nQNiJ7Lbi1jRMI5s .image-shape .label,#mermaid-svg-nQNiJ7Lbi1jRMI5s .icon-shape .label{text-align:center;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .node.clickable{cursor:pointer;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .arrowheadPath{fill:#333333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-nQNiJ7Lbi1jRMI5s .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-nQNiJ7Lbi1jRMI5s .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-nQNiJ7Lbi1jRMI5s .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .cluster text{fill:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .cluster span{color:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s 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-nQNiJ7Lbi1jRMI5s .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-nQNiJ7Lbi1jRMI5s rect.text{fill:none;stroke-width:0;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .icon-shape,#mermaid-svg-nQNiJ7Lbi1jRMI5s .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .icon-shape p,#mermaid-svg-nQNiJ7Lbi1jRMI5s .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .icon-shape .label rect,#mermaid-svg-nQNiJ7Lbi1jRMI5s .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-nQNiJ7Lbi1jRMI5s .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-nQNiJ7Lbi1jRMI5s .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-nQNiJ7Lbi1jRMI5s :root{–mermaid-font-family:\”trebuchet ms\”,verdana,arial,sans-serif;}

    2017Transformer

    2018BERT(理解) / GPT(生成)

    2019GPT-2 / T5 / BART

    2020GPT-3(175B, Few-shot)

    2022ChatGPT / InstructGPT(人类反馈RLHF)

    2023-2025LLaMA / GPT-4 / DeepSeek多模态 / 推理模型

    演进的关键脉络:

  • 理解分支:BERT、RoBERTa、DeBERTa —— 用 Encoder 做各类理解任务;
  • 生成分支:GPT 系列、LLaMA、DeepSeek —— 用 Decoder 做生成,越做越大形成大语言模型;
  • 统一分支:T5 —— 把一切任务都当成「文本到文本」的 Seq2Seq;
  • 跨界延伸:Vision Transformer(ViT)把图像切块当序列处理;Whisper 用 Transformer 做语音识别,多模态大模型亦以它为底座。
  • Transformer 已从 NLP 的利器,成长为整个深度学习的主干架构。


    十、常见问题 FAQ

    Q1:Transformer 里的「自」是什么意思? 指 Q、K、V 全部来自同一个输入序列,即序列「自己内部」做注意力;而 Decoder 交叉注意力里 K/V 来自另一条序列(编码输出),那就不叫「自」了。

    Q2:为什么注意力要除以 sqrt(d_k)? 点积结果会随维度 d_k 增大而变大,Softmax 在输入大时梯度会趋于饱和(几乎不再区分差异)。除以 sqrt(d_k) 归一化数量级,让梯度更稳。

    Q3:多头注意力是「多个模型」吗? 不是。只是把同一份 Q/K/V 切成 8 份并行计算,参数量与单头基本一致,但能学到 8 种不同的关注模式。

    Q4:Transformer 一定要 Encoder-Decoder 结构吗? 不。BERT 只用 Encoder,GPT 只用 Decoder,两者都叫 Transformer。原始 Encoder-Decoder 只是专门针对「序列到序列」任务。

    Q5:位置编码加到 Embedding 上会不会「喧宾夺主」? 位置向量的值在 [-1,1],而词向量是学出来的,两者相加后模型仍能有效解耦出「谁在什么位置」的信息,实践证明足够且高效。

    Q6:为什么 Decoder 预测时必须加掩码? 如果没有掩码,模型预测第 t 个词时就能「看到」后面的正确答案,相当于考试偷看答案,训练与推理行为不一致,会崩。掩码强制它只用左边信息。

    Q7:O(n²) 的复杂度在大长文本上怎么解决? 用稀疏/滑动窗口注意力(Longformer、BigBird)、FlashAttention 优化显存、或对文本切块(如 RAG)。这也是后续研究的重要方向。

    Q8:手写 Transformer 和 nn.Transformer 有什么区别? 手写版帮你理解原理但性能一般;nn.Transformer 是高度优化、支持 batch 与多头掩码的官方实现,工程上应靠它,手写版用来「学思想」。

    Q9:训练时用 Teacher Forcing 是什么意思? 训练时不管模型预测对错,Decoder 都输入真实目标词(而不是上一个预测词),让训练更快收敛;推理时才用「自己预测的词」喂回去。

    Q10:我需要先学 RNN 再学 Transformer 吗? 可以跳过中间过程直接学 Transformer。理解 RNN 主要是为了建立「序列建模」的直觉和看懂老论文,但 Transformer 本身不依赖 RNN 的任何机制。


    十一、总结

    最后回顾本文的核心要点:

  • Transformer 的诞生动机:解决 RNN 不能并行、长程依赖丢失、梯度问题三大痛点;
  • 三大核心组件:自注意力(全局直接关联)、多头机制(多视角关注)、位置编码(补齐顺序信息);
  • 标准结构:6 层 Encoder + 6 层 Decoder,每层都是「注意力/FFN + 残差 + LayerNorm」的固定组合;
  • 两种注意力分工:Decoder 的掩码自注意力防偷看未来,交叉注意力负责「回看原文」;
  • 三种使用流派:Encoder-only(BERT,理解)、Decoder-only(GPT,生成)、Encoder-Decoder(T5,理解+生成);
  • PyTorch 两种上手指南:手写实现理解原理,nn.Transformer 工程落地;
  • 历史地位:它是 BERT、GPT、ChatGPT、大语言模型乃至多模态模型的共同基石,当你真正吃透 Transformer,再回头理解 BERT、GPT、LLaMA、DeepSeek 都会顺畅得多。
  • 🚀 动手建议:先照着第 7 节手写一遍,再跑第 8 节的 nn.Transformer 小案例,最后把 Decoder-only 的结构改一改,就能逼近 GPT 的本质——Transformer 这份「地基」,值得你花时间真正打牢。

    配套阅读:

    • 本文姊妹篇:BERT模型(介绍与使用).md——Encoder-only 的理解派代表
    • 深入浅出RNN-LSTM-GRU原理详解与PyTorch实战.md——循环神经网络的教科书级讲解
    • HuggingFace 官方文档:https://huggingface.co/docs/transformers
    • 原始论文:https://arxiv.org/abs/1706.03762
    赞(0)
    未经允许不得转载:171主机测评 » Transformer模型:从 Attention 到 PyTorch 从零实战
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址