Transformer自注意力机制(self-attention)原理解析与计算过程详解

1次阅读
没有评论

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

image.webp

背景介绍

自注意力机制 (self-attention) 是 Transformer 模型的核心组件,它彻底改变了自然语言处理 (NLP) 领域的格局。相较于传统的 RNN 和 CNN,自注意力机制能够直接建模输入序列中任意两个位置之间的关系,无论它们相距多远。这种能力使得 Transformer 在处理长距离依赖、并行计算等方面展现出巨大优势,成为 BERT、GPT 等现代 NLP 模型的基础。

Transformer 自注意力机制 (self-attention) 原理解析与计算过程详解

数学原理

自注意力机制的计算过程可以分为以下几个步骤:

  1. 输入表示:对于输入序列 X∈R^(n×d_model),其中 n 是序列长度,d_model 是嵌入维度。

  2. QKV 投影 :通过三个不同的线性变换将输入映射到查询(Query)、键(Key) 和值 (Value) 空间:

  3. Q = XW_Q, W_Q∈R^(d_model×d_k)
  4. K = XW_K, W_K∈R^(d_model×d_k)
  5. V = XW_V, W_V∈R^(d_model×d_v)

  6. 注意力分数计算:通过点积计算查询和键的相似度:

  7. Attention(Q,K,V) = softmax(QK^T/√d_k)V
    其中√d_k 是缩放因子,防止点积结果过大导致 softmax 梯度消失。

  8. 多头注意力:为了捕捉不同子空间的信息,通常会使用多头注意力:

  9. MultiHead(Q,K,V) = Concat(head_1,…,head_h)W_O
    其中 head_i = Attention(QW_Q^i,KW_K^i,VW_V^i)

代码实现

下面是用 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), "Embedding size needs to be divisible by heads"

        # 定义 QKV 投影矩阵
        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):
        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"))

        # 缩放点积注意力
        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)

        # 最终线性变换
        out = self.fc_out(out)
        return out

优势分析

与 RNN 和 CNN 相比,自注意力机制具有以下优势:

  1. 并行计算:不像 RNN 那样需要顺序处理,自注意力可以同时计算所有位置的表示。

  2. 长距离依赖:直接建模任意两个位置的关系,不受距离限制,解决了 RNN 的梯度消失问题。

  3. 可解释性:通过注意力权重可以直观地看到模型关注了输入的哪些部分。

  4. 灵活性:可以轻松处理可变长度输入,不需要像 CNN 那样固定卷积核大小。

实际应用

在文本分类任务中,我们可以这样应用自注意力机制:

  1. 首先将输入文本通过嵌入层转换为向量表示。

  2. 然后通过多个自注意力层提取上下文相关的特征表示。

  3. 最后通过全连接层进行分类预测。

关键实现代码如下:

class TransformerClassifier(nn.Module):
    def __init__(self, vocab_size, embed_size, num_classes, heads):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_size)
        self.attention = SelfAttention(embed_size, heads)
        self.fc = nn.Linear(embed_size, num_classes)

    def forward(self, x):
        embedded = self.embedding(x)
        attended = self.attention(embedded, embedded, embedded, None)
        out = self.fc(attended.mean(dim=1))
        return out

避坑指南

在训练自注意力模型时,需要注意以下问题:

  1. 梯度爆炸:由于自注意力计算涉及矩阵乘法,梯度可能变得很大。解决方法:
  2. 使用梯度裁剪
  3. 适当的初始化(如 Xavier 初始化)

  4. 计算资源消耗:自注意力的计算复杂度是 O(n^2),对于长序列会消耗大量内存。解决方法:

  5. 使用稀疏注意力
  6. 分块处理长序列

  7. 过拟合:特别是当训练数据较少时。解决方法:

  8. 使用 dropout
  9. 添加 L2 正则化
  10. 数据增强

思考题

  1. 如何修改自注意力机制使其能够处理超过训练时见过的序列长度?

  2. 在多模态任务(如图文匹配)中,如何设计跨模态的自注意力机制?

  3. 自注意力机制在计算效率方面还有哪些优化空间?

总结

自注意力机制通过查询、键和值的交互,实现了强大的上下文建模能力。它不仅解决了传统序列模型的诸多限制,还为 NLP 领域带来了革命性的进步。理解自注意力机制的原理和实现,是掌握现代深度学习模型的重要基础。希望通过本文的讲解,读者能够更深入地理解这一关键技术,并在实际项目中灵活应用。

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