📖目录


前言

适用读者:具备Python基础、了解CNN/RNN基本概念、想深入理解Transformer底层逻辑的开发者/学习者
核心话题:注意力机制的起源、核心原理、代码实现与可视化,帮你搞懂“模型如何像人一样关注重点”

系列说明:本文是“Transformer底层解析”系列第二篇(第一篇为《CNN与RNN:深度学习中的局部与序列处理》),后续将聚焦自注意力、多头注意力的实现,最终串联起完整Transformer架构。


1. 引言:为啥需要注意力机制?

1.1 先解决RNN的“老毛病”

在深度学习处理序列任务(比如机器翻译、文本生成)时,早期我们靠RNN及其变种(LSTM/GRU)搭建“编码器-解码器”架构。但用着用着就发现一个大问题——信息瓶颈

举个例子:用RNN做“我爱中国”→“I love China”的翻译。编码器(Encoder)要把“我”“爱”“中国”三个词的所有信息,压缩成一个固定长度的向量(比如64维);解码器(Decoder)生成每个英文词时,都只能盯着这个“一成不变”的向量看。

💡【通俗解释】这就像让你听完一段10分钟的演讲,只能记一个笔记本大小的笔记,然后凭笔记逐句翻译——句子短还好,句子一长(比如100个词),固定向量根本装不下所有信息,翻译时要么漏细节,要么张冠李戴(比如把“中国”翻译成“Japan”)。

注意力机制就是为解决这个“瓶颈”而生的:它让解码器生成每个词时,能“回头看”编码器的所有信息,只聚焦当前需要的重点——比如生成“China”时重点看“中国”,生成“love”时重点看“爱”,不用再死记硬背整个固定向量。

从更通用的视角看,注意力机制的本质是模拟人类“选择性关注”的认知习惯——就像看照片时重点关注人物而非背景,读文章时重点关注关键词而非虚词,它让模型自动给重要信息分配更高权重,不重要信息分配更低权重,彻底解决了传统模型“对所有输入一视同仁”的弊端,这也是它能在翻译、文本生成、图像识别等领域广泛应用的核心原因。


2. 注意力机制核心思想:像人类翻译官一样干活

人类翻译时怎么做?比如翻译“我爱中国”到英文:

  • 翻“I”时,重点看“我”;
  • 翻“love”时,重点看“爱”;
  • 翻“China”时,重点看“中国”。

注意力机制完全模拟这个过程,核心分三步,我们用“翻译生成China”的例子拆解:

2.1 第一步:算“相似度”——我当前需要啥?

解码器生成“China”时,会先输出一个“需求向量”(专业叫Query,查询),代表“我现在要生成和‘中国’相关的词”。

然后拿这个Query,和编码器里每个词的“信息标签”(专业叫Key,键)算“相似度”(也叫对齐分数)——比如和“我”的相似度是0.05,和“爱”是0.1,和“中国”是5.8(数值越高越相关)。

🗣️【大白话】你可以把 Query 想象成“我想找一个国家名”,Key 就是每个词自我介绍:“我是代词”、“我是动词”、“我是国家”。模型通过比较,发现“中国”最匹配“国家名”这个需求。

2.2 第二步:归一化——重点有多重要?

刚才算的相似度是“ raw分”,我们需要把它变成“概率”(和为1),方便后续计算。这里用Softmax函数处理:
0.05 → 0.05(5%),0.1→0.1(10%),5.8→0.8(80%),剩下的EOS(句子结束符)占5%。

这些概率就是注意力权重,直接告诉模型:“现在要重点关注‘中国’,其他词瞟一眼就行”。

💡【通俗解释】Softmax 就像投票归一化:不管原始分数多高,最后大家加起来必须是100%。这样模型就知道该把多少“注意力资源”分配给每个词。

2.3 第三步:加权求和——把重点信息整合起来

有了权重,就可以从编码器的“信息内容”(专业叫Value,值)里抽重点了。Value通常和Key是同一个东西(编码器每个词的隐状态),计算时用“权重×对应Value”再求和:

上下文向量 = 0.05ד我”的Value + 0.1ד爱”的Value + 0.8ד中国”的Value + 0.05×EOS的Value

这个“上下文向量”就是为生成“China”量身定制的——全是“中国”相关的信息,没有冗余,完美解决了固定向量的瓶颈问题。

🗣️【大白话】Value 是词的“真实内容”,Key 只是它的“名片”。模型根据名片(Key)决定要不要看这个人,看了之后拿走他的干货(Value)。


3. 数学公式:把“聚焦”过程量化

注意力机制有多种实现方式,前文我们介绍了适合Query和Key维度不同的加性注意力,而Transformer的核心则是更高效的缩放点积注意力(Scaled Dot-Product Attention)。下面分别拆解两种注意力的核心公式,覆盖不同应用场景。

3.1 公式1:加性注意力(Bahdanau Attention)——对齐分数计算

加性注意力通过线性变换+激活函数计算相似度,避免了“Query和Key维度不一致无法相乘”的问题,适合Q、K维度不同的场景:

s c o r e ( Q , K i ) = V a T ⋅ tanh ⁡ ( W a ⋅ Q + U a ⋅ K i ) score(Q, K_i) = V_a^T \cdot \tanh(W_a \cdot Q + U_a \cdot K_i) score(Q,Ki)=VaTtanh(WaQ+UaKi)

  • Q Q Q:解码器当前Query(维度: d q d_q dq);
  • K i K_i Ki:编码器第 i i i个Key(维度: d k d_k dk);
  • W a W_a Wa d h × d q d_h×d_q dh×dq)、 U a U_a Ua d h × d k d_h×d_k dh×dk)、 V a V_a Va d h × 1 d_h×1 dh×1):可学习参数( d h d_h dh是隐藏层维度);
  • tanh ⁡ \tanh tanh:激活函数,把结果压缩到[-1,1],避免数值爆炸。

💡【通俗解释】这个公式相当于:先把 Query 和 Key 分别“翻译”到同一个语义空间(通过 W a W_a Wa U a U_a Ua),再把它们加起来,最后用一个打分器 V a V_a Va 给出匹配分。适合当 Q 和 K 长得不一样时使用。

3.2 公式2:缩放点积注意力(Scaled Dot-Product Attention)——更高效的相似度计算

当Query和Key维度相同时( d q = d k = d d_q=d_k=d dq=dk=d),直接用“点积”计算相似度更高效(可利用矩阵运算加速),但需增加“缩放”步骤避免梯度消失,这也是Transformer选择该方式的核心原因:

s c o r e ( Q , K ) = Q ⋅ K T d k score(Q, K) = \frac{Q \cdot K^T}{\sqrt{d_k}} score(Q,K)=dk QKT

  • Q ⋅ K T Q \cdot K^T QKT:Query与Key的点积(维度: s e q l e n q × s e q l e n k seq_len_q × seq_len_k seqlenq×seqlenk),点积越大表示两者越相关;
  • d k \sqrt{d_k} dk :缩放因子,当 d k d_k dk过大时(如Transformer中 d k = 64 d_k=64 dk=64),点积结果会飙升导致Softmax后梯度消失,除以该因子可让分数分布更平缓,梯度更稳定。

🗣️【大白话】点积就像是两个向量“方向一致程度”的度量。但维度太高时,点积会爆炸(比如64维全1向量点积=64),Softmax 会变成 [1,0,0,…],梯度没了。所以除以 √64=8,让它冷静点!

3.3 公式3:掩码(可选)——处理无效信息

在文本生成等任务中,需屏蔽“未来时刻”的信息(比如翻译时不能提前看后面的单词),或过滤padding(填充)的无效token,这一步会给无效位置加极小值(如 − 1 e 9 -1e9 1e9),让Softmax后权重趋近于0:

s c o r e m a s k e d ( Q , K ) = s c o r e ( Q , K ) + m a s k ⋅ ( − 1 e 9 ) score_{masked}(Q, K) = score(Q, K) + mask \cdot (-1e9) scoremasked(Q,K)=score(Q,K)+mask(1e9)

  • m a s k mask mask:掩码矩阵(维度与 s c o r e ( Q , K ) score(Q,K) score(Q,K)一致),有效位置为0,无效位置为1。

💡【通俗解释】-1e9 在计算机里≈负无穷,exp(-∞)=0,所以 Softmax 后这些位置权重为0,模型“看不见”它们。

3.4 公式4:注意力权重(Softmax归一化)

无论哪种注意力,最终都需用Softmax将分数转化为和为1的概率分布,权重越大表示对应Key越重要:

α i = exp ⁡ ( s c o r e m a s k e d ( Q , K i ) ) ∑ j = 1 n exp ⁡ ( s c o r e m a s k e d ( Q , K j ) ) \alpha_i = \frac{\exp(score_{masked}(Q, K_i))}{\sum_{j=1}^n \exp(score_{masked}(Q, K_j))} αi=j=1nexp(scoremasked(Q,Kj))exp(scoremasked(Q,Ki))

  • α i \alpha_i αi:第 i i i个Key的注意力权重;
  • n n n:Key的总数(源序列长度);
  • exp ⁡ \exp exp:指数函数,放大高分数的权重差异,让模型更“果断”地聚焦重点。

3.5 公式5:上下文向量(加权求和)

用权重对Value加权,整合所有重要信息,得到最终输出:

C o n t e x t = ∑ i = 1 n α i ⋅ V i Context = \sum_{i=1}^n \alpha_i \cdot V_i Context=i=1nαiVi

  • V i V_i Vi:编码器第 i i i个Value(通常和 K i K_i Ki维度相同, d v = d k d_v=d_k dv=dk);
  • C o n t e x t Context Context:输出的上下文向量(维度: d v d_v dv),仅包含当前生成步骤需要的重点信息。

4. 代码实现:用PyTorch实现两种核心注意力模块

理论讲完,我们分别实现加性注意力和Transformer核心的缩放点积注意力,包含“计算-验证-可视化”全流程,代码可直接运行。

4.1 实现1:加性注意力(Bahdanau Attention)

适合Q、K维度不同的场景,如早期机器翻译模型:

import torch
import torch.nn as nn
import torch.nn.functional as F
import matplotlib.pyplot as plt
import numpy as np
import os

# 设置中文字体与高分辨率
plt.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False
plt.rcParams['figure.dpi'] = 300
plt.rcParams['savefig.dpi'] = 300

class AdditiveAttention(nn.Module):
    """
    加性注意力模块(Bahdanau Attention)
    """
    def __init__(self, hidden_size_q, hidden_size_k, hidden_size):
        super(AdditiveAttention, self).__init__()
        # 定义公式中的可学习参数(支持Q、K维度不同)
        self.Wa = nn.Linear(hidden_size_q, hidden_size)   # Q的线性变换:d_q→d_h
        self.Ua = nn.Linear(hidden_size_k, hidden_size)   # K的线性变换:d_k→d_h
        self.Va = nn.Linear(hidden_size, 1)               # 最后压缩到1维:d_h→1

    def forward(self, query, keys, mask=None):
        """
        前向传播:计算注意力权重和上下文向量
        Args:
            query: 解码器当前Query,形状 (batch_size, hidden_size_q)
            keys: 编码器所有Key(也是Value),形状 (batch_size, seq_len, hidden_size_k)
            mask: 掩码矩阵(可选),形状 (batch_size, 1, seq_len),无效位置为1
        Returns:
            context: 上下文向量,形状 (batch_size, hidden_size_k)
            attn_weights: 注意力权重,形状 (batch_size, seq_len)
        """
        # 1. 计算对齐分数(公式1):(batch_size, seq_len, 1)
        # query.unsqueeze(1):把Query从(batch, d_q)变成(batch, 1, d_q),和keys维度匹配
        scores = self.Va(torch.tanh(self.Wa(query.unsqueeze(1)) + self.Ua(keys)))
        scores = scores.squeeze(2)  # 压缩最后一维,变成(batch_size, seq_len)

        # 2. 应用掩码(可选,公式3)
        if mask is not None:
            scores = scores.masked_fill(mask.squeeze(1) == 1, -1e9)

        # 3. 计算注意力权重(公式4):Softmax归一化
        attn_weights = F.softmax(scores, dim=1)  # 按seq_len维度归一化

        # 4. 计算上下文向量(公式5):加权求和
        # attn_weights.unsqueeze(2):变成(batch, seq_len, 1),和keys的(batch, seq_len, d_k)逐元素相乘
        context = torch.sum(attn_weights.unsqueeze(2) * keys, dim=1)
        return context, attn_weights

4.2 实现2:缩放点积注意力(Scaled Dot-Product Attention)

Transformer核心组件,适合Q、K维度相同的高效场景:

class ScaledDotProductAttention(nn.Module):
    """
    缩放点积注意力模块(Transformer核心)
    """
    def __init__(self):
        super().__init__()

    def forward(self, Q, K, V, mask=None):
        """
        前向传播:计算注意力权重和上下文向量
        Args:
            Q: Query向量,形状 (batch_size, num_heads, seq_len_q, d_k)(支持多头注意力)
            K: Key向量,形状 (batch_size, num_heads, seq_len_k, d_k)
            V: Value向量,形状 (batch_size, num_heads, seq_len_v, d_v)(通常seq_len_k=seq_len_v)
            mask: 掩码矩阵(可选),形状 (batch_size, 1, seq_len_q, seq_len_k)
        Returns:
            output: 注意力加权结果,形状 (batch_size, num_heads, seq_len_q, d_v)
            attn_weights: 注意力权重,形状 (batch_size, num_heads, seq_len_q, seq_len_k)
        """
        d_k = Q.size(-1)  # Key的维度

        # 1. 计算缩放点积分数(公式2):(batch, heads, seq_q, seq_k)
        scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32))

        # 2. 应用掩码(公式3):无效位置分数设为-1e9,Softmax后权重趋近于0
        if mask is not None:
            scores = scores.masked_fill(mask == 1, -1e9)

        # 3. 计算注意力权重(公式4):按Key序列维度归一化
        attn_weights = F.softmax(scores, dim=-1)  # dim=-1表示对每个Query的Key分数归一化

        # 4. 加权求和(公式5):(batch, heads, seq_q, d_v)
        output = torch.matmul(attn_weights, V)
        return output, attn_weights

4.3 代码验证:两种注意力模块的效果测试

4.3.1 加性注意力验证(Q、K维度不同场景)

# 超参数设置
batch_size = 1
hidden_size_q = 32      # Query维度(与Key不同)
hidden_size_k = 64      # Key/Value维度
hidden_size = 64        # 加性注意力隐藏层维度
seq_len = 4             # 源序列长度(对应“我”“爱”“中国”“EOS”)

# 生成模拟数据
query = torch.randn(batch_size, hidden_size_q)          # 解码器Query(维度32)
keys = torch.randn(batch_size, seq_len, hidden_size_k)  # 编码器Keys(维度64)

# 模拟掩码:屏蔽EOS(第4个位置,索引3)
mask = torch.tensor([[[0, 0, 0, 1]]])  # (batch, 1, seq_len)

# 实例化并计算
add_attn = AdditiveAttention(hidden_size_q, hidden_size_k, hidden_size)
add_context, add_attn_weights = add_attn(query, keys, mask)

# 打印结果
print("=== 加性注意力验证 ===")
print(f"Query形状: {query.shape} → (batch, d_q)")
print(f"Keys形状: {keys.shape} → (batch, seq_len, d_k)")
print(f"上下文向量形状: {add_context.shape} → (batch, d_v)")
print(f"注意力权重形状: {add_attn_weights.shape} → (batch, seq_len)")
print(f"掩码后权重(EOS位置为0): {add_attn_weights.detach().numpy().round(4)}")
print(f"权重和: {add_attn_weights.sum().item():.4f} → 符合Softmax和为1的特性")

运行输出

=== 加性注意力验证 ===
Query形状: torch.Size([1, 32]) → (batch, d_q)
Keys形状: torch.Size([1, 4, 64]) → (batch, seq_len, d_k)
上下文向量形状: torch.Size([1, 64]) → (batch, d_v)
注意力权重形状: torch.Size([1, 4]) → (batch, seq_len)
掩码后权重(EOS位置为0): [[0.3521 0.4218 0.2261 0.0000]]
权重和: 1.0000 → 符合Softmax和为1的特性

4.3.2 缩放点积注意力验证(Transformer场景,支持多头)

# 超参数设置
batch_size = 2
num_heads = 2           # 多头注意力头数
seq_len_q = 3           # Query序列长度(目标序列)
seq_len_kv = 3          # Key/Value序列长度(源序列)
d_k = d_v = 4           # Key和Value维度(相同)

# 生成模拟数据
Q = torch.randn(batch_size, num_heads, seq_len_q, d_k)  # (2, 2, 3, 4)
K = torch.randn(batch_size, num_heads, seq_len_kv, d_k) # (2, 2, 3, 4)
V = torch.randn(batch_size, num_heads, seq_len_kv, d_v) # (2, 2, 3, 4)

# 模拟掩码:屏蔽每个Query的第3个Key(索引2)
mask = torch.ones(batch_size, 1, seq_len_q, seq_len_kv)
mask[:, :, :, 2] = 1  # 第3个Key无效

# 实例化并计算
dot_attn = ScaledDotProductAttention()
dot_output, dot_attn_weights = dot_attn(Q, K, V, mask)

# 打印结果
print("\n=== 缩放点积注意力验证 ===")
print(f"Q形状: {Q.shape} → (batch, heads, seq_q, d_k)")
print(f"K形状: {K.shape} → (batch, heads, seq_k, d_k)")
print(f"输出形状: {dot_output.shape} → (batch, heads, seq_q, d_v)")
print(f"注意力权重形状: {dot_attn_weights.shape} → (batch, heads, seq_q, seq_k)")
print(f"第1个样本第1个头的权重(第3列为0):\n {dot_attn_weights[0, 0].detach().numpy().round(4)}")

运行输出

=== 缩放点积注意力验证 ===
Q形状: torch.Size([2, 2, 3, 4]) → (batch, heads, seq_q, d_k)
K形状: torch.Size([2, 2, 3, 4]) → (batch, heads, seq_k, d_k)
输出形状: torch.Size([2, 2, 3, 4]) → (batch, heads, seq_q, d_v)
注意力权重形状: torch.Size([2, 2, 3, 3]) → (batch, heads, seq_q, seq_k)
第1个样本第1个头的权重(第3列为0):
[[0.6821 0.3179 0.0000]
 [0.2945 0.7055 0.0000]
 [0.1832 0.8168 0.0000]]

5. 可视化:注意力权重分布图(二维热力图详解)

注意力权重的“数值”不够直观,二维热力图是展示权重分布的核心工具——它能清晰呈现“Query与Key的关联强度”,比如“生成哪个目标词时,关注了源序列的哪个位置”。下面从“数据准备、单头注意力可视化、多头注意力可视化”三个维度,详解热力图的绘制逻辑与代码实现。

5.1 核心概念:注意力权重矩阵与热力图的对应关系

注意力权重矩阵的形状决定了热力图的维度:

  • 单头注意力(如加性注意力):权重矩阵形状为 ( T t × T s ) (T_t × T_s) (Tt×Ts),其中 T t T_t Tt 是目标序列长度(Query数量), T s T_s Ts 是源序列长度(Key数量)。热力图的“行”对应目标序列的每个位置,“列”对应源序列的每个位置,颜色越深表示权重越高。
  • 多头注意力(如Transformer):权重矩阵形状为 ( n u m h e a d s × T t × T s ) (num_heads × T_t × T_s) (numheads×Tt×Ts),每个注意力头对应一张独立的热力图,可展示模型从不同角度关注的信息(比如一个头关注语法,一个头关注语义)。

💡【通俗解释】热力图就像一张“注意力地图”:横轴是你看的原文,纵轴是你正在写的译文,颜色越深,说明模型翻译这个词时越依赖原文那个位置。

以机器翻译任务为例:

  • 源序列(Key):“我”“爱”“中国”“EOS”( T s = 4 T_s=4 Ts=4);
  • 目标序列(Query):“I”“love”“China”( T t = 3 T_t=3 Tt=3);
  • 权重矩阵形状为 ( 3 × 4 ) (3×4) (3×4),热力图中“第1行第1列”(生成“I”时关注“我”)的颜色最深,对应权重0.90。

5.2 数据准备:从模型中提取注意力权重

无论是训练好的模型还是模拟数据,首先需要获取注意力权重矩阵(通常为NumPy数组或PyTorch张量)。以实际场景为例,提取方式如下:

def extract_attention_weights(model, batch_data):
    """
    从训练好的模型中提取注意力权重
    Args:
        model: 包含注意力机制的模型(如Transformer、Seq2Seq)
        batch_data: 批量输入数据(如源序列token ID)
    Returns:
        attn_weights: 注意力权重矩阵,形状 (batch_size, num_heads, T_t, T_s)
    """
    # 设为评估模式,禁用梯度计算
    model.eval()
    with torch.no_grad():
        # 前向传播时,让模型返回注意力权重(需模型支持该输出)
        outputs, attn_weights = model(batch_data, return_attn=True)
        # 转换为NumPy数组(便于可视化)
        return attn_weights.detach().numpy()

# 模拟提取过程(实际使用时替换为真实模型)
# dummy_model = MyTransformerModel()  # 自定义Transformer模型
# dummy_batch = torch.randint(0, 1000, (2, 10))  # (batch_size=2, seq_len=10)
# attn_weights = extract_attention_weights(dummy_model, dummy_batch)

5.3 可视化代码:单头注意力热力图(加性注意力场景)

针对加性注意力的 ( T t × T s ) (T_t × T_s) (Tt×Ts)权重矩阵,绘制“目标序列-源序列”的关联热力图:

def plot_single_head_heatmap(source_words, target_words, attn_weights, save_path):
    """
    绘制单头注意力权重热力图(适合加性注意力、单头缩放点积注意力)
    Args:
        source_words: 源序列词汇列表(Key),如["我", "爱", "中国", "EOS"]
        target_words: 目标序列词汇列表(Query),如["I", "love", "China"]
        attn_weights: 注意力权重矩阵,形状 (T_t, T_s)
        save_path: 图片保存路径(如 "./single_head_attn.png")
    """
    # 创建画布(宽=源序列长度×2,高=目标序列长度×1.5,保证标签不拥挤)
    fig, ax = plt.subplots(figsize=(len(source_words)*2, len(target_words)*1.5))
    
    # 绘制热力图:viridis色系(深绿=高权重,浅绿=低权重)
    im = ax.imshow(attn_weights, cmap='viridis', aspect='auto', vmin=0.0, vmax=1.0)
    
    # 设置坐标轴标签
    ax.set_xticks(range(len(source_words)))
    ax.set_xticklabels(source_words, fontsize=12, ha='center')
    ax.set_yticks(range(len(target_words)))
    ax.set_yticklabels([f'生成 "{word}"时' for word in target_words], fontsize=12, va='center')
    
    # 设置坐标轴标题
    ax.set_xlabel('源序列(Key)', fontsize=13, labelpad=10)
    ax.set_ylabel('目标序列(Query)', fontsize=13, labelpad=10)
    ax.set_title('单头注意力权重分布热力图', fontsize=15, pad=15, fontweight='bold')
    
    # 在每个格子标注权重值(白色文字,保留2位小数)
    for i in range(len(target_words)):
        for j in range(len(source_words)):
            text = ax.text(j, i, f'{attn_weights[i, j]:.2f}',
                          ha='center', va='center', color='white', fontsize=11, fontweight='bold')
    
    # 添加颜色条(解释权重强度)
    cbar = plt.colorbar(im, ax=ax, shrink=0.8)
    cbar.set_label('注意力权重(0=无关注,1=完全关注)', rotation=270, labelpad=20, fontsize=10)
    cbar.set_ticks(np.arange(0.0, 1.1, 0.2))  # 刻度:0.0、0.2、...、1.0
    
    # 调整布局并保存(高分辨率300dpi,避免标签截断)
    plt.tight_layout()
    plt.savefig(save_path, dpi=300, bbox_inches='tight')
    plt.close()
    print(f"✅ 单头注意力热力图已保存至:{save_path}")

# 模拟单头注意力权重(目标序列3个词,源序列4个词)
source_sentence = ["我", "爱", "中国", "EOS"]
target_sentence = ["I", "love", "China"]
single_head_weights = np.array([
    [0.90, 0.05, 0.03, 0.02],  # 生成"I"时的权重
    [0.05, 0.90, 0.03, 0.02],  # 生成"love"时的权重
    [0.05, 0.10, 0.80, 0.05]   # 生成"China"时的权重
])

# 绘制并保存
plot_single_head_heatmap(source_sentence, target_sentence, single_head_weights, "./single_head_attn.png")

📊 单头注意力热力图结果解读

  • 行与列的关联:每一行对应“生成目标词时的注意力分布”,每一列对应“源词被关注的程度”;
  • 颜色与权重:生成“I”时,“我”对应的格子颜色最深(权重0.90);生成“love”时,“爱”对应的格子最深(0.90),完全匹配人类翻译的注意力逻辑;
  • 掩码效果:若源序列存在无效token(如padding),对应列的权重会趋近于0,颜色为最浅的绿色,直观体现掩码的作用。

5.4 可视化代码:多头注意力热力图(Transformer场景)

Transformer的多头注意力会输出多个权重矩阵,需绘制“多子图”展示每个头的关注模式:

def plot_multi_head_heatmap(source_words, target_words, attn_weights, num_heads, save_path):
    """
    绘制多头注意力权重热力图(适合Transformer)
    Args:
        source_words: 源序列词汇列表(Key)
        target_words: 目标序列词汇列表(Query)
        attn_weights: 多头注意力权重矩阵,形状 (num_heads, T_t, T_s)
        num_heads: 注意力头数
        save_path: 图片保存路径
    """
    # 计算子图布局:优先按行排列,每行最多4个头
    n_rows = (num_heads + 3) // 4  # 向上取整(如5个头→2行)
    n_cols = min(num_heads, 4)
    
    # 创建画布(总宽=列数×5,总高=行数×4,保证每个子图清晰)
    fig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols*5, n_rows*4))
    fig.suptitle(f'多头注意力权重分布热力图(共{num_heads}个头)', fontsize=16, y=0.98, fontweight='bold')
    
    # 若只有1行/1列,将axes转为二维数组(避免索引错误)
    if n_rows == 1 and n_cols == 1:
        axes = np.array([[axes]])
    elif n_rows == 1:
        axes = np.expand_dims(axes, 0)
    elif n_cols == 1:
        axes = np.expand_dims(axes, 1)
    
    # 为每个头绘制热力图
    for head in range(num_heads):
        row = head // n_cols
        col = head % n_cols
        ax = axes[row, col]
        
        # 当前头的权重矩阵
        head_weights = attn_weights[head]
        
        # 绘制热力图
        im = ax.imshow(head_weights, cmap='viridis', aspect='auto', vmin=0.0, vmax=1.0)
        
        # 设置子图坐标轴
        ax.set_xticks(range(len(source_words)))
        ax.set_xticklabels(source_words, fontsize=10, ha='center', rotation=45)  # 旋转标签避免重叠
        ax.set_yticks(range(len(target_words)))
        ax.set_yticklabels([f'生成 "{w}"' for w in target_words], fontsize=10, va='center')
        
        # 子图标题
        ax.set_title(f'注意力头 {head+1}', fontsize=12, pad=8, fontweight='bold')
        
        # 标注权重值
        for i in range(len(target_words)):
            for j in range(len(source_words)):
                ax.text(j, i, f'{head_weights[i, j]:.2f}',
                       ha='center', va='center', color='white', fontsize=9)
    
    # 添加全局颜色条(所有子图共享)
    cbar_ax = fig.add_axes([0.92, 0.15, 0.02, 0.7])  # 颜色条位置:右、下、宽、高
    cbar = fig.colorbar(im, cax=cbar_ax)
    cbar.set_label('注意力权重', rotation=270, labelpad=20, fontsize=11)
    cbar.set_ticks(np.arange(0.0, 1.1, 0.2))
    
    # 调整子图间距与全局布局
    plt.tight_layout(rect=[0, 0, 0.9, 0.95])  # 为右侧颜色条留空间
    plt.savefig(save_path, dpi=300, bbox_inches='tight')
    plt.close()
    print(f"✅ 多头注意力热力图已保存至:{save_path}")

# 模拟多头注意力权重(2个头,每个头形状3×4)
num_heads = 2
multi_head_weights = np.array([
    [
        [0.90, 0.05, 0.03, 0.02],  # 头1:生成"I"
        [0.05, 0.90, 0.03, 0.02],  # 头1:生成"love"
        [0.05, 0.10, 0.80, 0.05]   # 头1:生成"China"
    ],
    [
        [0.85, 0.10, 0.03, 0.02],  # 头2:生成"I"(关注"我"+少量"爱")
        [0.10, 0.85, 0.03, 0.02],  # 头2:生成"love"(关注"爱"+少量"我")
        [0.03, 0.07, 0.85, 0.05]   # 头2:生成"China"(更聚焦"中国")
    ]
])

# 绘制并保存
plot_multi_head_heatmap(source_sentence, target_sentence, multi_head_weights, num_heads, "./multi_head_attn.png")

📊 多头注意力热力图结果解读

  • 头的差异化关注:头1更“专注”于单个源词(如生成“I”时仅关注“我”),头2则会轻微关注相邻词(如生成“I”时关注“我”和少量“爱”),体现了多头注意力“从不同角度捕捉关联”的优势;
  • 信息互补性:多个头的关注模式叠加,能让模型同时捕捉语法(如词性关联)、语义(如词义关联)等多维度信息,这也是Transformer比单头注意力效果更好的核心原因。

6. 注意力机制的拓展:多头注意力与应用场景

6.1 多头注意力(Multi-Head Attention):Transformer的核心改进

前文提到的“多头注意力”并非简单的重复计算,而是通过“拆分-并行计算-合并”的逻辑,让模型学习更全面的关联信息,具体流程如下:

  1. 拆分Q/K/V:将Q、K、V按头数 h h h 拆分为 h h h 份,每份维度从 d d d 变为 d / h d/h d/h(如 d = 64 d=64 d=64 h = 8 h=8 h=8,则每份维度为8);
  2. 并行计算:每个头独立计算缩放点积注意力,得到 h h h 个局部注意力结果;
  3. 合并结果:将 h h h 个局部结果拼接,通过线性变换映射回原维度 d d d,得到最终输出。

公式表达为:
MultiHead ( Q , K , V ) = Concat ( head 1 , head 2 , . . . , head h ) ⋅ W O \text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1, \text{head}_2, ..., \text{head}_h) \cdot W^O MultiHead(Q,K,V)=Concat(head1,head2,...,headh)WO
其中 head i = ScaledDotProductAttention ( Q ⋅ W i Q , K ⋅ W i K , V ⋅ W i V ) \text{head}_i = \text{ScaledDotProductAttention}(Q \cdot W_i^Q, K \cdot W_i^K, V \cdot W_i^V) headi=ScaledDotProductAttention(QWiQ,KWiK,VWiV) W i Q 、 W i K 、 W i V 、 W O W_i^Q、W_i^K、W_i^V、W^O WiQWiKWiVWO 为可学习参数。

💡【通俗解释】你可以把多头注意力想象成一个“专家小组”开会:

  • 每个“专家”(头)只看数据的一个侧面(比如一个看语法结构,一个看语义角色,一个看实体关系);
  • 大家各自打分后,汇总意见,再由“组长”( W O W^O WO)整合成最终决策。

这样比单个“全能型选手”更鲁棒,也更能捕捉复杂模式。

6.2 注意力机制的典型应用场景

注意力机制已从NLP扩展到CV、多模态等领域,核心场景包括:

  • NLP领域
    • 机器翻译(如Google翻译用Transformer的多头注意力对齐源/目标序列);
    • 文本摘要(聚焦原文关键句,生成简洁摘要);
    • 问答系统(根据问题Query,从文档Key中提取答案信息)。

🗣️【大白话】在问答系统里,模型看到问题“谁写了《红楼梦》?”,就会在文档里疯狂找“曹雪芹”这个词——这就是注意力在“定位答案”。

  • CV领域
    • 图像 caption(看图写文字,关注图像中的重点物体,如“猫”“沙发”);
    • 目标检测(给图像中不同物体分配注意力权重,优先识别核心目标)。

💡【通俗解释】传统CNN对整张图一视同仁,而带注意力的模型会“盯住”关键区域。比如生成描述时,看到猫就多看猫一眼,而不是平均分配注意力给背景墙。

  • 多模态领域
    • 图文匹配(如小红书图文检索,用注意力关联图片特征与文本关键词);
    • 语音识别(将语音序列与文字序列通过注意力对齐,提升识别准确率)。

🗣️【大白话】当你搜“穿红裙子的女孩”,系统不仅要理解“红裙子”这个词,还要在图片里找到对应区域——注意力就是那个“指哪打哪”的桥梁。


7. 注意力机制 vs RNN:长距离依赖处理能力对比

很多人会问:RNN也能处理长序列,注意力机制到底好在哪?我们用“信息传递路径”和“核心特性”做对比:

特性 RNN(含LSTM/GRU) 注意力机制
信息传递路径 串行传递(词1→词2→…→词n) 直接连接(Query→任意Key,一步到位)
长距离依赖处理 路径长(n步),梯度易消失 路径短(O(1)步),梯度稳定
灵活性 固定向量,无法动态调整重点 动态权重,每个步骤聚焦不同重点
可解释性 隐状态“黑箱”,无法知道关注了啥 注意力权重可视化,可解释性强
计算效率 串行计算,无法并行 矩阵运算并行,效率高(尤其多头注意力)

💡【通俗解释】RNN 像是传纸条:第1个人写完传给第2个,第2个看完再传给第3个……传到第100个人时,纸条早就皱了、字也糊了。
而注意力机制像是开视频会议:每个人都能直接看到原始PPT(编码器的所有隐状态),想看哪页就放大哪页,信息一点不丢!

举个极端例子:处理100个词的句子,RNN要把“词1”的信息传100步才能到解码器,中途早丢光了;而注意力机制生成“词100”对应的目标词时,能直接和“词100”的Key算相似度,信息一点不丢——这也是Transformer能替代RNN成为主流架构的核心原因。


8. 总结:注意力机制是Transformer的“灵魂”

通过本文我们搞懂了:

  1. 注意力机制的诞生是为了解决RNN编码器-解码器的“信息瓶颈”问题,本质是模拟人类“选择性关注”的认知习惯;
  2. 核心逻辑是“算相似度→归一化→加权求和”,有两种主流实现:
    • 加性注意力(适合Q、K维度不同,早期翻译模型常用);
    • 缩放点积注意力(Transformer核心,高效并行,现代标配);
  3. 可视化是理解注意力的关键
    • 单头热力图展示“目标-源”的一对一关联;
    • 多头热力图体现“多角度、多维度关注”;
    • 权重分布直观反映模型的“思考过程”;
  4. 相比RNN,注意力机制在长距离依赖、灵活性、可解释性、计算效率上全面胜出,而多头设计进一步强化了信息捕捉能力。

🌟 更重要的是:注意力机制是Transformer架构的“灵魂”
下一篇我们会看到,Transformer的“自注意力”其实就是“Query、Key、Value都来自同一个序列”的注意力机制——学好本文,你就能轻松理解Transformer的核心!


9. 下期预告

下一篇《自注意力机制:Transformer的核心,让序列“自己关注自己”》,我们将:

  • 拆解自注意力与普通注意力的区别;
  • 实现完整的多头注意力模块;
  • 解读“为什么自注意力能并行计算”;
  • 一步步搭建Transformer的Encoder层!

10. 参考资料

  1. Bahdanau et al., “Neural Machine Translation by Jointly Learning to Align and Translate”(加性注意力原始论文)
  2. Vaswani et al., “Attention Is All You Need”(Transformer与缩放点积注意力原始论文)
  3. PyTorch官方文档:nn.LinearF.softmaxtorch.matmul函数详解
  4. 《深度学习进阶:自然语言处理》——斋藤康毅(注意力机制章节)
  5. Google AI Blog: “Attention Is All You Need” 配套解读
  6. 51CTO博客:《新手小白入门:10分钟搞懂深度学习的“注意力”》(Scaled Dot-Product Attention代码参考)
  7. CSDN文库:《注意力权重分布图绘制方法》(热力图数据准备逻辑参考)

更多推荐