实现 Grouped Query Attention(GQA)算法。
给定输入序列 X(L×d)、Query 头数 hq、KV 头数 hkv(hq 为 hkv 的整数倍),以及投影矩阵,按以下步骤计算输出:
第一行包含四个整数 L,d,hq,hkv(1≤L≤8,2≤d≤8,hkv≤hq≤d,hqmodhkv=0,dmodhq=0)。
接下来 L 行,每行 d 个浮点数,表示输入序列 X。
接下来 d 行,每行 d 个浮点数,表示 Wq。
接下来 d 行,每行 dkv 个浮点数,表示 Wk(dkv=d⋅hkv/hq)。
接下来 d 行,每行 dkv 个浮点数,表示 Wv。
接下来 d 行,每行 d 个浮点数,表示 Wo。
输出 L 行,每行 d 个浮点数,表示 GQA 输出,保留 4 位小数。
输入
2 4 4 2
1.0 0.0 1.0 0.0
0.0 1.0 0.0 1.0
1.0 0.0 0.0 0.0
0.0 1.0 0.0 0.0
0.0 0.0 1.0 0.0
0.0 0.0 0.0 1.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 0.0
0.0 0.0
1.0 0.0 0.0 0.0
0.0 1.0 0.0 0.0
0.0 0.0 1.0 0.0
0.0 0.0 0.0 1.0
输出
0.7311 0.5000 0.7311 0.5000
0.5000 0.7311 0.5000 0.7311
说明
hq=4,hkv=2,分组比例 g=2。
Query 拆分为 4 头,每头维度 1。KV 拆分为 2 头,每头维度 1。Query 头 0,1 共享 KV 头 0;Query 头 2,3 共享 KV 头 1。
© CodeFun2000 · 使用条款
Scan the QR code below with WeChat to sign in
First-time scan will create your account automatically
请使用微信扫描下方二维码完成注册