8头自注意力机制入门指南:从原理到PyTorch实战

1次阅读
没有评论

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

image.webp

自注意力机制是 Transformer 架构的核心组件,它通过计算序列元素间的相关性权重实现全局依赖建模。相比 RNN 的串行计算缺陷,自注意力能并行处理所有位置且不受长距离衰减影响。多头设计则让模型同时关注不同子空间的特征模式,如同多视角观察数据。

8 头自注意力机制入门指南:从原理到 PyTorch 实战

多头机制的分头计算原理

8 头自注意力将输入拆分为 8 组独立的注意力计算单元,每组维护独立的 Q /K/ V 投影矩阵。具体维度拆分如下:

  • 输入张量形状:[batch, seq_len, d_model=512]
  • 分头后形状:[batch, seq_len, num_heads=8, head_dim=64]
  • 计算公式:d_head = d_model // num_heads

矩阵运算流程示意图:

[输入] -> Q/K/ V 投影 -> 分头 -> 8 组注意力 -> 拼接 -> 输出投影

PyTorch 完整实现

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

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, num_heads=8):
        super().__init__()
        assert d_model % num_heads == 0
        self.d_head = d_model // num_heads
        self.qkv_proj = nn.Linear(d_model, d_model*3)  # 合并 QKV 投影提升效率
        self.out_proj = nn.Linear(d_model, d_model)

    @torch.jit.script_method
    def forward(self, x: torch.Tensor, mask: torch.Tensor = None):
        # x: [batch, seq_len, d_model]
        batch_size, seq_len, _ = x.shape

        # 投影并分头 [batch, seq_len, num_heads, 3*d_head]
        qkv = self.qkv_proj(x)
        q, k, v = rearrange(qkv, 'b s (n h d) -> n b h s d', n=3, h=self.num_heads).unbind(0)

        # Scaled Dot-Product Attention [batch, num_heads, seq_len, seq_len]
        attn_scores = einsum(q, k, 'b h i d, b h j d -> b h i j') / (self.d_head ** 0.5)
        if mask is not None:
            attn_scores = attn_scores.masked_fill(mask == 0, -1e9)
        attn_weights = torch.softmax(attn_scores, dim=-1)

        # 加权求和并拼接 [batch, seq_len, d_model]
        context = einsum(attn_weights, v, 'b h i j, b h j d -> b h i d')
        context = rearrange(context, 'b h s d -> b s (h d)')
        return self.out_proj(context)

性能优化实践

  1. 显存占用分析
  2. 8 头比单头多消耗约 15% 显存,主要来自中间 attention 矩阵
  3. 建议头数选择 2 的幂次(如 4 /8/16)以利用 GPU 并行特性

  4. FlashAttention 集成

    # 替换原始 softmax 计算
    from flash_attn import flash_attn_qkvpacked
    context = flash_attn_qkvpacked(torch.stack([q,k,v], dim=2),
        dropout_p=0.1,
        causal=self.is_causal
    )

常见问题避坑

  • 梯度消失
  • 初始化 Q / K 投影矩阵方差设为 1 /√d_head
  • 使用 Xavier 初始化时选择gain=0.02

  • 维度错误

  • 分头后务必检查d_head * num_heads == d_model
  • 拼接时注意最后一维必须是 h*d 的顺序

拓展思考

  1. 头部分析实验设计
  2. 对不同头计算的 attention 矩阵进行聚类分析
  3. 观察特定头在语法结构(如括号匹配)和语义角色(主谓宾)上的激活模式

  4. 长序列处理策略

  5. 采用相对位置编码(如 ALiBi)替代绝对位置编码
  6. 对超长序列使用局部注意力窗口(如滑动 128 个 token)

通过上述实现,我们完整构建了支持 mask 处理的 8 头自注意力层。建议读者使用 PyTorch 的 autograd profiler 分析各环节耗时,并尝试可视化不同头的注意力模式来直观理解其工作原理。

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