按题面顺序实现缩放点积自注意力:先用投影矩阵得到 Q,K,V,再算 S=QK⊤/d,对 j>i 的分数置 0,按行把分数除以该行和得到权重(行和为 0 则该行全 0,不使用 exp),最后 Y=AV。每一步的矩阵元素与 d 都四舍五入到两位小数,因果掩码发生在分数舍入之后、归一化之前。
实现方法:矩阵乘法在点积完成后立刻舍入;用自定义的两位四舍五入(0.5 向远离 0 进位),不要依赖语言默认的银行家舍入。归一化时对行和再舍入一次,避免浮点误差把全 0 行判成非 0。
时间复杂度 O(n2d+nd2),空间复杂度 O(n2+nd+d2)。在 n≤80、d≤16 下可以直接模拟。
某自回归解码器对一段已位置编码的 token 序列计算缩放点积自注意力,并施加因果掩码,输出每个位置的融合向量。
记序列长度为 n,隐向量维数为 d。输入给出 token 矩阵 X(n×d)与投影矩阵 WQ、WK、WV(均为 d×d)。令 Q=XWQ,K=XWK,V=XWV, S=QK⊤/d。 对 S 施加因果掩码:下标从 1 开始,若 j>i,将 Si,j 置为 0。再按行归一化得到注意力 A:第 i 行元素之和记为 Zi;若 Zi=0,则该行全为 0,否则 Ai,j=Si,j/Zi。本题归一化不使用 exp。最后 Y=AV。
矩阵乘法的每个输出元素、d、分数 S、行和 Zi、权重 A 与结果 Y 的每个元素,均在该步计算完成后四舍五入到小数点后 2 位;因果掩码在分数舍入之后、归一化之前执行。
第一行两个正整数 n 和 d。
接下来 n 行,每行 d 个实数,为 X 的一行。
接下来 d 行,每行 d 个实数,为 WQ。
接下来 d 行,每行 d 个实数,为 WK。
接下来 d 行,每行 d 个实数,为 WV。
1≤n≤80
1≤d≤16
X、WQ、WK、WV 的元素绝对值不超过 5,小数位不超过 2 位。
输出 n 行,每行 d 个实数,表示 Y,每个数保留两位小数。
输入
2 2
1 0
0 1
1 1
0 1
1 0
1 1
1 1
0.5 1
输出
1.00 1.00
0.50 1.00
说明
2 舍入为 1.41。Q 为 [[1.00,1.00],[0.00,1.00]],K 为 [[1.00,0.00],[1.00,1.00]],V 为 [[1.00,1.00],[0.50,1.00]]。 S 为 [[0.71,1.42],[0.00,0.71]]。因果掩码将 S1,2 置 0,得 [[0.71,0.00],[0.00,0.71]]。 第 1 行和为 0.71,归一化得 [1.00,0.00];第 2 行同理得 [0.00,1.00]。再与 V 相乘得到输出。
输入
3 2
1 0
1 1
0 1
1 0
0.5 1
0 1
1 0
1 0.5
0 1
输出
0.00 0.00
1.00 1.21
0.83 1.09
说明
2 仍为 1.41。Q 为 [[1.00,0.00],[1.50,1.00],[0.50,1.00]],K 为 [[0.00,1.00],[1.00,1.00],[1.00,0.00]],V 为 [[1.00,0.50],[1.00,1.50],[0.00,1.00]]。 S 为 [[0.00,0.71,0.71],[0.71,1.77,1.06],[0.71,1.06,0.35]]。掩码后第 1 行变为全 0,行和为 0,该行注意力全 0,故 Y 第 1 行为 0.00 0.00。 第 2 行掩码后为 [0.71,1.77,0.00],行和 2.48,权重 [0.29,0.71,0.00],与 V 相乘得 1.00 1.21。 第 3 行权重 [0.33,0.50,0.17],得 0.83 1.09。
In following contests:
Scan the QR code below with WeChat to sign in
First-time scan will create your account automatically
请使用微信扫描下方二维码完成注册