MLX 高性能 KV Cache 编写指南:告别逐 token 拼接,用预分配 + 原地更新榨干自回归生成性能
2026/9/11 6:29:06 网站建设 项目流程

MLX 高性能 KV Cache 编写指南:告别逐 token 拼接,用预分配 + 原地更新榨干自回归生成性能

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

导读

在 MLX 中做自回归生成(LLM 推理)时,每生成一个 token 就要往 Key/Value 缓存中追加一个位置,而“追加”这一动作的实现策略往往决定了 KV Cache 的整体性能:mx.concatenate式的朴素拼接会让步耗时随上下文长度持续上涨,而“预分配固定块 + 原地更新”的策略可以把步耗时压到近似常数。本文以 docs/src/usage/kv_cache.rst 为核心骨架,结合 MLX 源码(slice_update的实现、CUDA 后端 SDPA 融合路径、缓冲区缓存BufferCache)深入讲解两种写法的差异、背后的内存与分配原理,以及 chunk size 的选取规则,读完即可在自己的推理脚本中直接落地。

为什么 KV Cache 的写入方式如此关键

自回归解码是一个逐 token 迭代的过程:每一步根据上一步的输出预测下一个 token,并把新 token 的 Key、Value 向量追加进缓存。用伪代码表示就是:

for x in steps: cache = mx.concatenate([cache, x], axis=1) # 每步都"重建"缓存 mx.eval(cache)

问题在于:append 策略本身就主导了 KV Cache 的性能。随着上下文增长,缓存越来越大,每次拼接都要复制全部旧数据,这种 O(n²) 的复制开销会让推理延迟线性恶化。文档明确指出,正确的做法不是换一种拼接方式,而是从根本上避免每步拼接

反面教材:朴素mx.concatenate拼接

先看文档给出的反例:

# Avoid this cache = mx.zeros((1, 0, d)) for x in steps: cache = mx.concatenate([cache, x], axis=1) mx.eval(cache)

这段代码每一步都执行concatenate:MLX 会为结果分配一个全新的数组,把旧缓存与新 token 全部拷贝进去。这里有两个隐藏成本:

  1. 数据复制成本:每步复制整个已有缓存,追加n个位置总共要复制约量级的元素;
  2. 缓冲区复用被破坏:缓存的形状每步都在变大,每次分配的尺寸都不同,MLX 的缓冲池无法命中,导致每步都从驱动层申请新内存。

更隐蔽的是,这种分配工作发生在CPU 侧。在性能剖析中,它表现为 kernel 之间的 GPU 空闲间隙,而不是某个慢 kernel——这让定位问题变得困难,甚至会被误判为模型本身是瓶颈(文档原文明确指出了这一现象)。

从源码看BufferCache的复用约束

MLX 对释放的设备缓冲区采用池化复用,其实现位于 mlx/backend/common/buffer_cache.h。核心逻辑在reuse_from_cache(第 30-46 行):它用lower_bound(size)在缓冲池中寻找最接近请求尺寸的缓冲区,且只有当候选块尺寸满足it->first < min(2 * size, size + 2 * page_size_)时才允许复用。

这意味着:

  • MLX不会把一个大缓冲区拆成小请求使用,也不会把多个小缓冲区合并成大请求;
  • 当一个逐 token 增长的缓存每步都有新尺寸时,每次分配几乎必然从驱动(driver)申请,而池中释放的旧缓冲区因为尺寸不再匹配而一直闲置;
  • 累积下来就是“每步一次新分配 + 旧块持续滞留”的恶性循环。

这也解释了为什么文档称“拼接会同时造成拷贝与阻碍缓冲区复用”两重代价。

推荐做法:预分配固定块 +slice_update原地更新

正确的姿势是预先按固定 chunk 分配缓存容量,之后只做原地切片更新

chunk = 256 cache = mx.zeros((1, chunk, d)) offset = 0 for x in steps: if offset == cache.shape[1]: cache = mx.concatenate([cache, mx.zeros((1, chunk, d))], axis=1) cache = mx.slice_update(cache, x, mx.array(offset), (1,)) offset += 1 mx.eval(cache) keys = cache[:, :offset]

要点拆解:

  • 容量预分配:缓存按chunk的整数倍提前分配,形状在大部分时间内保持固定;
  • 增量扩容:仅当offset == cache.shape[1](容量耗尽)时才一次性追加一个新的 chunk 块,而不是每步追加一个 token——扩容频率从每步一次降到每 chunk 一次;
  • 原地更新mx.slice_update(cache, x, mx.array(offset), (1,))x写入cache的第offset个位置,不复制旧缓存;
  • 最终切片:解码结束后用keys = cache[:, :offset]取出实际已用的有效部分。

这样既避免了反复拷贝,也避免了每步变化的分配尺寸破坏缓冲池复用:缓存形状固定后,分配只发生一次,后续slice_update全部命中同一块缓冲区。

两种写法在 M4 Max 上的实测对比

文档在 M4 Max 上对 20 个bfloat16、形状为[1, 4, N, 512]的缓存做了逐 token 追加测量(20 个缓存代表 20 层的典型 LLM 配置,N为上下文长度),结果如下:

ContextConcatenatePreallocate + update
5120.90 ms/step0.24 ms/step
10241.11 ms/step0.21 ms/step
40963.73 ms/step0.22 ms/step

两个清晰的结论:

  1. 预分配 + 原地更新的步耗时随上下文增长几乎不变(0.21~0.24 ms);
  2. 拼接写法的步耗时随上下文快速恶化(512→4096 时从 0.90 ms 涨到 3.73 ms,涨幅超过 4 倍),上下文越长差距越悬殊。

更自然的写法:索引赋值

如果觉得slice_update的可读性一般,文档还提供了等价的索引赋值语法:

cache[:, offset : offset + 1, :] = x

MLX 的索引赋值在底层同样落到slice_update系列算子。从 python/src/indexing.cpp 可以看到,a[idx] = v这类赋值会被改写为slice_update(第 968-971 行),而+=*=等复合赋值则分别映射到slice_update_addslice_update_prod等变体(第 996-998、1016 行、1034 行等),max/min也有对应变体(第 1052、1070 行)。因此两种写法性能等价,选自己喜欢的一种即可。

slice_update的 Python 签名在 python/src/ops.cpp 中定义:

def slice_update(a: array, update: array, start_indices: array, axes: Sequence[int], *, stream: StreamOrDevice = None) -> array

返回与输入a同形状、同类型的新数组(在 MLX 惰性求值下,它只是写入计算图的一个节点,实际内存写入发生在mx.eval时)。

为什么拼接慢:深入剖析两重代价

代价一:O(n²) 的重复拷贝

concatenate每次都会创建一个新数组并复制旧缓存。追加n个 token 时,第 1 步复制 1 个位置,第 2 步复制 2 个……总共复制的元素量级为1 + 2 + ... + n ≈ n²/2,即O(n²) 的拷贝总量。而slice_update只写入新到达的那一个 token,总拷贝量是 O(n)。

代价二:破坏缓冲区复用(再谈BufferCache

结合 mlx/backend/common/buffer_cache.h 的reuse_from_cache逻辑可以更精确地理解:缓冲池按“尺寸最接近”的原则复用,且拒绝拆分/合并(不满足it->first < min(2*size, size + 2*page_size_)即视为不可复用)。逐 token 增长的缓存每一步都是新尺寸,池里那些略小或略大的旧块统统不可用,于是:

  • 每步都向驱动申请新内存(CPU 侧开销);
  • 旧块滞留池中,显存/内存占用虚高。

预分配则同时消除两重代价:缓存形状固定后既没有反复拷贝,也没有每步一次的新分配。

一个重要提示:瓶颈藏在 GPU 空闲时间

因为分配与簿记工作在 CPU 侧完成,profiling 中看不到慢 kernel,只看到 kernel 之间的 GPU 空闲间隙拉长。若用mx.metal的捕获工具或 CUDA profiler 观察,遇到“kernel 间隔异常增长 + 模型 kernel 本身并不慢”的情况,优先检查 KV Cache 是否还在用拼接写法。

如何选择 Chunk Size:256 的倍数不是玄学

文档给出了硬性建议:chunk size 取 256 的倍数。理由有两条,且两者都直接对应实现层面的约束:

  1. 摊薄扩容开销:扩容频率与 chunk 成反比,chunk 越大,concatenate触发的次数越少;
  2. 启用 CUDA 后端的融合 cuDNN 注意力 kernel:这是关键约束。

源码证据:CUDA SDPA 对 256 的硬依赖

在 mlx/backend/cuda/scaled_dot_product_attention.cpp 中:

  • 第 85 行定义了constexpr int kv_cache_step = 256; // number is from mlx-lm,即融合路径的 KV 步进单位就是 256;
  • 第 86-88 行:若k.shape(2) < kv_cache_step(KV 序列长度不足 256),直接返回false,走慢路径;
  • 第 96-106 行的is_slice检查:通过 strides 反推 pre-sliced 的序列长度T_kv = kv.strides(1) / kv.strides(2),并要求T_kv % kv_cache_step == 0,同时验证kv.buffer_size()能容纳整个连续缓存(即 KV 必须是连续缓存的一个切片);
  • 第 110-122 行的unslice_kv则通过共享 buffer、反向构造 strides 来还原完整的 KV 视图,供融合 kernel 使用。

也就是说,在 CUDA 后端,单 token 注意力要走上融合 cuDNN 路径,必须同时满足:

  • KV 数组是连续缓存(contiguous cache)的切片
  • 缓存容量是256 的倍数
  • 已占用位置至少达到 256

不满足这些条件时,代码会静默回退到较慢的路径(文档明确提示“Other chunk sizes silently use a slower path”)。此外第 180 行还揭示了环境变量MLX_CUDA_SDPA_CACHE_SIZE(默认容量 256)用于控制 cuDNN SDPA kernel 缓存条数,与 KV chunk 无直接关系但同属 SDPA 调优范畴,这里不展开。

因此在实际代码里,chunk 取 256 是“性能下限的保证线”,取 512、1024 等更大的倍数同样合法,并能在扩容时获得更低的扩容频率。

完整可运行示例:把策略拼装成推理循环

把上面的要素组合起来,一个可直接套用的 KV Cache 更新循环如下(B为 batch,H为 head 数,d为 head 维度,L为层数):

import mlx.core as mx def make_cache(B, H, d, chunk=256, layers=1): # 每层独立维护 offset,容量按 256 的倍数预分配 return { "capacity": chunk, "offset": [0] * layers, "keys": [mx.zeros((B, H, chunk, d)) for _ in range(layers)], "values": [mx.zeros((B, H, chunk, d)) for _ in range(layers)], } def update_cache(cache, layer, k, v): chunk = cache["capacity"] off = cache["offset"][layer] if off == cache["keys"][layer].shape[2]: # 扩容:一次性追加一个 chunk,而不是一个 token cache["keys"][layer] = mx.concatenate( [cache["keys"][layer], mx.zeros_like(cache["keys"][layer])], axis=2) cache["values"][layer] = mx.concatenate( [cache["values"][layer], mx.zeros_like(cache["values"][layer])], axis=2) # 原地写入第 off 个位置(两种写法等价,任选其一) cache["keys"][layer][:, :, off : off + 1, :] = k cache["values"][layer] = mx.slice_update( cache["values"][layer], v, mx.array(off), (2,)) cache["offset"][layer] = off + 1 mx.eval([cache["keys"][layer], cache["values"][layer]]) # 解码结束后取出有效部分 keys = cache["keys"][0][:, :, : cache["offset"][0], :]

关键提醒:

  • 这里axis=2对应文档示例中[1, chunk, d]axis=1,只是形状约定不同,机制完全一致;
  • 每步都要mx.eval以触发实际计算(MLX 惰性求值下不 eval 不执行);
  • 初始容量、chunk 大小、[B, H, 256 倍数, d]的布局(d在最后一维、stride 为 1)都是为了满足 CUDA 融合 SDPA 的连续性要求——这也与prepare_sdpa_input中“last dim's stride 为 1、指针 16 字节对齐”的前置条件(scaled_dot_product_attention.cpp)相吻合。

总结与最佳实践清单

维度朴素拼接预分配 + 原地更新
每步数据拷贝O(n²) 总量仅写入新 token
每步分配每步新尺寸,缓冲池不可复用形状固定,一次分配
CPU 侧开销高(表现为 GPU 空闲间隙)
步耗时随上下文线性恶化(512→4096 时 0.90→3.73 ms)近似常数(约 0.21~0.24 ms)
CUDA 融合 SDPA不满足切片/容量条件,走慢路径chunk 为 256 倍数时可启用

最终落地建议:

  1. 永远不要在循环里用mx.concatenate逐 token 追加 KV Cache;
  2. 优先cache[:, offset:offset+1, :] = xmx.slice_update做原地写入,容量耗尽时按 chunk 一次性扩容;
  3. chunk 取 256 的倍数(256/512/1024…),以保证 CUDA 后端能进入融合 cuDNN 注意力路径,避免静默回退到慢 kernel;
  4. 排查推理变慢时,先看 profiling 里 kernel 之间的 GPU 空闲时间是否异常增长——那往往是 KV Cache 分配开销的信号,而不是模型 kernel 本身的问题。

以上全部内容以 docs/src/usage/kv_cache.rst 为主干,并交叉验证了 mlx/backend/common/buffer_cache.h、mlx/backend/cuda/scaled_dot_product_attention.cpp、python/src/indexing.cpp 与 python/src/ops.cpp 等源码,相关代码可直接在 examples/python 的推理示例基础上改造使用。

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询