从零实现自注意力


文档摘要

从零实现自注意力 本节摘要:自注意力(Self-Attention)是一张查找表,表里每一个词都在问「谁对我重要?」,然后从训练中学到答案。RNN 把信息一个一个 token 地压缩,长程依赖被挤成渣;2017 年的论文把问题推到极致——如果唯一的机制就是注意力,没有循环、没有卷积,会怎样?答案就是 Transformer。本节将带你只用 NumPy 从零实现缩放点积注意力(Scaled Dot-Product Attention):三个可学习投影矩阵把每个 token 映射成查询(Query)、键(Key)、值(Value),注意力矩阵记录 token 两两之间的关系,softmax 把分数变成权重,加权求和得到输出。

从零实现自注意力

本节摘要:自注意力(Self-Attention)是一张查找表,表里每一个词都在问「谁对我重要?」,然后从训练中学到答案。RNN 把信息一个一个 token 地压缩,长程依赖被挤成渣;2017 年的论文把问题推到极致——如果唯一的机制就是注意力,没有循环、没有卷积,会怎样?答案就是 Transformer。本节将带你只用 NumPy 从零实现缩放点积注意力(Scaled Dot-Product Attention):三个可学习投影矩阵把每个 token 映射成查询(Query)、键(Key)、值(Value),注意力矩阵记录 token 两两之间的关系,softmax 把分数变成权重,加权求和得到输出。你会亲手看到为什么要在 Q·Kᵀ 后除以 √dk(防止 softmax 饱和),以及整条流水线如何浓缩成一行公式:softmax(QKᵀ/√dk)V。读完本节,你能把这条公式逐项解释清楚,并跑通一个会打印注意力热力图的最小实现。

学习目标

阅读完本节,你应当能够:

  1. 只用 NumPy 从零实现缩放点积自注意力,包含 Q/K/V 投影与 softmax 加权求和。
  2. 解释注意力矩阵如何刻画 token 两两关系,并说清为什么除以 √dk 能防止 softmax 饱和。
  3. 用**因果掩码(Causal Mask)**把双向注意力改写成自回归(解码器式)注意力。
  4. 把整套流水线浓缩成一行公式,并逐项解释几何与概率含义。

一、问题与直觉

RNN 一次处理一个 token,等你走到第 50 个 token 时,第 1 个 token 的信息已经被挤过 50 道压缩。2014 年的 Bahdanau 注意力论文给出了修法:让解码器回头看每一个编码器位置,自己决定哪几个重要。但它仍是绑在 RNN 上的。2017 年的论文问了一个更尖锐的问题:如果注意力是唯一的机制呢?

自注意力让序列中每个位置在一次并行步骤内同时关注所有其他位置。这就是 Transformer 快、能 scale、且统治一切的原因。

数据库查找的类比

把注意力想象成一次「软」数据库查找:

传统数据库: 查询:「法国的首都」 --> 精确匹配 --> 「巴黎」 注意力: 查询:「法国的首都」 --> 与所有键算相似度 --> 所有值的加权混合

每个 token 生成三个向量:

  • 查询(Query, Q):「我在找什么?」
  • 键(Key, K):「我包含什么?」
  • 值(Value, V):「如果我被选中,我提供什么信息?」

查询与所有键的点积产生注意力分数,分数高表示「这个键匹配我的查询」。这些分数对值做加权,输出就是值的加权和。

三步成型

每个 token 的嵌入通过三个可学习权重矩阵投影:

X = [x1, x2, ..., xn] 形状 (n, d) 输入嵌入 Wq 形状 (d, dk) Q = X @ Wq 形状 (n, dk) 每个token的查询 Wk 形状 (d, dk) K = X @ Wk 形状 (n, dk) 每个token的键 Wv 形状 (d, dv) V = X @ Wv 形状 (n, dv) 每个token的值 分数 = Q @ K^T 形状 (n, n) 缩放 = 分数 / sqrt(dk) 权重 = softmax(缩放) 输出 = 权重 @ V

一行公式浓缩全部:

Attention(Q, K, V) = softmax( Q @ K^T / sqrt(dk) ) @ V

为什么除以 √dk?

点积随维度 dk 增长。若 dk=64,点积可能到几十,把 softmax 推进梯度几乎为零的饱和区。除以 √dk 把数值压回 softmax 能产生有效梯度的区间。这是一个数值稳定性技巧,不改变注意力「谁匹配谁」的排序,只改变概率分布的尖锐程度。

⚠️ 常见误解:有人以为 √dk 是为了归一化数值范围防止溢出。真正的目的是防止 softmax 饱和——饱和后梯度消失,模型学不动。这两件事后果不同。

注意力矩阵长什么样

分数矩阵 Q @ Kᵀn×n 的,每一行是一个 token 对全序列的注意力。softmax 按行归一化后,行内加起来等于 1。看一个 5 token 的小例子:

k1 k2 k3 k4 k5 +-----+-----+-----+-----+-----+ q1 | 2.1 | 0.3 | 0.1 | 0.8 | 0.2 | <- q1 对每个键的关注度 +-----+-----+-----+-----+-----+ q2 | 0.4 | 1.9 | 0.7 | 0.1 | 0.3 | ... softmax 后 q1 这一行:[0.52, 0.09, 0.07, 0.14, 0.08] (和约为 1.0)

最终输出:output_1 = 0.52·v1 + 0.09·v2 + 0.07·v3 + 0.14·v4 + 0.08·v5

二、从零实现

Step 1:从零写 softmax

为数值稳定,先减去最大值再取指数。完整代码见原课程 phases/07-transformers-deep-dive/02-self-attention-from-scratch/code/main.py

import numpy as np def softmax(x): shifted = x - np.max(x, axis=-1, keepdims=True) exp_x = np.exp(shifted) return exp_x / np.sum(exp_x, axis=-1, keepdims=True)

Step 2:缩放点积注意力

核心函数,接收 Q、K、V,返回输出和权重矩阵:

def scaled_dot_product_attention(Q, K, V): dk = Q.shape[-1] scores = Q @ K.T / np.sqrt(dk) weights = softmax(scores) output = weights @ V return output, weights

Step 3:带可学习投影的自注意力类

用类 Xavier 缩放初始化 Wq/Wk/Wv:

class SelfAttention: def __init__(self, d_model, dk, dv, seed=42): rng = np.random.default_rng(seed) scale = np.sqrt(2.0 / (d_model + dk)) self.Wq = rng.normal(0, scale, (d_model, dk)) self.Wk = rng.normal(0, scale, (d_model, dk)) scale_v = np.sqrt(2.0 / (d_model + dv)) self.Wv = rng.normal(0, scale_v, (d_model, dv)) def forward(self, X): Q = X @ self.Wq K = X @ self.Wk V = X @ self.Wv return scaled_dot_product_attention(Q, K, V)

Step 4:在一句话上跑起来

造假的嵌入,看注意力权重。把权重映射成 ASCII 字符( ░▒▓█)就能快速画一张热力图——对角线通常偏亮(每个 token 多少关注自己),但跨 token 的亮块会揭示出哪些词在「互相看」。

💡 观察重点:同一个 SelfAttention 实例喂两句话长相同但内容不同的句子,对角线亮块的位置会变——注意力是内容驱动的,不是位置驱动的(位置信息要靠下一节的位置编码注入)。

三、框架对比:PyTorch 的 nn.MultiheadAttention

PyTorch 的 nn.MultiheadAttention 做的事和我们手写的一模一样,外加多头切分和输出投影:

import torch.nn as nn mha = nn.MultiheadAttention(embed_dim=8, num_heads=2, batch_first=True) X = torch.randn(1, 6, 8) # (batch, seq, d_model) output, attn_weights = mha(X, X, X) # Q=K=V=X 即自注意力

关键差别在于多头:并行跑多个注意力函数,每个用 dk = d_model / n_heads 大小的独立投影,最后拼接。这让模型同时关注不同类型的关系(下一节详述)。

维度 手写 NumPy PyTorch nn.MultiheadAttention
计算核心 Q@K.T/√dk→softmax→@V 完全相同
多头 需自己切分拼接 内置 num_heads
反向传播 无(需手写) autograd 自动
因果掩码 自己加 -inf attn_mask

四、可复用产物

原课程产出 outputs/prompt-attention-explainer.md:一个解释提示词,用「数据库软查找」的类比向非专业读者解释注意力机制。喂给它任何一段注意力代码或论文片段,它产出一段通俗解释,适合做教学或文档。

五、练习

  1. (Easy)scaled_dot_product_attention 加一个可选的 mask 参数,在 softmax 前把掩码位置设成 -inf。这就是因果(解码器)掩码的实现方式。
  2. (Medium) 从零实现多头注意力:把 Q、K、V 切成 n_heads 份,每份单独跑注意力,拼接后过一个最终投影矩阵 Wo
  3. (Hard) 取两个长度相同、内容不同的句子,过同一个 SelfAttention 实例,比较它们的注意力模式。什么变了?什么没变?为什么对角线不一定总是最亮?

本节要点回顾

  1. 自注意力 = 软数据库查找:每个 token 生成 Q(找什么)、K(含什么)、V(提供什么),查询与所有键的点积产生分数。
  2. 三个投影矩阵 Wq/Wk/Wv 是唯一的可学习参数,把同一个嵌入投影成三种角色。
  3. 注意力矩阵是 n×n:行 = 一个 token 对全序列的关注度,softmax 按行归一化成概率分布。
  4. 除以 √dk 防 softmax 饱和:点积随维度增长,不缩放会掉进梯度消失区——这是数值稳定性的关键,不是普通的归一化。
  5. 一行公式浓缩全部:Attention(Q,K,V) = softmax(QKᵀ/√dk)V
  6. 输出是值的加权和:output_i = Σⱼ a_ij·v_j,权重 a_ij 由相似度决定。
  7. 注意力是内容驱动的:同一个实例对不同句子给出不同的注意力模式——这正是它需要位置编码(下一节)的原因。

下一节,我们将把单个注意力头升级成多头注意力(Multi-Head Attention)——并行跑多组独立的 Q/K/V 投影,让模型同时捕捉不同类型的关系。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U