基于PyTorch矩阵运算实现自定义注意力公式的技术求助
看起来你正在实现一个带额外位置依赖项的多头注意力机制,你的思路已经对了大半,只是最后计算z_i的时候的维度映射没处理对,导致输出维度不符合预期。我来帮你梳理下问题并修正代码~
原始公式与张量尺寸
你要实现的注意力公式如下:
$$
\begin{align*}
e_{ij} &= \frac{X_i W^Q (X_j W^K + A^K_{ij}) }{\sqrt{D_z}} \
\alpha_{ij} &= softmax(e_{ij}) \
z_{i} &= \sum_j \alpha_{ij} (X_j W^V + A^V_{ij})
\end{align*}
$$
各张量的尺寸定义:
X: [B, S, H, D] each W: [H, D, D] each A: [S, S, H, D]
你的现有代码与问题点
你已经完成了XW_Q、XW_K的核心投影计算,以及e_ij的初步推导,但最后计算z_i时,einsum的维度映射有误:既没有在计算e_ij时完成特征维度的内积求和,也没有在最终求和时正确压缩j维度,导致输出变成了[B, S, S, H, D],而我们需要的是对j维度求和后的[B, S, H, D]。
修正后的完整代码与逐步解释
import torch import torch.nn.functional as F # 初始化示例张量(补全X的定义) B, S, H, D = 2, 4, 8, 16 # 可根据需求调整尺寸 X = torch.randn(B, S, H, D) # 初始化注意力权重矩阵与额外偏置项 W_Q = torch.randn(H, D, D) W_K = torch.randn(H, D, D) W_V = torch.randn(H, D, D) a_K = torch.randn(S, S, H, D) a_V = torch.randn(S, S, H, D) d_z = D # 按你的假设d_z等于D # 1. 计算Q、K、V投影:X与W的批量矩阵乘法 XW_Q = torch.einsum('bshd,hde->bshe', X, W_Q) # 输出:[B, S, H, D] XW_K = torch.einsum('bshd,hde->bshe', X, W_K) # 输出:[B, S, H, D] XW_V = torch.einsum('bshd,hde->bshe', X, W_V) # 输出:[B, S, H, D] # 2. 计算注意力分数e_ij:完成特征维度的内积求和 # XW_Q.unsqueeze(2) → [B, S, 1, H, D],XW_K.unsqueeze(1)+a_K → [B, 1, S, H, D] # 对D维度内积求和后,得到每个(i,j,h)对应的注意力分数标量 e_ij_numerator = torch.einsum('bshd,bshjd->bshj', XW_Q, XW_K.unsqueeze(1) + a_K) e_ij = e_ij_numerator / torch.sqrt(torch.tensor(d_z, dtype=torch.float32)) # 输出:[B, S, S, H] # 3. 计算注意力权重alpha_ij:对j维度做softmax # 每个i对应的所有j的权重和为1,所以dim=2(对应j所在的维度) alpha = F.softmax(e_ij, dim=2) # 输出:[B, S, S, H] # 4. 计算最终输出z_i:对j维度加权求和 # alpha与(XW_V.unsqueeze(1)+a_V)相乘后,对j维度求和,得到每个i的最终特征 z_i = torch.einsum('bshj,bshjd->bshd', alpha, XW_V.unsqueeze(1) + a_V) # 输出:[B, S, H, D] # 验证输出维度 print("z_i的尺寸:", z_i.shape) # 预期输出:torch.Size([2, 4, 8, 16])(对应你设置的B,S,H,D)
关键修正说明
e_ij的维度修正:
原始代码中没有对特征维度D做内积求和,导致e_ij保留了D维度,这不符合注意力分数的定义——注意力分数应该是每个头h下i与j的相似度标量,因此必须通过einsum的bshd,bshjd->bshj完成D维度的求和。softmax的维度选择:
修正后的e_ij维度为[B, S, S, H],需要对每个i(第1个S维度)对应的j(第2个S维度)做softmax,因此指定dim=2,确保每个i的注意力权重之和为1。z_i的求和逻辑修正:
使用einsum的bshj,bshjd->bshd,明确对j维度(j标记)求和,将每个i对应的所有j的加权特征累加,最终得到符合预期的[B, S, H, D]尺寸。
备注:内容来源于stack exchange,提问作者GeraniumCat

