来源:互联网 更新时间:2026-08-25 14:12
今天来聊一个在深度学习领域堪称里程碑的模型——Transformer。毫不夸张地说,它彻底改变了自然语言处理(NLP)的面貌,从机器翻译到文本生成,从语义理解到如今横扫一切的LLM,背后都有它的影子。

Transformer 是一种基于注意力机制的深度学习模型,最早由 Vaswani 等人在2017年的论文《Attention is All You Need》中提出。相比传统的序列模型(比如RNN、LSTM),它最大的突破在于:能够并行处理整个序列,而不是一个词一个词地串行跑。这一变化带来的效率提升是质的飞跃,也使得训练更大、更深的模型成为可能。
Transformer 的架构由两大块组成:
编码器的任务是给输入序列做“特征提取”。每一层编码器包含两个子层:
每个子层后面都跟着残差连接(Residual Connection)和层归一化(Layer Normalization),这有助于缓解深层网络的梯度问题,让训练更稳定、收敛更快。
解码器结构与编码器类似,但多了一个子层,所以每层共有三个子层:
每个子层同样有残差连接和层归一化。
下面我们逐一拆解 Transformer 中那些关键的设计。
输入嵌入就是把文本中的单词或子词映射成高维向量。具体来说,文本先被切分成 token,然后每个 token 通过一个查找表(嵌入矩阵)转换成固定长度的向量。这个嵌入矩阵是在训练过程中学习得到的。说白了,就是把离散的符号变成计算机能理解的连续向量。
由于 Transformer 不像 RNN 那样天然具备顺序信息(RNN 是按时间步依次处理的),它必须通过额外的机制告诉模型“这个词在句子中的第几个位置”。位置编码(Positional Encoding)就是用来干这个的。原文使用了正弦和余弦函数来生成位置向量,公式如下:
对于位置 pos 和嵌入维度中的第 2i+1 个维度:
其中 pos 是位置索引,i 是维度索引,dₘₒₐₗ 是嵌入向量的维度。这种设计让模型能够编码位置信息,并且可以外推到更长的序列。
自注意力是整个 Transformer 的核心创新。它允许模型在计算某个位置的表示时,关注输入序列中所有其他位置的信息,而不是只局限于局部。具体计算分三步:
公式流程如下:
单靠一个自注意力机制可能还不够——不同角度的特征都值得关注。多头注意力(Multi-Head Attention)就是并行地执行多次自注意力,每个头有自己独立的 Q、K、V 变换,然后把所有头的输出拼接起来,再经过一个线性层得到最终结果。公式如下:
其中 Wⁱ_Q, Wⁱ_K, Wⁱ_V 是第 i 个头的权重矩阵。将所有头的输出拼接后再经过 W_O 线性变换:
多头注意力的好处很明显:它能捕捉不同子空间中的语义关系,让模型表达能力更强。
每个编码器和解码器层中的 FFN 是一个两层的全连接网络,对每个位置的表示独立做非线性变换。公式如下:
其中 W₁、W₂ 是权重矩阵,b₁、b₂ 是偏置。
为了防止网络变深后梯度消失或爆炸,每个子层(注意力或 FFN)后面都接一个残差连接,再跟上层归一化。计算公式:
这里的 Sublayer(x) 可以是多头注意力或前馈网络的输出。
在解码器中,生成序列时绝不能看到未来的单词。所以需要在计算注意力时对未来的位置加掩码——把那些位置的分数设为负无穷,这样 Softmax 后权重就变成零。这就是 Masked Multi-Head Attention 的设计。
解码器中的第二个注意力子层叫做 Encoder-Decoder Attention。它的 Query 来自解码器前一层(自注意力的输出),而 Key 和 Value 则来自编码器的输出。这样解码器就能根据输入序列来生成合理的输出序列。
下面用 PyTorch 实现一个简单的 Transformer 模型,用来演示序列到序列任务(比如机器翻译)的基本流程。
import torch
import torch.nn as nn
import torch.optim as optim
import math
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super(PositionalEncoding, self).__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0).transpose(0, 1)
self.register_buffer('pe', pe)
def forward(self, x):
x = x + self.pe[:x.size(0), :]
return x
class TransformerModel(nn.Module):
def __init__(self, input_dim, output_dim, d_model=512, nhead=8, num_encoder_layers=6, dim_feedforward=2048, dropout=0.1):
super(TransformerModel, self).__init__()
self.model_type = 'Transformer'
self.embedding = nn.Embedding(input_dim, d_model)
self.pos_encoder = PositionalEncoding(d_model)
encoder_layers = nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout)
self.transformer_encoder = nn.TransformerEncoder(encoder_layers, num_encoder_layers)
self.d_model = d_model
self.decoder = nn.Linear(d_model, output_dim)
self.init_weights()
def init_weights(self):
initrange = 0.1
self.embedding.weight.data.uniform_(-initrange, initrange)
self.decoder.bias.data.zero_()
self.decoder.weight.data.uniform_(-initrange, initrange)
def forward(self, src, src_mask):
src = self.embedding(src) * math.sqrt(self.d_model)
src = self.pos_encoder(src)
output = self.transformer_encoder(src, src_mask)
output = self.decoder(output)
return output
def generate_square_subsequent_mask(sz):
mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)
mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))
return mask
# 使用示例
input_dim = 1000 # 词汇表大小
output_dim = 1000 # 输出大小
seq_length = 10 # 序列长度
model = TransformerModel(input_dim=input_dim, output_dim=output_dim)
src = torch.randint(0, input_dim, (seq_length, 32)) # (seq_len, batch_size)
src_mask = generate_square_subsequent_mask(seq_length)
output = model(src, src_mask)
print(output.shape) # 预期输出: [seq_len, batch_size, output_dim]
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# 简单训练循环
for epoch in range(10):
optimizer.zero_grad()
output = model(src, src_mask)
loss = criterion(output.view(-1, output_dim), src.view(-1))
loss.backward()
optimizer.step()
print(f"Epoch {epoch+1}, Loss: {loss.item()}")