4.5 注意力机制:让张量互相打分


文档摘要

4.5 注意力机制:让张量互相打分 本节摘要:注意力机制把"哪个位置重要"变成可微分的打分问题:每个位置的查询向量与所有位置的键向量算相似度,softmax 归一成权重,再对值向量加权求和。它取消了 RNN 的逐步传递——任意两位置直接交互、全部时步一次并行算完,长序列的调度成本结构被彻底改写。本节拆解查询、键、值三元组与缩放点积的每一步,手写一个最小注意力层,并说明它如何嵌入编码器-解码器结构与自注意力。这是通向现代大模型结构族的门槛章节。 读完你应当能做到 阅读完本节,你应当能够: 说清查询、键、值三者的角色与缩放点积公式每一步的形状; 手写一个多头注意力的最小实现并用形状测试验证; 解释自注意力与交叉注意力在调度上的差别; 说明注意力相对 RNN 的并行优势与计算复杂度代价。

4.5 注意力机制:让张量互相打分

本节摘要:注意力机制把"哪个位置重要"变成可微分的打分问题:每个位置的查询向量与所有位置的键向量算相似度,softmax 归一成权重,再对值向量加权求和。它取消了 RNN 的逐步传递——任意两位置直接交互、全部时步一次并行算完,长序列的调度成本结构被彻底改写。本节拆解查询、键、值三元组与缩放点积的每一步,手写一个最小注意力层,并说明它如何嵌入编码器-解码器结构与自注意力。这是通向现代大模型结构族的门槛章节。

读完你应当能做到

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

  1. 说清查询、键、值三者的角色与缩放点积公式每一步的形状;
  2. 手写一个多头注意力的最小实现并用形状测试验证;
  3. 解释自注意力与交叉注意力在调度上的差别;
  4. 说明注意力相对 RNN 的并行优势与计算复杂度代价。

打分:从"顺序传递"到"全局加权"

RNN 的信息传递是接力:位置 t 想知道位置 1 的信息,必须等状态一站站传过来。注意力改成广播:每个位置都带着一份查询(我想找什么),每个位置也都亮着自己的键(我是什么),查询与所有键做点积得相似度,softmax 把相似度变成总和为一的权重,最后按权重把各位置的值向量加权汇合。位置 1 的信息直达位置 t,距离不再产生衰减——长距离依赖从"传不回来"变成"一次打分"。

缩放点积注意力的公式四步:相似度等于查询乘键的转置;除以键维度的平方根(防止点积值过大把 softmax 推进饱和区、梯度消失);softmax 归一;权重乘值求和。形状账要亲手过一遍:

import tensorflow as tf # 形状演练:批次 2,8 个位置,每个位置 16 维表示 X = tf.random.normal([2, 8, 16]) # 输入序列表示 Wq = tf.keras.layers.Dense(8) # 投影出查询,8 维 Wk = tf.keras.layers.Dense(8) # 投影出键 Wv = tf.keras.layers.Dense(8) # 投影出值 Q, K, V = Wq(X), Wk(X), Wv(X) # 各 (2, 8, 8) scores = tf.matmul(Q, K, transpose_b=True) # (2, 8, 8):每行是本位置对全部位置的打分 scaled = scores / tf.sqrt(tf.cast(8, tf.float32)) # 缩放:除以键维度平方根 weights = tf.nn.softmax(scaled, axis=-1) # 每行归一,总和为 1 out = tf.matmul(weights, V) # (2, 8, 8):加权汇合后的新表示 print(weights[0, 0].numpy().round(2), weights[0, 0].numpy().sum().round(3)) # 输出示例:[0.09 0.13 0.11 0.12 0.14 0.10 0.16 0.15] 1.0 # 权重行总和恒为 1——每个位置把注意力分配给所有位置

注意一个调度的关键观察:matmul(Q, K, transpose_b=True) 是一次矩阵乘法,8 个位置的打分全部并行完成——对比 RNN 的 8 次串行单元调用,这就是"并行取代串行"的字面含义。

图 20 注意力打分的信号流

图 20 注意力打分的信号流

手写一个多头注意力层

单头注意力只有一种"看序列的角度"。多头把表示切成若干份,各自打分再拼接,让不同头关注不同类型的依赖(语法、指代、位置模式等)。最小实现放在自定义层里:

class MiniMultiHeadAttention(tf.keras.layers.Layer): """多头缩放点积注意力的最小实现""" def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads == 0 self.h = num_heads self.depth = d_model // num_heads # 每头的键维度 self.wq = tf.keras.layers.Dense(d_model) self.wk = tf.keras.layers.Dense(d_model) self.wv = tf.keras.layers.Dense(d_model) self.wo = tf.keras.layers.Dense(d_model) # 拼接后的输出投影 def _split_heads(self, x, batch): # (batch, seq, d_model) 切成 (batch, heads, seq, depth) x = tf.reshape(x, (batch, -1, self.h, self.depth)) return tf.transpose(x, perm=[0, 2, 1, 3]) def call(self, x): batch = tf.shape(x)[0] q = self._split_heads(self.wq(x), batch) k = self._split_heads(self.wk(x), batch) v = self._split_heads(self.wv(x), batch) scaled = tf.matmul(q, k, transpose_b=True) / \ tf.sqrt(tf.cast(self.depth, tf.float32)) weights = tf.nn.softmax(scaled, axis=-1) out = tf.matmul(weights, v) # (b, h, s, d) out = tf.transpose(out, perm=[0, 2, 1, 3]) out = tf.reshape(out, (batch, -1, self.h * self.depth)) return self.wo(out) attn = MiniMultiHeadAttention(d_model=32, num_heads=4) y = attn(tf.random.normal([2, 10, 32])) print(y.shape) # 输出:(2, 10, 32) # 输入输出同形——注意力层可以像 Dense 一样串联堆叠

实现里最值得记住的是形状技巧:切头用 reshape 加 transpose,先拆维度再换轴;每头的缩放分母是 depth(8)而不是 d_model(32)——用全维度缩放是常见的隐性错误,会让 softmax 进入饱和区。

两种接线:自注意力与交叉注意力

自注意力(上例):查询、键、值来自同一序列,"序列内部互相看"——编码器的标配。交叉注意力:查询来自序列 A(解码器当前状态),键与值来自序列 B(编码器输出),"A 去 B 里找相关信息"——翻译类任务的解码标配,也是多模态对齐的基本件。两者的数学完全一样,差别只在 Q 与 KV 的来源。

# 交叉注意力:换掉 Q 的来源即可 enc_out = tf.random.normal([2, 20, 32]) # 编码器输出:20 个位置 dec_state = tf.random.normal([2, 6, 32]) # 解码器状态:6 个位置 wq = tf.keras.layers.Dense(32) wk = tf.keras.layers.Dense(32) wv = tf.keras.layers.Dense(32) scores = tf.matmul(wq(dec_state), wk(enc_out), transpose_b=True) / 5.66 weights = tf.nn.softmax(scores, axis=-1) out = tf.matmul(weights, wv(enc_out)) print(out.shape) # 输出:(2, 6, 32) # 解码器 6 个位置各从编码器 20 个位置加权取材——对齐就是这张权重表

成本账:并行的收益与代价

注意力把 RNN 的 O(n) 串行步换成一次并行矩阵乘,训练吞吐随序列长度近乎线性扩展;代价是打分矩阵是序列长度的平方——n 等于 1000 时就是一百万对的打分表,显存与计算都按平方涨。所以工程上有分界:几百步以内的序列,注意力全面碾压;上万步的超长序列,各种稀疏化、分块近似(线性注意力、滑动窗口)应运而生,把平方压回近似线性。RNN 家族也没有死——流式、超长、资源受限的场景里,串行的低显存反而是优点。

⚠️ 常见坑:给注意力层的输入忘了加掩码。变长序列补零后,零位置不该被打分(softmax 会分走权重),要用掩码把补零位置的分值压成负无穷再 softmax。短序列任务感知不明显,长任务与生成任务的输出会微妙走样。

💡 关键直觉:注意力的本质是把"信息该从哪来"交给数据自己算——打分权重是学出来的路由表。理解了这一点,Transformer 的一堆组件(多头、掩码、位置编码)都只是给这张路由表加工程约束。

本节要点回顾

  • 三元组分工:查询找什么、键是什么、值带什么,打分即路由。
  • 四步公式:点积、缩放、归一、加权;缩放分母是每头维度。
  • 多头机制:reshape 加 transpose 切头,多头各看各的角度。
  • 两种接线:自注意力序列内互看,交叉注意力跨序列取材。
  • 成本结构:并行换掉串行,但打分表按位置数平方增长。

下节把网格先验换成拓扑先验:图神经网络。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U