从零手搓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

几个关键区别:

  1. 使用nn.Linear代替手动矩阵乘法,自动处理参数初始化与梯度计算
  2. torch.bmm用于批量矩阵乘法(batch matrix multiply)
  3. 维度可以灵活配置,不要求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分离的设计带来了三个核心优势:

  1. 解耦功能:查询、匹配、信息传递各司其职
  2. 灵活建模:允许不同语义空间的投影
  3. 表达能力强:比单一矩阵能捕获更复杂的关系

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. 工程实践中的常见陷阱

在实际项目中,有几点需要特别注意:

  1. 数值稳定性

    • Softmax计算需要防溢出
    • 混合精度训练时注意尺度
  2. 内存消耗

    • 注意力矩阵是O(n²)的,长序列会爆内存
    • 解决方案:内存高效的注意力实现
  3. 初始化策略

    # 通常用较小的初始化防止训练初期饱和
    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)
    
  4. 批处理效率

    • 使用torch.bmm而非循环
    • 当序列长度不一致时需要padding或mask

6. 扩展思考:Self-Attention的变体与进化

基础Self-Attention只是起点,现代Transformer使用了许多改进版本:

  1. 多头注意力:并行多个注意力头,捕获不同关系

    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
            # 投影矩阵初始化...
    
  2. 相对位置编码:解决绝对位置编码的局限性

  3. 稀疏注意力:降低O(n²)计算复杂度

在视觉Transformer中,还需要处理二维空间关系;在图神经网络中,注意力机制用于学习节点间的重要性。这些变体都建立在本文介绍的基础之上。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐