别再死记公式了!用NumPy和PyTorch手搓Self-Attention,搞懂Transformer核心
从零手搓Self-Attention:用NumPy和PyTorch拆解Transformer核心
第一次看到Self-Attention的公式时,那些Q、K、V矩阵和复杂的Softmax计算确实让人望而生畏。但当我真正动手实现它时,才发现这些看似复杂的操作背后,隐藏着极其优雅的设计思想。本文将带你用两种方式实现Self-Attention——先用NumPy理解计算本质,再用PyTorch掌握工程实现,最终让你不仅能写出代码,更能"发明"出这套机制。
1. 为什么需要Q、K、V?
在传统的序列建模中,每个位置的信息处理都是孤立的。而Self-Attention的核心突破在于:让每个位置都能直接关注到序列中所有相关的位置。这种全局视角带来了惊人的表达能力,但也需要一套巧妙的机制来管理这种信息流动。
想象你在阅读一篇文章时,大脑会自然地在不同段落间建立联系——某些词需要参考前文的定义,某些句子需要呼应后文的结论。Q(query)、K(key)、V(value)正是模拟这种动态关联的三大组件:
- Query:当前位置"想问"的问题
- Key:其他位置"能回答"的能力
- Value:实际传递的信息内容
它们的矩阵乘法本质上是在计算:当前query与所有key的匹配程度,然后按此权重聚合values。下面这段伪代码揭示了核心逻辑:
attention_scores = query @ key.T # 计算匹配度
attention_weights = softmax(attention_scores / sqrt(d_k)) # 归一化权重
output = attention_weights @ value # 加权聚合
2. NumPy实现:剥开Self-Attention的数学内核
让我们先用NumPy实现最基础的Self-Attention,避开框架的封装,直面计算本质。假设我们有一个包含32个token的序列,每个token的特征维度是256:
import numpy as np
from numpy.random import randn
# 输入序列:256维特征,32个token
d, n = 256, 32
x = randn(d, n) # 形状(256, 32)
2.1 生成Q、K、V矩阵
首先需要三个投影矩阵,将原始输入分别映射到query、key和value空间:
# 初始化投影矩阵
w_q = randn(d, d) # query权重
w_k = randn(d, d) # key权重
w_v = randn(d, d) # value权重
# 计算Q、K、V
q = w_q @ x # (256,256)@(256,32)→(256,32)
k = w_k @ x # 同上
v = w_v @ x # 同上
这里的关键理解点:投影不是降维而是特征重组。虽然维度保持不变,但经过投影后,每个位置的向量被赋予了特定的语义角色。
2.2 注意力分数计算
接下来计算query和key的匹配程度:
A = k.T @ q # (32,256)@(256,32)→(32,32)
A /= np.sqrt(d) # 缩放防止梯度消失
这个32x32的矩阵A就是注意力分数矩阵,其中A[i,j]表示第i个token对第j个token的关注程度。缩放因子√d是关键技巧——当维度d很大时,点积结果会变得极大,导致softmax饱和。
2.3 Softmax归一化
将原始分数转换为概率分布:
def softmax(x):
e_x = np.exp(x - np.max(x)) # 防溢出
return e_x / e_x.sum(axis=0)
A_hat = softmax(A) # 形状(32,32)
现在A_hat的每列和为1,表示每个query对所有key的注意力分布。有趣的是,这种归一化是逐列进行的,因为每个query需要独立决定关注哪些key。
2.4 加权聚合Value
最后用注意力权重聚合values:
output = v @ A_hat # (256,32)@(32,32)→(256,32)
这个输出就是Self-Attention的最终结果——每个位置的新表示,都是全局信息的有选择聚合。整个过程可以用下表总结:
| 步骤 | 操作 | 输入形状 | 输出形状 | 物理意义 |
|---|---|---|---|---|
| 投影 | Q=W_q·X | (256,32) | (256,32) | 生成查询向量 |
| 投影 | K=W_k·X | (256,32) | (256,32) | 生成键向量 |
| 投影 | V=W_v·X | (256,32) | (256,32) | 生成值向量 |
| 匹配 | A=Kᵀ·Q | (256,32)×(256,32) | (32,32) | 计算注意力分数 |
| 归一化 | Â=softmax(A/√d) | (32,32) | (32,32) | 生成注意力权重 |
| 聚合 | Output=V·Â | (256,32)×(32,32) | (256,32) | 生成最终表示 |
3. PyTorch实现:工程化的Self-Attention模块
理解了数学本质后,我们来看PyTorch如何实现可训练的Self-Attention模块。与NumPy版本相比,这里需要处理批量输入和自动微分。
3.1 定义Self-Attention层
import torch
import torch.nn as nn
from math import sqrt
class SelfAttention(nn.Module):
def __init__(self, input_dim, dim_k, dim_v):
super().__init__()
self.dim_k = dim_k
self.q = nn.Linear(input_dim, dim_k) # query投影
self.k = nn.Linear(input_dim, dim_k) # key投影
self.v = nn.Linear(input_dim, dim_v) # value投影
self.softmax = nn.Softmax(dim=-1)
def forward(self, x):
# x形状: (batch_size, seq_len, input_dim)
Q = self.q(x) # (batch, seq, dim_k)
K = self.k(x) # (batch, seq, dim_k)
V = self.v(x) # (batch, seq, dim_v)
# 计算注意力分数
scores = torch.bmm(Q, K.transpose(1,2)) / sqrt(self.dim_k)
# scores形状: (batch, seq, seq)
# 计算注意力权重
attn_weights = self.softmax(scores)
# 加权聚合
output = torch.bmm(attn_weights, V) # (batch, seq, dim_v)
return output
几个关键区别:
- 使用
nn.Linear代替手动矩阵乘法,自动处理参数初始化与梯度计算 torch.bmm用于批量矩阵乘法(batch matrix multiply)- 维度可以灵活配置,不要求Q、K、V同维度
3.2 实际使用示例
让我们用随机数据测试这个模块:
batch_size = 4
seq_len = 10
input_dim = 64
dim_k = 128
dim_v = 160
# 生成随机输入
x = torch.randn(batch_size, seq_len, input_dim)
# 初始化Self-Attention层
self_attn = SelfAttention(input_dim, dim_k, dim_v)
# 前向传播
output = self_attn(x)
print(output.shape) # 输出: torch.Size([4, 10, 160])
在实际Transformer中,这种Self-Attention会进一步扩展为多头机制,让模型能够同时关注不同子空间的信息。
4. 从实现到理解:关键问题解析
通过上述实现,我们可以深入回答一些常见困惑:
4.1 为什么需要三个不同的矩阵?
Q、K、V分离的设计带来了三个核心优势:
- 解耦功能:查询、匹配、信息传递各司其职
- 灵活建模:允许不同语义空间的投影
- 表达能力强:比单一矩阵能捕获更复杂的关系
4.2 缩放因子√d_k的数学意义
缩放操作看似简单,实则解决了一个关键问题:当维度d_k很大时,点积结果的方差会增大,导致softmax进入饱和区(某些位置概率接近1,其余接近0)。缩放保持梯度稳定,使模型更容易训练。
4.3 注意力矩阵的可视化意义
假设我们处理一句话:"The animal didn't cross the street because it was too tired"。下图展示了"it"这个词的注意力权重分布:
it → [0.02, 0.01, 0.88, ..., 0.03, 0.01]
可以看到,模型正确地让"it"高度关注"animal"(权重0.88),这正是语言理解的体现。
5. 工程实践中的常见陷阱
在实际项目中,有几点需要特别注意:
-
数值稳定性:
- Softmax计算需要防溢出
- 混合精度训练时注意尺度
-
内存消耗:
- 注意力矩阵是O(n²)的,长序列会爆内存
- 解决方案:内存高效的注意力实现
-
初始化策略:
# 通常用较小的初始化防止训练初期饱和 nn.init.xavier_uniform_(self.q.weight, gain=0.02) nn.init.xavier_uniform_(self.k.weight, gain=0.02) nn.init.xavier_uniform_(self.v.weight, gain=0.02) -
批处理效率:
- 使用
torch.bmm而非循环 - 当序列长度不一致时需要padding或mask
- 使用
6. 扩展思考:Self-Attention的变体与进化
基础Self-Attention只是起点,现代Transformer使用了许多改进版本:
-
多头注意力:并行多个注意力头,捕获不同关系
class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads == 0 self.d_k = d_model // num_heads self.num_heads = num_heads # 投影矩阵初始化... -
相对位置编码:解决绝对位置编码的局限性
-
稀疏注意力:降低O(n²)计算复杂度
在视觉Transformer中,还需要处理二维空间关系;在图神经网络中,注意力机制用于学习节点间的重要性。这些变体都建立在本文介绍的基础之上。
更多推荐


所有评论(0)