深入解析Transformer中的自注意力机制:计算过程与核心优势

1次阅读
没有评论

共计 2540 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

传统 RNN 的局限性

在自然语言处理领域,传统 RNN(循环神经网络)在处理长序列时面临着显著的挑战。RNN 通过隐藏状态传递信息,理论上可以捕捉任意长度的依赖关系。然而在实际应用中,RNN 存在两个主要问题:

深入解析 Transformer 中的自注意力机制:计算过程与核心优势

  1. 梯度消失 / 爆炸问题 :随着序列长度增加,反向传播时梯度会指数级衰减或增长,导致模型难以学习长期依赖关系。
  2. 顺序计算限制 :RNN 必须按时间步顺序处理输入,无法充分利用现代 GPU 的并行计算能力。

这些局限性促使研究者寻找更好的序列建模方法,自注意力机制应运而生。

自注意力机制详解

1. Query、Key、Value 矩阵

自注意力机制的核心思想是让序列中的每个位置都能直接关注到其他所有位置。为了实现这一点,首先需要将输入序列转换为三种表示:

  1. Query(查询):表示当前关注的位置
  2. Key(键):表示被查询的位置
  3. Value(值):包含实际的信息内容

数学上,这些矩阵通过可学习的权重矩阵对输入进行线性变换得到:

Q = X @ W_Q  # Query 矩阵
K = X @ W_K  # Key 矩阵
V = X @ W_V  # Value 矩阵 

其中 X 是输入序列(形状为 [序列长度, 特征维度]),W_Q、W_K、W_V 是可训练参数。

2. 注意力分数计算

计算每个 Query 与所有 Key 的点积,得到注意力分数。分数越高表示两个位置的相关性越强:

attention_scores = Q @ K.T  # 形状 [序列长度, 序列长度]

为了防止点积结果过大导致 softmax 梯度消失,通常会对分数进行缩放(除以 Key 维度的平方根):

d_k = K.shape[-1]  # Key 的维度
scaled_scores = attention_scores / (d_k ** 0.5)

3. Softmax 归一化

对每一行(对应一个 Query)应用 softmax 函数,将分数转换为概率分布:

attention_weights = torch.softmax(scaled_scores, dim=-1)

4. 上下文向量生成

最后用注意力权重对 Value 矩阵加权求和,得到每个位置的输出表示:

output = attention_weights @ V  # 形状 [序列长度, 特征维度]

5. 多头注意力

为了捕捉不同子空间的信息,实际应用中会使用多头注意力:

  1. 将 Q、K、V 分别拆分到 h 个头
  2. 在每个头上独立计算注意力
  3. 拼接所有头的输出并通过线性变换

代码实现

以下是一个简化版的自注意力层实现(使用 PyTorch):

import torch
import torch.nn as nn
import torch.nn.functional as F

class SelfAttention(nn.Module):
    def __init__(self, embed_size, heads):
        super(SelfAttention, self).__init__()
        self.embed_size = embed_size
        self.heads = heads
        self.head_dim = embed_size // heads

        assert self.head_dim * heads == embed_size, "Embed size needs to be divisible by heads"

        self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)
        self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)
        self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)
        self.fc_out = nn.Linear(heads * self.head_dim, embed_size)

    def forward(self, values, keys, query, mask=None):
        N = query.shape[0]  # 批大小
        value_len, key_len, query_len = values.shape[1], keys.shape[1], query.shape[1]

        # 拆分输入到多个头
        values = values.reshape(N, value_len, self.heads, self.head_dim)
        keys = keys.reshape(N, key_len, self.heads, self.head_dim)
        queries = query.reshape(N, query_len, self.heads, self.head_dim)

        # 计算注意力分数
        energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])

        if mask is not None:
            energy = energy.masked_fill(mask == 0, float("-1e20"))

        # 缩放和 softmax
        attention = torch.softmax(energy / (self.embed_size ** (1/2)), dim=3)

        # 加权求和
        out = torch.einsum("nhql,nlhd->nqhd", [attention, values])
        out = out.reshape(N, query_len, self.heads * self.head_dim)

        return self.fc_out(out)

技术优势分析

  1. 并行计算 :与 RNN 不同,自注意力可以同时计算所有位置的表示,充分利用 GPU 并行能力。
  2. 长距离依赖 :任何两个位置的距离都是 O(1),解决了 RNN 的长距离依赖问题。
  3. 可解释性 :注意力权重直观显示模型关注了输入的哪些部分。

计算复杂度方面,自注意力是 O(n²)(n 为序列长度),比 RNN 的 O(n) 更高,但现代硬件对矩阵运算有良好优化,实际运行时间可能更短。

避坑指南

  1. 注意力分数缩放 :不缩放会导致 softmax 进入饱和区,梯度消失。
  2. 处理不同序列长度 :使用 mask 忽略填充位置。
  3. 内存优化 :对于极长序列,可以考虑稀疏注意力或分块计算。

思考题

自注意力机制最初为 NLP 设计,但它可以应用于任何需要建模元素间关系的任务。例如:

  • 计算机视觉(图像补全、目标检测)
  • 推荐系统(用户 - 商品交互建模)
  • 时序预测(传感器数据)

关键在于如何定义合适的 Query、Key 和 Value 表示。

正文完
 0
评论(没有评论)