本题对应的算法是 基于注意力重要性的 KV 缓存稀疏化(Sparse KV / Token Eviction),实现上就是一次 全局贪心 Top-K:先锁定正在生成的位置 C,再按注意力分数对其余 KV 排序并取前若干个。不是聚类,也不是按层、按头均分预算。
在全局预算 M 下,为每一层每一个注意力头选出要保留的 KV 位置,规则已经写死,按 Top-K 实现即可。
某推理服务在长序列生成时,会为每一层、每一个注意力头保存历史 Token 的 KV。序列变长后,全部留下会超出显存。系统给出全局名额上限 M,要求在各层、各头上选出要留下的 Token 下标。不同层、不同头留下的个数可以不同。
层数记为 L,每层 H 个头,当前序列长度为 N。正在生成的位置为 C(0≤C<N)。每个头在每个位置上都有一个注意力分数。
规则如下:
第一行四个整数 L、H、N、M,依次表示层数、头数、序列长度和名额上限,满足 M≥L×H。
随后 L 行。第 i 行有 H×N 个浮点数,为空格分隔的该层全部头的注意力分数:每个头连续 N 个数,并按头下标从 0 到 H−1 依次拼接。
最后一行一个整数 C,表示正在生成的位置。
1≤L,H≤32
1≤N≤200
L×H≤M≤L×H×N
0≤C<N
注意力分数为浮点数,满足 0≤score≤1
输出 L 行。第 i 行对应第 i 层,含 H 个字符串,相邻字符串之间用一个空格分隔。
每个字符串是该层该头留下的 Token 下标:按下标升序排列,用逗号连接;只有一个下标时不写逗号。
输入
2 2 4 9
0.2 0.5 0.1 0.8 0.4 0.4 0.7 0.3
0.6 0.9 0.2 0.1 0.5 0.3 0.8 0.8
2
输出
1,2,3 2
0,1,2 2,3
说明
当前位置 C=2。每层每头必留位置 2,共 4 个,剩余名额 9−4=5。
去掉位置 2 后,其余分数按「分数从高到低,相同则 l、h、t 从小到大」排列:
0.9 (1,0,1),0.8 (0,0,3),0.8 (1,1,3),0.6 (1,0,0),0.5 (0,0,1),0.5 (1,1,0),0.4 (0,1,0),0.4 (0,1,1),0.3 (0,1,3),0.3 (1,1,1),0.2 (0,0,0),0.1 (1,0,3)。
括号内依次为层、头、Token 下标。取前 5 项:(1,0,1)、(0,0,3)、(1,1,3)、(1,0,0)、(0,0,1)。再加上必留的 2,各头集合为:
下标升序后即得输出。
输入
2 2 3 6
0.5 0.2 0.5 0.5 0.1 0.3
0.4 0.9 0.5 0.5 0.6 0.5
1
输出
0,1,2 1
1 1
说明
当前位置 C=1。必留 2×2=4 个位置 1,剩余名额 6−4=2。
去掉位置 1 后排序:0.5 (0,0,0),0.5 (0,0,2),0.5 (0,1,0),0.5 (1,0,2),0.5 (1,1,0),0.5 (1,1,2),0.4 (1,0,0),0.3 (0,1,2)。
分数同为 0.5 时按 l、h、t 升序,因此前 2 项是 (0,0,0) 与 (0,0,2),都落在第 0 层头 0。若按层从大到小取,会选到第 1 层,结果不同。
加上必留位置 1 后:第 0 层头 0 为 {0,1,2},其余三个头均只留 {1}。
Scan the QR code below with WeChat to sign in
First-time scan will create your account automatically
请使用微信扫描下方二维码完成注册