三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

从零实现Transformer:自注意力机制与PyTorch实战详解

从零实现Transformer:自注意力机制与PyTorch实战详解

1. 项目概述:从“注意力”到“变革者”

如果你在过去几年里稍微关注过人工智能,尤其是自然语言处理领域,那么“Transformer”这个词一定如雷贯耳。它早已不是一个简单的模型名称,而是一个时代的标志。从最初的机器翻译任务中脱颖而出,到如今驱动着像GPT、BERT、文心一言、通义千问等几乎所有顶尖的大语言模型,Transformer架构彻底重塑了我们对序列建模的认知。很多人第一次接触它时,会被论文《Attention Is All You Need》中复杂的结构图和数学公式劝退,觉得这玩意儿是“天书”。但事实上,它的核心思想非常直观有力,甚至可以说,它解决了一个困扰了RNN和LSTM等前辈模型多年的根本性难题。

简单来说,Transformer是一个完全基于“自注意力机制”构建的深度学习模型架构。它放弃了传统的循环神经网络(RNN)那种按时间步顺序处理序列的方式,转而允许序列中的任意两个位置直接建立联系,无论它们相距多远。这种设计带来了两个革命性的优势:一是极强的并行计算能力,训练速度大幅提升;二是能够更好地捕捉长距离的依赖关系,理解“巴黎是法国的首都”中“巴黎”和“法国”的关联,与它们中间隔了多少个词无关。这个项目,我们就来亲手拆解这个“变革者”,从最基础的概念入手,把它内部每一个齿轮是如何咬合的,都看得清清楚楚。无论你是刚入门深度学习的新手,还是想巩固基础的老手,这篇内容都将带你绕过那些晦涩的论文表述,用最直白的语言和可运行的代码,把Transformer的里里外外讲明白。

2. 核心思想:为什么是“Attention Is All You Need”?

在Transformer出现之前,序列建模的王者是RNN及其变体LSTM、GRU。它们的工作方式很像我们阅读:从左到右,一个字一个字地看,同时脑子里记住前面看过的内容(隐藏状态),以此来理解整个句子的意思。这种方式很符合直觉,但存在明显的瓶颈。

2.1 RNN/LSTM的固有缺陷

首先,并行化困难。因为必须等第t步计算完,才能计算第t+1步,这就像一条单行道,无法让多辆车同时通过。在如今动辄使用数十甚至上百个GPU进行训练的时代,这种串行计算是巨大的效率瓶颈。

其次,长距离依赖捕捉能力弱。尽管LSTM通过门控机制缓解了梯度消失问题,但信息在漫长的序列中逐层传递,仍然会不可避免地衰减或混杂。当一个句子开头的信息需要影响到句子末尾时,这个“信号”需要穿越很多步,很容易变得模糊不清。

Transformer的论文标题“Attention Is All You Need”就像一份宣言,它指出:要理解一个词,你不需要按顺序记住前面所有的词,你只需要让模型学会“注意”当前句子中所有与之相关的词,无论它们在前还是在后。这种机制就是“自注意力”。

2.2 自注意力机制的精髓

想象一下你在阅读一段复杂的文章。当你看到“它”这个代词时,你会本能地向前回溯,寻找它所指代的那个名词(比如“苹果公司”)。你的注意力在句子内的不同位置间跳跃。自注意力机制让模型学会了做同样的事情。

它的计算过程可以概括为三步:

  1. 为每个词生成三把“钥匙”:查询向量(Query)、键向量(Key)、值向量(Value)。你可以把Query理解为“我(当前词)想知道什么”,Key是“我(其他词)有什么信息”,Value是“我(其他词)的实际内容”。
  2. 计算注意力分数:用当前词的Query去和序列中所有词的Key做点积(衡量相似度)。这样,当前词就和所有词都进行了一次“亲密程度”打分。
  3. 加权求和:将这些分数通过Softmax函数归一化为权重(所有权重和为1),然后用这些权重对所有的Value向量进行加权求和。最终得到的向量,就是融合了全局相关信息的、新的当前词表示。

这个过程是同时对所有词进行的,完美实现了并行计算。并且,因为任意两个词都直接计算了关联分数,所以长距离依赖问题迎刃而解。

注意:这里说的“词”在NLP中通常是“词元”,可能是单词、子词或字符。在视觉任务中,则是图像分块后的“图块”。

3. Transformer架构的逐层拆解

理解了自注意力这个核心发动机,我们来看Transformer整台机器的蓝图。它遵循经典的编码器-解码器结构,但内部组件全部焕然一新。

3.1 整体框架:编码器与解码器堆叠

一个标准的Transformer模型由N个相同的编码器层堆叠而成,以及N个相同的解码器层堆叠而成。原论文中N=6。编码器负责将输入序列(如一句英文)编码成一个富含上下文信息的中间表示;解码器则利用这个中间表示,并结合之前已生成的输出,自回归地生成目标序列(如对应的中文)。

编码器层包含两个子层:

  1. 多头自注意力层
  2. 前馈神经网络层 每个子层外面都包裹着“残差连接”和“层归一化”。公式可以简化为:LayerOutput = LayerNorm(x + Sublayer(x))。残差连接让梯度更容易流动,缓解深层网络训练中的梯度消失问题。

解码器层包含三个子层:

  1. 带掩码的多头自注意力层(确保预测时看不到未来信息)
  2. 多头交叉注意力层(让解码器关注编码器的输出)
  3. 前馈神经网络层 同样,每个子层也都有残差连接和层归一化。

3.2 输入处理:词嵌入与位置编码

模型首先需要把离散的符号(词)变成连续的向量。这就是词嵌入层的工作。它将每个词元映射到一个高维向量(例如512维)。但是,自注意力机制本身不具备感知词序的能力,打乱输入词的顺序,得到的注意力输出是一样的。这显然不符合语言规律。

为了解决这个问题,Transformer引入了位置编码。它为序列中的每个位置(第1个词,第2个词...)生成一个独一无二的、与词嵌入同维度的向量,然后直接加到词嵌入向量上。这样,模型就能同时知道“这个词是什么”以及“这个词在什么位置”。

位置编码的生成使用了正弦和余弦函数:PE(pos, 2i) = sin(pos / 10000^(2i/d_model))PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))其中pos是位置,i是维度索引。这种函数形式能让模型轻松地学习到相对位置关系(例如,位置pos+k的编码可以由位置pos的编码线性表示)。

3.3 核心中的核心:多头注意力机制详解

这是Transformer最精彩的部分。与其只做一次自注意力计算,为什么不并行地做多次呢?这就是“多头”的由来。

3.3.1 单头注意力计算过程我们先把公式摆出来,再一步步解释:Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V

假设输入序列有seq_len个词,每个词的嵌入向量是d_model维。我们通过三个不同的线性变换矩阵W_Q,W_K,W_V,将每个词的输入向量x投影到d_kd_kd_v维,得到Q, K, V。通常d_k = d_v = d_model / h,其中h是头的数量。

  1. QK^T:计算一个seq_len * seq_len的矩阵,其中每个元素(i, j)代表第i个词对第j个词的注意力分数。
  2. 除以 sqrt(d_k):这是一个非常关键的缩放操作。因为点积的结果会随着维度d_k增大而增大,经过Softmax后,梯度会变得非常小。缩放可以稳定梯度。
  3. Softmax:对每一行进行归一化,使得当前词对所有词的注意力权重和为1。
  4. 乘以V:用归一化后的权重对V矩阵进行加权求和,得到每个词新的表示。

3.3.2 如何实现“多头”?所谓“多头”,就是准备h套不同的W_Q,W_K,W_V矩阵。同一份输入,分别用这h套矩阵进行投影,并行地计算出h组不同的Q, K, V,然后独立进行上面的注意力计算,得到h个输出矩阵。每个头的d_k都较小,例如d_model=512,h=8,则d_k=64

这h个输出矩阵(每个形状为[seq_len, d_k])被拼接起来,形成一个[seq_len, h*d_k = d_model]的大矩阵。最后,再通过一个可学习的线性投影矩阵W_O将其映射回d_model维,作为多头注意力层的最终输出。

为什么需要多头?这相当于让模型同时从多个不同的“表示子空间”来关注信息。有的头可能更关注语法结构(如主谓一致),有的头可能更关注语义关联(如同义词),有的头可能更关注指代关系。这种分工协作使得模型的表示能力更加强大。

3.4 前馈神经网络与归一化

注意力层负责融合信息,而前馈神经网络层则负责对每个位置的特征进行独立、非线性的变换和增强。它是一个两层全连接网络,中间有一个ReLU激活函数:FFN(x) = max(0, xW1 + b1)W2 + b2值得注意的是,这个网络对序列中的每个位置是独立、相同地应用的,这又是一处可以高度并行化的设计。

层归一化是Transformer稳定训练的另一个关键。它不像批归一化那样对一个批次内所有样本的同一特征进行归一化,而是对单个样本的所有特征进行归一化。这对于变长序列处理尤其友好,使得模型对批次大小的变化不敏感。

4. 从零开始:动手实现一个微型Transformer

理解了原理,最好的巩固方式就是动手实现。我们将使用PyTorch框架,构建一个超小规模的Transformer,用于一个简单的任务:学习复制输入序列。这个任务能直观地检验模型是否学会了关注输入。

4.1 环境准备与超参数定义

首先,确保你安装了PyTorch。然后,我们定义模型的核心超参数。

import torch import torch.nn as nn import torch.optim as optim import math # 超参数定义 d_model = 128 # 词嵌入和模型内部特征的维度 num_heads = 8 # 注意力头的数量 num_layers = 3 # 编码器和解码器的层数 d_ff = 512 # 前馈网络中间层的维度 dropout_rate = 0.1 # Dropout比率,防止过拟合 max_seq_len = 100 # 最大序列长度 vocab_size = 100 # 词汇表大小(假设我们只有100个不同的符号)

4.2 实现位置编码

根据公式实现正弦位置编码。

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) # [max_len, 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) # 偶数维度用sin pe[:, 1::2] = torch.cos(position * div_term) # 奇数维度用cos pe = pe.unsqueeze(0) # [1, max_len, d_model],方便广播 self.register_buffer('pe', pe) # 这不是模型参数,不参与梯度更新 def forward(self, x): # x: [batch_size, seq_len, d_model] seq_len = x.size(1) x = x + self.pe[:, :seq_len, :] # 直接相加 return x

4.3 实现多头注意力层

这是核心组件,需要仔细实现缩放点积注意力。

class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout=0.1): super(MultiHeadAttention, self).__init__() assert d_model % num_heads == 0, "d_model must be divisible by num_heads" self.d_model = d_model self.num_heads = num_heads self.d_k = d_model // num_heads # 定义Q, K, V和最终输出的线性投影层 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) self.dropout = nn.Dropout(dropout) def forward(self, query, key, value, mask=None): batch_size = query.size(0) # 1. 线性投影并分头 Q = self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # [B, h, seq_len, d_k] K = self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V = self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # 2. 计算缩放点积注意力 scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # [B, h, seq_len, seq_len] if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) # 将mask为0的位置填充为负无穷 attn_weights = torch.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) context = torch.matmul(attn_weights, V) # [B, h, seq_len, d_k] # 3. 合并多头 context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) # [B, seq_len, d_model] # 4. 最终线性投影 output = self.W_o(context) return output, attn_weights

4.4 实现前馈网络与编码器层

class PositionwiseFeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout=0.1): super(PositionwiseFeedForward, self).__init__() self.linear1 = nn.Linear(d_model, d_ff) self.linear2 = nn.Linear(d_ff, d_model) self.dropout = nn.Dropout(dropout) self.relu = nn.ReLU() def forward(self, x): return self.linear2(self.dropout(self.relu(self.linear1(x)))) class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super(EncoderLayer, self).__init__() self.self_attn = MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward = PositionwiseFeedForward(d_model, d_ff, dropout) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout1 = nn.Dropout(dropout) self.dropout2 = nn.Dropout(dropout) def forward(self, x, mask=None): # 子层1:多头自注意力 + 残差 & 层归一化 attn_output, _ = self.self_attn(x, x, x, mask) x = x + self.dropout1(attn_output) x = self.norm1(x) # 子层2:前馈网络 + 残差 & 层归一化 ff_output = self.feed_forward(x) x = x + self.dropout2(ff_output) x = self.norm2(x) return x

4.5 组装完整Transformer模型

我们将实现一个仅包含编码器的简化版Transformer,用于完成“复制序列”任务。

class CopyTransformer(nn.Module): def __init__(self, vocab_size, d_model, num_heads, num_layers, d_ff, max_seq_len, dropout=0.1): super(CopyTransformer, self).__init__() self.token_embedding = nn.Embedding(vocab_size, d_model) self.positional_encoding = PositionalEncoding(d_model, max_seq_len) self.encoder_layers = nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.final_layer = nn.Linear(d_model, vocab_size) # 输出层,预测下一个词的概率分布 self.dropout = nn.Dropout(dropout) def forward(self, src_tokens, src_mask=None): # 1. 词嵌入 + 位置编码 x = self.token_embedding(src_tokens) # [B, seq_len] -> [B, seq_len, d_model] x = self.positional_encoding(x) x = self.dropout(x) # 2. 通过多层编码器 for layer in self.encoder_layers: x = layer(x, src_mask) # 3. 投影到词汇表空间 logits = self.final_layer(x) # [B, seq_len, vocab_size] return logits

4.6 训练与验证:学习复制序列

我们创建一个简单的训练循环,让模型学会输出与输入完全相同的序列。

def train_simple_copy_task(): # 初始化模型、优化器和损失函数 model = CopyTransformer(vocab_size, d_model, num_heads, num_layers, d_ff, max_seq_len, dropout_rate) optimizer = optim.Adam(model.parameters(), lr=0.0001, betas=(0.9, 0.98), eps=1e-9) criterion = nn.CrossEntropyLoss(ignore_index=0) # 假设0是填充符<PAD> model.train() for epoch in range(50): total_loss = 0 # 模拟一个简单的批次数据:随机生成长度为10的序列 batch_size = 32 seq_len = 10 src_data = torch.randint(1, vocab_size, (batch_size, seq_len)) # 忽略0(PAD) tgt_data = src_data.clone() # 目标就是复制输入 # 前向传播 logits = model(src_data) # [B, seq_len, vocab_size] # 计算损失时,我们将logits和target都reshape成二维 loss = criterion(logits.view(-1, vocab_size), tgt_data.view(-1)) # 反向传播与优化 optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() if (epoch + 1) % 10 == 0: print(f'Epoch [{epoch+1}/50], Loss: {loss.item():.4f}') # 简单推理测试 model.eval() with torch.no_grad(): test_input = torch.randint(1, vocab_size, (1, seq_len)) output_logits = model(test_input) predicted_ids = output_logits.argmax(dim=-1) # 取概率最大的词元ID print(f"Input: {test_input.squeeze().tolist()}") print(f"Output: {predicted_ids.squeeze().tolist()}") print(f"Match: {torch.all(test_input == predicted_ids).item()}") if __name__ == '__main__': train_simple_copy_task()

运行这段代码,你会看到模型损失逐渐下降,并且在训练结束后,对于简短的输入序列,它应该能近乎完美地复制出来。这证明了我们的微型Transformer已经学会了基本的“注意力”和序列映射能力。

5. 关键技巧与实战避坑指南

在理论理解和基础实现之上,要让Transformer在实际任务中发挥威力,还需要掌握一系列工程技巧。这些往往是论文里一笔带过,但实践中至关重要的部分。

5.1 注意力掩码的艺术

掩码是控制注意力范围的关键工具,主要有两种:

  1. 填充掩码:在处理变长序列时,我们会用<PAD>符号将批次内的序列补齐到相同长度。在计算注意力时,需要屏蔽这些填充位置,防止模型关注无意义的<PAD>。通常生成一个布尔矩阵,<PAD>位置为False(或0)。
  2. 序列掩码(因果掩码):在解码器的自注意力层中,必须确保在预测第t个位置时,只能看到第1t-1个位置的信息,不能“偷看”未来。这通过一个下三角矩阵(主对角线及以上为-inf,以下为0)来实现。
# 生成因果掩码的示例 def generate_causal_mask(seq_len): mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool() # triu返回上三角矩阵,diagonal=1表示不包括主对角线 # 结果为:未来位置是True(需要被mask掉) return mask # 在注意力计算中:scores.masked_fill(mask, -1e9)

5.2 学习率预热与衰减策略

Transformer模型对学习率非常敏感。常见的策略是使用“预热”学习率调度器:在训练初期用一个较小的学习率线性预热,达到一个峰值后再按步数或轮次的平方根倒数进行衰减。这有助于模型在初期稳定,后期精细调优。

# 类似原始论文的Warmup调度器 def get_warmup_scheduler(optimizer, d_model, warmup_steps=4000): def lr_lambda(step): # step从1开始计数 return min(step ** -0.5, step * (warmup_steps ** -1.5)) return optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

5.3 梯度裁剪与标签平滑

Transformer层数多,容易产生梯度爆炸。梯度裁剪是标准操作,将梯度向量的范数限制在一个阈值内。

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

标签平滑是一种正则化技术,在计算交叉熵损失时,不是给正确标签1,其他0,而是给正确标签1 - epsilon,其他标签平均分配epsilon / (vocab_size - 1)。这可以防止模型对训练数据过度自信,提升泛化能力。PyTorch的CrossEntropyLoss可以通过设置label_smoothing参数实现。

5.4 初始化与层归一化的位置

参数的初始化对训练稳定性影响巨大。Transformer通常使用Xavier均匀初始化或更复杂的初始化方法。另一个细节是层归一化的位置。原始Transformer使用“后归一化”(在残差连接之后进行层归一化)。但后续研究发现,“前归一化”(在子层计算之前进行层归一化,即Pre-LN)能使训练更稳定、更容易收敛,这已成为许多现代Transformer变体的标准配置。

6. Transformer的演进与变体

自2017年诞生以来,Transformer本身也在不断进化,衍生出众多高效、专用的变体。

6.1 编码器派系:BERT与它的朋友们

BERT只使用了Transformer的编码器部分,通过“掩码语言模型”和“下一句预测”两个任务进行预训练,学习强大的双向语言表示。它催生了“预训练+微调”的范式革命。后续的RoBERTa、ALBERT、ELECTRA等都在此基础上进行优化,比如移除NSP任务、使用更大的批次和更长的训练时间、用更高效的预训练任务等。

6.2 解码器派系:GPT系列与自回归生成

GPT系列模型则专注于Transformer的解码器部分(严格来说是去掉了交叉注意力层的解码器堆叠),通过自回归的方式,给定上文预测下一个词。从GPT-1到GPT-3、ChatGPT,其核心架构思想一脉相承,但模型规模、训练数据和训练技巧发生了指数级增长,最终涌现出惊人的理解和生成能力。

6.3 视觉Transformer:当注意力遇见图像

ViT首次证明,将图像分割成固定大小的图块,线性投影后加上位置编码,直接送入标准Transformer编码器,就能在图像分类任务上取得媲美CNN的效果。这打破了计算机视觉领域CNN的长期统治。随后的Swin Transformer引入了“滑动窗口”和“分层下采样”思想,让Transformer能够像CNN一样高效处理多尺度特征,并计算复杂度与图像大小呈线性关系,成为视觉领域的里程碑。

6.4 高效Transformer:解决计算与内存瓶颈

标准自注意力的计算复杂度是序列长度的平方级(O(n²)),这限制了其处理超长序列(如长文档、高分辨率图像)的能力。为此,研究者提出了多种高效注意力变体:

  • 稀疏注意力:如Longformer、BigBird,只计算所有注意力对中的一部分(如滑动窗口、全局注意力)。
  • 线性化注意力:如Linformer、Performer,通过核函数技巧将注意力计算近似为线性复杂度。
  • 分块/递归注意力:如Reformer使用局部敏感哈希将相似的键值分到同一桶中;Transformer-XL引入循环机制处理超长文本。

7. 常见问题与调试心得

在实际实现和训练Transformer时,你几乎一定会遇到下面这些问题。

7.1 模型不收敛或损失为NaN

这是最常见的问题,可能的原因和排查步骤:

  1. 检查学习率:这是首要怀疑对象。尝试将学习率调低1-2个数量级(例如从1e-3调到1e-4或1e-5)。务必使用预热策略。
  2. 检查梯度:在反向传播后、优化器更新前,打印梯度的范数。如果出现NaN或无穷大,说明计算过程有问题。
    total_norm = 0 for p in model.parameters(): if p.grad is not None: param_norm = p.grad.data.norm(2) total_norm += param_norm.item() ** 2 total_norm = total_norm ** 0.5 print(f"Gradient norm: {total_norm}")
  3. 检查数据:确保输入数据中没有异常值(如NaN, inf),标签是否在有效范围内。
  4. 检查初始化:使用标准的初始化方法,如nn.init.xavier_uniform_
  5. 启用梯度裁剪:设置一个合理的阈值(如1.0或5.0)。

7.2 训练速度慢,GPU利用率低

  1. 增大批次大小:在GPU内存允许的范围内,尽可能使用更大的批次。这能提高并行度和硬件利用率。
  2. 检查序列长度:过长的序列会导致注意力计算量剧增。对于训练,可以尝试截断或使用动态批处理(将相似长度的样本放在同一批)。
  3. 使用混合精度训练:利用PyTorch的AMP(自动混合精度)功能,可以显著减少内存占用并加速计算,尤其在大模型上效果明显。
  4. 分析瓶颈:使用PyTorch Profiler或简单的计时,找出是数据加载慢还是模型计算慢。

7.3 过拟合与欠拟合

  1. 过拟合:模型在训练集上表现好,验证集差。
    • 增加正则化:提高Dropout比率,尝试更多的Dropout位置(如注意力权重后、前馈网络中间)。
    • 数据增强:对于NLP任务,可以使用回译、同义词替换、随机删除等。对于CV任务,使用标准的图像增强。
    • 早停:监控验证集损失,当连续多个epoch不再下降时停止训练。
    • 减小模型规模:如果数据量有限,一个更小的模型可能更合适。
  2. 欠拟合:模型在训练集上表现就很差。
    • 增加模型容量:增加d_modelnum_headsnum_layers
    • 降低正则化:减小Dropout比率。
    • 检查特征:确保输入特征包含了足够的信息。
    • 训练更长时间:Transformer通常需要较长的训练周期才能充分收敛。

7.4 注意力权重可视化与解释性

理解模型在“看”哪里是调试和解释模型行为的重要手段。在实现多头注意力时,我们已经返回了attn_weights。可以将其可视化:

import matplotlib.pyplot as plt import seaborn as sns def plot_attention_weights(attention_weights, source_tokens, target_tokens, head_idx=0): """ attention_weights: [batch, num_heads, target_len, source_len] """ attn = attention_weights[0, head_idx].cpu().detach().numpy() # 取第一个样本,第head_idx个头 plt.figure(figsize=(10, 8)) sns.heatmap(attn, xticklabels=source_tokens, yticklabels=target_tokens, cmap='viridis', cbar_kws={'label': 'Attention Weight'}) plt.xlabel('Source Tokens') plt.ylabel('Target Tokens') plt.title(f'Attention Weights (Head {head_idx})') plt.tight_layout() plt.show()

通过观察不同层、不同头的注意力图,你可以看到模型是否学会了关注语法结构、语义关联或指代关系。例如,在翻译任务中,你可能会看到目标语言的动词清晰地关注到源语言中对应的动词。

从我个人的多次实现和调试经验来看,Transformer就像一个精密的仪器,每一个部件(初始化、学习率、归一化、掩码)都必须调整到位,它才能稳定高效地运转。最开始实现时,最容易忽略的是缩放因子sqrt(d_k)正确的掩码应用,这两个细节出错会导致模型完全无法学习。另一个深刻的体会是,从一个小型任务(如复制序列)开始验证你的实现是否正确,远比直接在一个复杂任务上调试要高效得多。当你看到这个微型模型能完美复制输入时,你就获得了继续构建更复杂应用的坚实基础。

← 返回列表