Transformer API 详解笔记

基于 ch06_transformer/2_api_test.ipynb 的逐节讲解,用通俗易懂的方式讲清楚 PyTorch nn.Transformer 的使用方法。


目录

  1. 构造参数
  2. 前向传播(最简形式)
  3. 掩码(Mask)机制
  4. 带掩码的完整前向传播
  5. 编码器单独使用
  6. 解码器单独使用

1. 构造参数

导入库

import torch
from torch import nn
  • torch:PyTorch 核心库,提供张量运算、GPU 加速等功能
  • torch.nn:神经网络模块,包含 Transformer、Embedding 等预定义层

创建 Transformer 模型

model = nn.Transformer(
    d_model=64,
    nhead=8,
    num_encoder_layers=2,
    num_decoder_layers=2,
    dim_feedforward=256,
    batch_first=True,
)

参数详解

参数 含义
d_model 64 模型中所有层的输入/输出特征维度,可以理解为每个词用64个数字来表示
nhead 8 多头注意力的头数,将64维分成8个子空间,每个子空间8维(64÷8=8)
num_encoder_layers 2 编码器中 Transformer 层的数量(堆叠2层)
num_decoder_layers 2 解码器中 Transformer 层的数量(堆叠2层)
dim_feedforward 256 前馈神经网络的隐藏层维度,通常是 d_model 的4倍(64×4=256)
batch_first True 输入输出张量的批次维度在第0位,即 (batch, seq_len, d_model)

打个比方

把 Transformer 想象成一个翻译工厂:

  • d_model=64:工厂里每个工人处理的信息宽度是64个数字
  • nhead=8:有8个小组同时从不同角度分析句子(多头注意力)
  • num_encoder_layers=2:编码车间有2道工序
  • num_decoder_layers=2:解码车间有2道工序
  • dim_feedforward=256:每道工序中,工人先把信息扩展到256维进行思考,再压缩回64维
  • batch_first=True:材料的摆放顺序是「批次 × 序列长度 × 特征」

2. 前向传播(最简形式)

定义超参数

d_model = 64       # 特征维度,与模型定义一致
batch_size = 32    # 批次大小:一次处理32个样本
src_len = 12       # 源序列长度:如英文句子有12个词
tgt_len = 7        # 目标序列长度:如中文句子有7个字

定义词表大小

src_vocab_size = 1000   # 源语言词表大小(如英文有1000个不同的词)
tgt_vocab_size = 1500   # 目标语言词表大小(如中文有1500个不同的字)

生成随机输入序列

src_seqs = torch.randint(src_vocab_size, (batch_size, src_len))  # shape: (32, 12)
tgt_seqs = torch.randint(tgt_vocab_size, (batch_size, tgt_len))  # shape: (32, 7)
  • torch.randint(high, size):生成 [0, high) 范围内的随机整数
  • src_seqs:32个样本,每个样本12个词ID,模拟英文输入
  • tgt_seqs:32个样本,每个样本7个字ID,模拟中文目标

定义词嵌入层

src_embedding = nn.Embedding(num_embeddings=src_vocab_size, embedding_dim=d_model)  # (1000, 64)
tgt_embedding = nn.Embedding(num_embeddings=tgt_vocab_size, embedding_dim=d_model)  # (1500, 64)
  • nn.Embedding 本质上是一张查表:输入词ID,输出对应的向量
  • src_embedding:一张 1000×64 的表,每个英文词对应一个64维向量
  • tgt_embedding:一张 1500×64 的表,每个中文词对应一个64维向量

词嵌入转换

src = src_embedding(src_seqs)  # shape: (32, 12, 64)
tgt = tgt_embedding(tgt_seqs)  # shape: (32, 7, 64)

转换过程:

词ID序列 (32, 12)  →  查Embedding表  →  词向量序列 (32, 12, 64)

每个词ID被替换成一个64维的向量,所以序列从「32行12列的整数」变成了「32×12×64的浮点数」。

最简前向传播

output = model(src, tgt)  # shape: (32, 7, 64)

内部发生了什么:

src (32, 12, 64)  →  编码器  →  memory (32, 12, 64)
                                    ↓
tgt (32, 7, 64)  →  解码器(memory)  →  output (32, 7, 64)
  1. 编码器:对源序列 src 进行双向自注意力计算,输出 memory
  2. 解码器:对目标序列 tgt 进行自注意力 + 交叉注意力(以 memory 为条件),输出 output

输出形状 (32, 7, 64) 的含义:32个样本,每个样本7个位置,每个位置64维向量。后续加一个线性层就能映射到词表大小,得到每个位置的词概率分布。


3. 掩码(Mask)机制

为什么需要掩码?

在实际翻译任务中,不同句子长度不同,短句子需要填充(padding)到相同长度。如果不告诉模型哪些是填充的,模型会把填充位置当作有效信息来处理,干扰注意力计算。

3.1 Padding Mask(填充掩码)

padding_id = 0  # 填充位置的ID为0

src_key_padding_mask = (src_seqs == padding_id)  # shape: (32, 12)
tgt_key_padding_mask = (tgt_seqs == padding_id)  # shape: (32, 7)
memory_key_padding_mask = src_key_padding_mask     # shape: (32, 12)
原理
# src_seqs 中每个值与 0 比较:
# 等于 0 → True(是填充,应该被忽略)
# 不等于 0 → False(是有效词)
示例

假设一个批次中某条数据的 src_seqs[0] 是:

[256, 89, 423, 7, 0, 0, 0, 0, 0, 0, 0, 0]
 ↑有效 ↑有效 ↑有效 ↑有效 ↑填充 ↑填充 ...

对应的 src_key_padding_mask[0] 是:

[False, False, False, False, True, True, True, True, True, True, True, True]
三种 padding mask
掩码 形状 用在哪 含义
src_key_padding_mask (32, 12) 编码器自注意力 忽略源序列中的填充位置
tgt_key_padding_mask (32, 7) 解码器自注意力 忽略目标序列中的填充位置
memory_key_padding_mask (32, 12) 解码器交叉注意力 忽略 memory 中对应源序列填充的位置(与 src 一致)

3.2 Causal Mask(因果掩码 / 后续掩码)

tgt_mask = model.generate_square_subsequent_mask(tgt_len).bool()
# shape: (7, 7)
原理

因果掩码是一个上三角矩阵,保证解码器在预测第 i 个词时,只能看到第 0 到第 i-1 个词,不能看到未来的词。

tensor([[False,  True,  True,  True,  True,  True,  True],
        [False, False,  True,  True,  True,  True,  True],
        [False, False, False,  True,  True,  True,  True],
        [False, False, False, False,  True,  True,  True],
        [False, False, False, False, False,  True,  True],
        [False, False, False, False, False, False,  True],
        [False, False, False, False, False, False, False]])
  • False = 可以看到(当前位置及之前的位置)
  • True = 不能看到(未来位置,被遮挡)
解读
位置 能看到的位置 被遮挡的位置
第0个词 只能看自己 第1~6个词
第1个词 第0~1个词 第2~6个词
第2个词 第0~2个词 第3~6个词
第6个词 第0~6个词(全部)
为什么需要因果掩码?

训练时,我们一次性输入整个目标序列。如果没有因果掩码,预测第3个词时模型就能偷看到第4、5、6个词的答案,这叫信息泄露

推理时,我们是逐个词生成的(自回归),所以天然不会看到未来。但训练时为了效率一次性并行处理,必须用因果掩码来模拟这个逐词生成的过程。


4. 带掩码的完整前向传播

output = model(
    src,                              # 源序列嵌入 (32, 12, 64)
    tgt,                              # 目标序列嵌入 (32, 7, 64)
    tgt_mask=tgt_mask,                # 因果掩码 (7, 7)
    src_key_padding_mask=src_key_padding_mask,            # 源序列填充掩码 (32, 12)
    tgt_key_padding_mask=tgt_key_padding_mask,            # 目标序列填充掩码 (32, 7)
    memory_key_padding_mask=memory_key_padding_mask,      # memory填充掩码 (32, 12)
)
print(output.shape)  # torch.Size([32, 7, 64])

内部流程图

src (32, 12, 64)
    │
    ▼ 编码器(2层 Transformer Encoder Layer)
    │ 每层包含:
    │   1. 多头自注意力(用 src_key_padding_mask 忽略填充)
    │   2. 前馈神经网络
    │
memory (32, 12, 64)
    │
    │          ┌──────────────────────────────────┐
    ▼          │  tgt (32, 7, 64)                 │
    │          │    │                             │
    │          │    ▼ 解码器(2层 Transformer Decoder Layer)│
    │          │    │ 每层包含:                    │
    │          │    │   1. 带因果掩码的自注意力     │
    │          │    │      (tgt_mask +             │
    │          │    │       tgt_key_padding_mask)  │
    │          │    │   2. 交叉注意力(memory)       │
    │          │    │      (memory_key_padding_mask)│
    │          │    │   3. 前馈神经网络             │
    │          │    ▼                              │
    │          │  output (32, 7, 64)              │
    │          └──────────────────────────────────┘

掩码总结

掩码参数 形状 类型 作用
tgt_mask (7, 7) 布尔型,上三角True 防止解码器看到未来词
src_key_padding_mask (32, 12) 布尔型,填充位置True 编码器忽略源序列填充
tgt_key_padding_mask (32, 7) 布尔型,填充位置True 解码器忽略目标序列填充
memory_key_padding_mask (32, 12) 布尔型,填充位置True 交叉注意力忽略源序列填充

5. 编码器单独使用

有时我们只需要编码器部分(如 BERT 只有编码器,或提取句子表示)。

基本用法

memory = model.encoder(src)  # shape: (32, 12, 64)

编码器对输入 src 进行双向自注意力计算:每个位置都能看到所有其他位置(包括自己和之后的位置),这与解码器不同。

带填充掩码

memory = model.encoder(src, src_key_padding_mask=src_key_padding_mask)
# shape: (32, 12, 64)

加上填充掩码后,编码器在计算注意力时会忽略填充位置(ID=0的位置),使注意力完全聚焦在有效词上。

编码器 vs 解码器的注意力区别

编码器 解码器
自注意力类型 双向:能看到所有位置 单向(因果):只能看到当前及之前的位置
掩码 只需要 padding mask 需要 causal mask + padding mask
交叉注意力 有(关注编码器输出 memory)

6. 解码器单独使用

解码器需要两个输入:目标序列的嵌入 tgt 和编码器的输出 memory

基本用法

output = model.decoder(tgt, memory)  # shape: (32, 7, 64)
  • tgt:目标序列的嵌入向量 (32, 7, 64)
  • memory:编码器对源序列的编码结果 (32, 12, 64)

解码器内部会做两件事:

  1. 自注意力:目标序列各位置之间互相注意
  2. 交叉注意力:目标序列关注源序列的编码结果(memory)

带所有掩码

output = model.decoder(
    tgt,                                          # 目标序列嵌入
    memory,                                       # 编码器输出
    tgt_mask=tgt_mask,                            # 因果掩码
    tgt_key_padding_mask=tgt_key_padding_mask,    # 目标序列填充掩码
    memory_key_padding_mask=memory_key_padding_mask,  # 源序列填充掩码
)
# shape: (32, 7, 64)

输出的后续处理

output 的形状是 (batch_size, tgt_len, d_model) = (32, 7, 64)

要得到每个位置的词概率分布,需要加一个线性层映射到目标词表大小:

output_projection = nn.Linear(d_model, tgt_vocab_size)  # (64 → 1500)
logits = output_projection(output)  # shape: (32, 7, 1500)
probs = torch.softmax(logits, dim=-1)  # 每个位置的词概率分布
output (32, 7, 64)  →  Linear(64→1500)  →  logits (32, 7, 1500)  →  取最大概率的词ID

核心要点总结

概念 一句话理解
nn.Transformer PyTorch 提供的完整 Transformer,包含编码器和解码器
d_model 每个位置用多少维向量表示
nhead 多头注意力的头数,将 d_model 分成多个子空间
batch_first=True 输入张量的第0维是批次大小
Padding Mask 告诉模型忽略填充位置,True=忽略
Causal Mask 防止解码器偷看未来的词,上三角矩阵为True
编码器 双向自注意力,输出 memory
解码器 因果自注意力 + 交叉注意力,输出每个位置的表示
自回归 推理时逐个词生成,已生成的词作为下一步输入

完整流程图

┌─────────────────────────────────────────────────────────────────────┐
│                    Transformer 完整 Pipeline                          │
└─────────────────────────────────────────────────────────────────────┘

源序列词ID:  [256, 89, 423, 7, 0, 0, 0, 0, 0, 0, 0, 0]
    │
    ▼  nn.Embedding (查表)
源序列向量:  (32, 12, 64)
    │
    ▼  model.encoder (双向自注意力)
memory:      (32, 12, 64)  ← 源序列的语义编码

目标序列词ID: [101, 567, 3, 890, 0, 0, 0]
    │
    ▼  nn.Embedding (查表)
目标序列向量: (32, 7, 64)
    │
    ▼  model.decoder (因果自注意力 + 交叉注意力)
               ↑ 需要 memory 和各种 mask
output:       (32, 7, 64)
    │
    ▼  Linear (64 → 词表大小)
logits:       (32, 7, 1500)
    │
    ▼  argmax (取概率最大的词)
翻译结果:     [101, 567, 3, 890, 42, 17, 5]
Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐