注意力与 Transformer:参考答案#
说明#
以下答案与正文10道题逐题对应。前四题给出主要推导与计算;第 5--7 题给出可以运行的程序和检查方法;第 8--10 题说明比较时需要保持哪些条件相同,以及实验结果能够说明什么、不能说明什么。
本题的查询、键和值矩阵统一写为 \(\boldsymbol{Q}=[\boldsymbol{q}_1\trans;\ldots;\boldsymbol{q}_{T_q}\trans]\)、\(\boldsymbol{K}=[\boldsymbol{k}_1\trans;\ldots;\boldsymbol{k}_{T_k}\trans]\) 和 \(\boldsymbol{V}=[\boldsymbol{v}_1\trans;\ldots;\boldsymbol{v}_{T_k}\trans]\)。因 \(d_k=2\),缩放得分为
\[\begin{split}\boldsymbol{S} =\frac{\boldsymbol{Q}\cdot\boldsymbol{K}\trans}{\sqrt2} =\begin{bmatrix}1/\sqrt2&0\\0&1/\sqrt2\end{bmatrix}.\end{split}\]令 \(p=\exp(1/\sqrt2)/[\exp(1/\sqrt2)+1]\approx0.6698\),则 \(1-p\approx0.3302\),所以
\[\begin{split}\boldsymbol{A} \approx\begin{bmatrix}0.6698&0.3302\\0.3302&0.6698\end{bmatrix}, \qquad \boldsymbol{A}\cdot\boldsymbol{V} \approx\begin{bmatrix}1.3395&1.3210\\0.6605&2.6790\end{bmatrix}.\end{split}\]Softmax 必须逐行计算;最终输出矩阵写为 \(\boldsymbol{C}=[\boldsymbol{c}_1\trans;\ldots;\boldsymbol{c}_{T_q}\trans]\)。若沿查询轴归一化,本例因矩阵对称可能不易暴露错误,因此测试还应使用非对称矩阵。
对查询 \(i\) 的有效键集合 \(\mathcal K_i\),权重为
\[\alpha_{ij} =\frac{\exp(s_{ij})}{\sum_{r\in\mathcal K_i}\exp(s_{ir})}, \qquad j\in\mathcal K_i.\]指数和分母均为正,因此 \(\alpha_{ij}>0\),且直接求和得到 \(\sum_{j\in\mathcal K_i}\alpha_{ij}=1\)。若各值表示均为同一个列向量,即 \(\boldsymbol{v}_j=\boldsymbol{v}\),则
\[\sum_{j\in\mathcal K_i}\alpha_{ij}\boldsymbol{v}_j =\left(\sum_{j\in\mathcal K_i}\alpha_{ij}\right)\boldsymbol{v} =\boldsymbol{v}.\]当 \(\mathcal K_i\) 为空时分母是空和 0,归一化没有定义;程序必须屏蔽该查询输出或提前报错,不能依赖
softmax(-inf,...)。列向量 \(\boldsymbol{q}\) 与 \(\boldsymbol{k}\) 的点积为 \(S=\boldsymbol{q}\trans\cdot\boldsymbol{k}=\sum_{r=1}^{d_k}q_rk_r\)。独立和零均值给出 \(\mathbb E(q_rk_r)=0\);又因 \(\mathbb E(q_r^2)=\mathbb E(k_r^2)=1\),有 \(\operatorname{Var}(q_rk_r)=1\)。各项独立时
\[\mathbb E(S)=0, \qquad \operatorname{Var}(S)=d_k, \qquad \operatorname{Var}\!\left(\frac{S}{\sqrt{d_k}}\right)=1.\]真实网络的分量可能相关、均值和方差可能偏离假设,查询与键也由训练后的投影得到。因此缩放用于控制典型量级,并不保证每一层每个头的得分方差精确等于 1。
每头维度为 \(d_k=d/H=32\)。采用头轴显式存储时,\(\boldsymbol{Q},\boldsymbol{K},\boldsymbol{V}\) 的维度均为 \(8\times8\times128\times32\),得分维度为 \(8\times8\times128\times128\),各头拼接后输出为 \(8\times128\times256\)。
得分元素数为 \(8(8)(128)(128)=1\,048\,576\),float32 单个得分张量约占 \(4\,194\,304\) 字节,即 4 MiB;后向传播还需保存其他中间量,真实峰值更高。标准注意力得分和加权求和的主要时间为 \(O(BT^2d)\),得分内存为 \(O(BHT^2)\);线性投影还包含 \(O(BTd^2)\) 成本。
一个可核验的 PyTorch 骨架如下。数组
q[b, h, i, :]、k[b, h, j, :]和v[b, h, j, :]分别存放列向量 \(\boldsymbol{q}_i\)、\(\boldsymbol{k}_j\) 和 \(\boldsymbol{v}_j\) 的转置:import math import torch def scaled_attention(q, k, v, key_valid=None, causal=False): if q.ndim != 4 or k.ndim != 4 or v.ndim != 4: raise ValueError("输入应为 (B, H, T, D)") if q.shape[:2] != k.shape[:2] or k.shape[:3] != v.shape[:3]: raise ValueError("批次、头或键轴不一致") if q.shape[-1] != k.shape[-1]: raise ValueError("查询和键维度不一致") scores = q @ k.transpose(-1, -2) / math.sqrt(q.shape[-1]) allowed = torch.ones_like(scores, dtype=torch.bool) if key_valid is not None: allowed &= key_valid[:, None, None, :].bool() if causal: tq, tk = q.shape[-2], k.shape[-2] allowed &= torch.arange(tk, device=q.device)[None, :] \ <= torch.arange(tq, device=q.device)[:, None] if torch.any(~allowed.any(dim=-1)): raise ValueError("存在没有有效键的查询") masked = scores.masked_fill(~allowed, float("-inf")) weights = torch.softmax(masked, dim=-1) weights = weights.masked_fill(~allowed, 0.0) return weights @ v, weights
key_valid的含义是键位置是否有效;若还允许填充查询,应另外返回查询掩码并在输出和损失中使用,而不能混淆两个轴。用第 1 题的小矩阵检查全部中间量;再构造非对称得分以暴露错误 Softmax 轴。每种掩码都检查:被屏蔽权重精确为 0、有效行和在容差内为 1、输出维度正确、所有数有限。与库函数对照时先统一布尔掩码的真假语义,因为不同接口可能用
True表示允许或禁止。使用
float64小输入比较输出与对 \(\boldsymbol{Q},\boldsymbol{K},\boldsymbol{V}\) 的梯度,例如rtol=1e-7, atol=1e-9;极大得分测试验证稳定性。融合库算子可能采用不同累积顺序或较低精度,误差阈值应结合数据类型设定。循环版对每个查询显式计算所有键得分、稳定 Softmax 和值加权和;批量版使用矩阵乘法。固定同一输入,用
assert_close比较得分、权重和输出,预热后在多个 \(T\) 上重复计时。批量版一般能利用优化矩阵运算,但它与循环版都可能显式使用 \(T^2\) 存储。对置换矩阵 \(\boldsymbol{P}\),无位置表示且同步置换输入时,应有 \(f(\boldsymbol{P}\cdot\bX)=\boldsymbol{P}\cdot f(\bX)\)。测试可随机打乱词元轴并同步处理掩码。若把固定位置编码重新加到打乱后的存储位置,输入语义发生变化,通常不再满足这一等变关系。
两种模型使用相同的数据划分、词表、随机种子、训练样本、更新次数和每次更新的样本数;调整隐藏维度,使参数量大致接近。对于两种模型,都在同一个验证集上尝试相同的一组学习率,并按照相同标准进行选择。报告测试准确率或 F1、训练时间、固定 512 个样本的预测时间、峰值内存、参数量和网络层运行次数,并按短、中、长序列分组报告结果。
Transformer 训练可并行处理位置,但标准注意力内存随长度二次增长;GRU 沿时间顺序计算但单步存储较小。具体速度取决于设备、融合算子和长度,不能仅凭复杂度符号断言哪个模型必然更快。
固定 \(d\) 后改变头数,每头维度随 \(H\) 改变;保证 \(d\) 能被 \(H\) 整除。所有实验共享划分、种子集合、层数、前馈维度、更新次数、批量和评价程序,学习率用同一验证网格选择。记录测试指标、训练时间、固定批量预测时间、峰值内存、参数量和注意力调用次数。
在标准投影方式下总参数量通常变化很小,但内核效率和每头表达维度会变化。若多头结果更好,只能说明该任务和设置下的分解更合适,不能据注意力图给出未经验证的因果解释。
全局与局部注意力必须使用同一划分、种子集合、模型宽度、层数、优化器、更新次数、总处理词元数和评价代码;局部窗口只改变允许连接集合。对固定批量和固定长度计时,记录测试指标、训练时间、预测时间、峰值内存、实际得分元素数和网络函数计算次数。
局部窗口把每层注意力连接从 \(T^2\) 降到约 \(Tw\),但跨越窗口的依赖需要多层传播,任务若依赖远距离位置可能损失质量。若使用的实现仍构造完整矩阵再遮蔽,理论稀疏性不会自动转化为时间或内存收益,必须同时报告实现细节。