2D多头自注意力计算流程图解:从零实现Transformer核心模块

1次阅读
没有评论

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

image.webp

1. Transformer 与自注意力机制简介

2017 年诞生的 Transformer 架构彻底改变了 NLP 领域,其核心创新就是 自注意力机制。与传统 RNN 不同,自注意力能直接捕捉序列中任意两个元素的关系,解决了长距离依赖问题。举个简单例子:

2D 多头自注意力计算流程图解:从零实现 Transformer 核心模块

  • 在句子 ”The animal didn’t cross the street because it was too tired” 中,自注意力能自动发现 ”it” 与 ”animal” 的高关联度,而无需像 RNN 那样逐步传递信息

2. 2D 多头注意力分步拆解

2.1 整体计算流程图

graph TD
    A[输入序列 X] --> B[线性投影得到 Q,K,V]
    B --> C[拆分为多头 Q,K,V]
    C --> D[缩放点积注意力计算]
    D --> E[多头结果拼接]
    E --> F[最终输出]

2.2 关键公式与解释

  1. QKV 投影
    $$\begin{aligned}
    Q = XW_Q, \quad K = XW_K, \quad V = XW_V
    \end{aligned}$$

  2. 这三个矩阵的维度通常为(seq_len, d_model)

  3. 投影权重 $W_Q/W_K/W_V$ 是可训练参数

  4. 多头拆分(以头数 h = 8 为例):
    $$\text{MultiHead}(Q,K,V) = \text{Concat}(head_1,…,head_h)W_O$$

  5. 使用 einops.rearrange 实现优雅的维度变换:

    q = rearrange(q, 'b s (h d) -> b h s d', h=self.num_heads)

  6. 缩放点积注意力
    $$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$

  7. 缩放因子 $\sqrt{d_k}$ 防止梯度消失

  8. 计算复杂度为 $O(n^2)$,n 为序列长度

3. PyTorch 完整实现

import torch
import torch.nn as nn
from einops import rearrange

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, num_heads=8):
        super().__init__()
        self.d_model = d_model
        self.num_heads = num_heads
        assert d_model % num_heads == 0, "d_model 必须能被 num_heads 整除"

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

    def forward(self, x, mask=None):
        b, s, _ = x.shape  # batch_size, seq_len, d_model

        # 1. 线性投影
        q = self.w_q(x)  # (b,s,d)
        k = self.w_k(x)
        v = self.w_v(x)

        # 2. 拆分为多头
        q = rearrange(q, 'b s (h d) -> b h s d', h=self.num_heads)
        k = rearrange(k, 'b s (h d) -> b h s d', h=self.num_heads)
        v = rearrange(v, 'b s (h d) -> b h s d', h=self.num_heads)

        # 3. 缩放点积注意力
        scores = torch.matmul(q, k.transpose(-2,-1)) / (self.d_model ** 0.5)
        if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)
        attn = torch.softmax(scores, dim=-1)

        # 4. 多头拼接
        output = torch.matmul(attn, v)  # (b,h,s,d)
        output = rearrange(output, 'b h s d -> b s (h d)')
        return self.w_o(output)

4. 性能优化实战

4.1 计算复杂度分析

  • 空间复杂度:存储 $QK^T$ 矩阵需要 $O(n^2)$ 内存
  • 处理 1024 长度序列时,显存占用已达 1GB(float32)

4.2 头数选择经验

头数 计算速度 显存占用 效果
4 最快 最低 一般
8 适中 中等 推荐
16 较慢 较高 提升有限

5. 避坑指南

5.1 梯度爆炸处理

  • 出现 NaN 值时,尝试调大缩放因子:
    scale = (d_model / num_heads) ** 0.5  # 可调整为 1.0~2.0

5.2 变长序列技巧

# 创建 padding mask 示例
mask = (x != 0).unsqueeze(1).unsqueeze(2)  # (b,1,1,s)

6. 拓展思考:Flash Attention

最新提出的 Flash Attention 通过以下方式优化:
1. 分块计算避免存储完整 $QK^T$ 矩阵
2. 融合 kernel 减少内存访问
3. 支持半精度计算

实现示例:

from flash_attn import flash_attention
output = flash_attention(q, k, v)

建议读者尝试将本文实现迁移到 Flash Attention,比较两者的速度和内存差异。

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