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

多头机制的分头计算原理
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)
性能优化实践
- 显存占用分析:
- 8 头比单头多消耗约 15% 显存,主要来自中间 attention 矩阵
-
建议头数选择 2 的幂次(如 4 /8/16)以利用 GPU 并行特性
-
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的顺序
拓展思考
- 头部分析实验设计:
- 对不同头计算的 attention 矩阵进行聚类分析
-
观察特定头在语法结构(如括号匹配)和语义角色(主谓宾)上的激活模式
-
长序列处理策略:
- 采用相对位置编码(如 ALiBi)替代绝对位置编码
- 对超长序列使用局部注意力窗口(如滑动 128 个 token)
通过上述实现,我们完整构建了支持 mask 处理的 8 头自注意力层。建议读者使用 PyTorch 的 autograd profiler 分析各环节耗时,并尝试可视化不同头的注意力模式来直观理解其工作原理。
正文完
发表至: 未分类
近一天内
