实现 Cross Attention(交叉注意力)算法。
给定查询序列 Xq(Lq×d)、键值序列 Xkv(Lkv×d)和投影矩阵 Wq,Wk,Wv(各 d×d),按以下步骤计算输出:
第一行包含三个整数 Lq,Lkv,d(1≤Lq,Lkv≤10,1≤d≤8),分别表示查询序列长度、键值序列长度、特征维度。
接下来 Lq 行,每行 d 个浮点数,表示查询序列 Xq。
接下来 Lkv 行,每行 d 个浮点数,表示键值序列 Xkv。
接下来 d 行,每行 d 个浮点数,表示 Wq。
接下来 d 行,每行 d 个浮点数,表示 Wk。
接下来 d 行,每行 d 个浮点数,表示 Wv。
输出 Lq 行,每行 d 个浮点数,表示 Cross Attention 输出,保留 4 位小数。
输入
2 3 2
1.0 0.0
1.0 1.0
1.0 0.0
0.0 0.0
-1.0 1.0
1.0 0.0
0.0 1.0
1.0 0.0
0.0 1.0
1.0 0.0
0.0 1.0
输出
0.4259 0.7160
0.0000 0.8022
说明
查询序列长度 Lq=2,键值序列长度 Lkv=3。
Q 来自 Xq(2×2),K,V 来自 Xkv(3×2)。注意力矩阵维度为 (Lq,Lkv)=(2,3),每个查询位置关注所有键值位置。
输入
1 4 3
1.0 2.0 3.0
0.5 0.5 0.5
1.0 1.0 1.0
-1.0 0.0 1.0
0.0 -1.0 0.0
0.0 1.0 0.0
0.0 0.0 0.0
0.0 1.0 0.0
0.0 0.0 1.0
0.0 1.0 0.0
0.0 0.0 0.0
0.0 1.0 0.0
1.0 0.0 0.0
0.0 1.0 0.0
0.0 0.0 1.0
输出
0.7691 0.8387 0.9235
说明
单个查询位置关注 4 个键值位置,输出为 Value 的加权和。
© CodeFun2000 · 使用条款
Scan the QR code below with WeChat to sign in
First-time scan will create your account automatically
请使用微信扫描下方二维码完成注册