KV 缓存每生成一个 token 都要追加一块东西,Flash 又是个写一次少一块寿命的介质。这两件事摆在一起,逼出一个具体问题:每次追加到底写了多少字节?这篇我们动手实现一个静态字典加稀疏索引的最小可运行版本,跑一遍你就能看见那 15 倍流量是怎么省出来的。
第一步,先把完整 KV 向量原样写进去,跑一遍。KV 缓存在这里就是一段顺序追加的存储:每来一个 token,把当前层的键向量塞进去。我们不管多层、多头那套结构,只模拟单头单层,d_model 取 128——8B 这档模型常见的单头维度。
def naive_append(store, vector):
# vector: list[float],长度 d_model
store.append(vector) # 原样存下整条向量
return len(vector) * 2 # float16 每个元素 2 字节
跑一遍:生成 64 个 token,每个 token 的向量用同一个占位值,统计写入量。
store = []
for _ in range(64):
vec = [0.1] * 128
bytes_written = naive_append(store, vec)
if _ == 0:
print(bytes_written)
输出是 256。每追加一个 token,这层缓存要往 Flash 上写 256 字节;64 个 token 就是 16KB。这个数单看不大,但真实推理里 32 层、32 头、batch 8 一乘,写量直接放大三个数量级。Flash 的写入耐久按 TBW 算,这种量级烧下去介质先坏掉,不是性能问题,是寿命问题。
往下挖一层。要少写,就得改“追加什么”。《LLM Inference in a Flash!》里的思路是把每个 KV 向量表示成几个静态字典向量的线性组合——字典提前训练好、固定住,追加的时候只写几个索引和系数。我们先照这个思路把最粗糙的压缩函数搭出来:从一个预置 codebook 里挑 N_COEFS 个原子,系数量化成 int8,返回索引和系数,不再返回原始向量。
import numpy as np
DICT_SIZE = 2048
N_COEFS = 6
codebook = np.random.randn(DICT_SIZE, 128).astype(np.float16)
def kv_compress(vector):
v = np.array(vector, dtype=np.float16)
scores = np.abs(codebook @ v.T) # 分数用点积绝对值,选最相关的原子
idxs = scores.argsort()[-N_COEFS:].tolist()
coefs = [(codebook[i] @ v.T).astype(np.int8) for i in idxs]
return idxs, coefs
这个选择故意粗糙——真字典不会用随机 codebook,选原子也不是简单取点积最大。但机制的形状就是这样:用索引替掉整段向量。论文没给完整实现,这部分骨架是我按摘要补的,别把它当论文原文。
第二步,换成字典查找,只写索引。写入函数改成接索引和系数,索引按 16 bit 存,系数 int8。
INDEX_BITS = 16
def compressed_append(store, idxs, coefs):
entry = (idxs, coefs)
store.append(entry)
return len(idxs) * INDEX_BITS // 8 + len(coefs) * 1
第三步,追一次执行,看写入量掉在哪。同样 64 个 token,同样的占位向量。
store = []
for _ in range(64):
vec = [0.1] * 128
idxs, coefs = kv_compress(vec)
bytes_written = compressed_append(store, idxs, coefs)
if _ == 0:
print(bytes_written)
输出 18。每 token 从 256 字节掉到 18 字节,大约 14 倍。论文报告的 15 倍,在这个量级的参数下对得上。参数换一换数字会动,但省字节的机制不变:写进去的东西从 128 个浮点变成 6 个索引加 6 个整数。
到这一步就清楚了。所谓“利用 Flash 的大容量”,不是把 KV 缓存原样搬到 Flash 上,而是让缓存本身变成小的追加项。Flash 怕的是高频率写大块数据,对偶尔追加几十字节的索引不算太要命;字典向量是静态的,提前写进去一次,之后不参与动态流量。真正的写入压力——动态 KV 缓存流——被这几个索引压掉了。
对照真实实现,论文里还多一条端到端整数量化,用来消除浮点计算。我们上面这个骨架只动了 KV 压缩这一条线,系数已经 int8 了,但模型权重和激活的量化是另一块地。两条线最终叠在一起,才让模型权重和 KV 缓存都能利用 Flash 的计算特性。我们一开始也以为它只是把 KV 缓存压一下,翻完摘要发现权重那边的量化是消除高精度浮点计算的前提——Flash 设备不支持高精度浮点,这个约束不解决,光压 KV 缓存跑不起来。
这里扯了一句,但顺带把另一条硬约束点到了。回到 KV 这条线:压缩之所以成立,是因为写进去的东西从 d_model 个浮点变成几个索引和系数。想再深一层,去翻论文里字典的训练方式和更新策略——静态字典在长上下文里会不会漂,这是我最想追的一处。骨架已经搭好,参数换一换就能测。
