通过率 100% · 提交 1 · 通过 1
给定同一序列的 Q、K、V 三个矩阵,请计算单头因果自注意力输出。第 i 个 token 只能关注编号不大于 i 的 token。注意力分数为 dot(Q[i], K[j]) / sqrt(d),权重为这些分数的 Softmax,输出为对应 V 的加权和。
这类题属于算法机考高频题型中「华为 AI 岗 / Attention」方向的高频题型,通常考察对「华为 AI 岗 / Attention」的建模能力与边界条件处理。掌握本题的解题思路后,可举一反三应对同类真题方向,稳步提升机考通过率。
第一行输入 n d。接下来 n 行输入 Q 矩阵,再接下来 n 行输入 K 矩阵,最后 n 行输入 V 矩阵。每行 d 个实数。
输出 n 行,每行 d 个实数,表示注意力输出。每个数保留四位小数,绝对值小于 0.00005 时输出 0.0000。
示例 1
输入示例
2 1 0 0 0 0 2 4
输出示例
2.0000 3.0000
一维两 token 的平均效果
示例 2
输入示例
3 1 1 1 1 0 0 10 1 2 100
输出示例
1.0000 1.5000 99.9911
后面的 token 不得影响前面输出
时间限制 3000 ms · 内存限制 256 MB
本平台为独立第三方培训机构,与华为技术有限公司无任何关联;课程的服务内容与权益以购买协议为准,学习效果因个人情况而异。「华为 OD」「华为可信」等仅为对岗位与考试方向的客观描述,相关商标归各自权利人所有。
这些是真正决定能不能 AC、但通用题解里常被略过的点。
单头因果自注意力的逐步计算:打分、掩码、稳定 Softmax、加权和。四步全是循环和四则运算,考的是把「因果」和「数值稳定」两条规则写对。
对第 i 个 token(0 ≤ i < n):
1. 打分:只对 j ≤ i 的位置算分数 s_j = dot(Q[i], K[j]) / √d。因果掩码的意思就是生成第 i 个 token 时看不到未来——j > i 的位置直接不参与,等价于分数为 −∞。 2. 减 max:取 m = max(s_0..s_i),全部分数减去 m。这是稳定 Softmax 的标准动作:分数绝对值可以到很大,exp(几百) 直接溢出,减掉最大值后指数的最大参数是 0,永不溢出,而且数学上和不减完全等价(分子分母同乘一个常数)。 3. Softmax:w_j = exp(s_j − m) / Σ exp(s_t − m)。 4. 加权和:输出第 i 行 = Σ w_j · V[j],逐维累加。
掩码的实现建议直接收窄循环范围(j 只跑到 i),比先算全矩阵再放 −∞ 干净,也省一半计算。√d 别忘除——它让分数尺度不随维度膨胀,漏掉后数值全错。求和按 j 从小到大累加,与判题口径一致。
每个数保留四位小数,绝对值小于 0.00005 输出 0.0000。加权和可能出现 −0.00003 这类值,负零钳制不做必挂一档用例。
时间 O(n²·d),空间 O(n·d)。n、d 规模内纯循环轻松通过。
-0.0000 未钳制。样例 1:n=2、d=1,Q=[0,0]、K=[0,0]、V=[2,4]。
输出 2.0000 与 3.0000。第 1 行是「平均」不是「只看自己」——掩码放行的是过去与现在,j ≤ i 里那个等号别丢。
手撕注意力是大模型岗的招牌题,写之前先把四步口述一遍:打分、掩码、归一、加权——分数是 Q 与 K 的点积除以 √d,掩码保证自回归生成不偷看未来,softmax 前减 max 保证数值稳定,最后用权重混合 V。追问几乎必到「为什么除以 √d」:d 越大点积的方差越大,分数一拉开 softmax 就趋近 one-hot,梯度消失;除以 √d 把方差拉回常数量级。能主动指出「max 要在掩码内取」这种实现细节,比背出完整公式更像真手写过的人。
第一行输出永远只依赖第 0 个位置,先看它对不对——第一行就错,问题在读入或 √d;第一行对、后面错,问题在掩码或减 max 的范围。溢出报错直接指向没减 max;输出全体偏移一个比例,八成是 √d 忘除或除成了 d。
顺带一个值得知道的联系:逐行计算时,K 和 V 的每一行都被后面的行反复读取——真实推理框架把已算好的 K、V 缓存下来复用,就是 KV cache 的由来。手算过一遍这道题,缓存省掉的是哪部分计算量,一眼就能看清。
# 因果注意力:第 i 行只对 j<=i 打分 dot(Q[i],K[j])/√d,行内减 max 再 softmax,加权和 V;
# 掩码用收窄循环范围实现,max 也只在掩码内取;输出四位小数注意负零钳制。
import math
import sys
def solve() -> None:
data = list(map(float, sys.stdin.buffer.read().split()))
if not data:
return
at = 0
n = int(data[at]); at += 1
d = int(data[at]); at += 1
def read_matrix():
nonlocal at
matrix = []
for _ in range(n):
matrix.append(data[at:at + d])
at += d
return matrix
q = read_matrix()
k = read_matrix()
v = read_matrix()
scale = math.sqrt(d)
ans = []
for i in range(n):
scores = []
for j in range(i + 1):
scores.append(sum(q[i][p] * k[j][p] for p in range(d)) / scale)
mx = max(scores)
weights = [math.exp(score - mx) for score in scores]
total = sum(weights)
row = [0.0] * d
for j, weight in enumerate(weights):
weight /= total
for p in range(d):
row[p] += weight * v[j][p]
ans.append(" ".join(f"{0.0 if abs(value) < 0.00005 else value:.4f}" for value in row))
print("\n".join(ans))
if __name__ == "__main__":
solve()
登录后可查看你在本题的历史提交,以及每次的各用例通过情况。
© 2026 广州慕课网络科技有限公司 · 吴师兄学算法官网 版权所有