ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

如何为 MLX 自回归生成编写快速 KV cache:预分配 chunk 与就地更新

如何为 MLX 自回归生成编写快速 KV cache:预分配 chunk 与就地更新 如何为 MLX 自回归生成编写快速 KV cache预分配 chunk 与就地更新【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx在 MLX 中做自回归生成时每生成一个 token 都要向 key 和 value 数组追加一个位置。这个追加策略会直接决定 KV cache 的性能用mx.concatenate逐次拼接时单步耗时随上下文长度线性变差改用固定大小 chunk 预分配加就地更新后单步耗时在上下文增长时基本保持恒定。本文给出文档中的实现代码、chunk 大小选择依据和实测数据用于替换生成循环中的 KV cache 写法。为什么逐次 concatenate 会变慢最直观的写法是每个 step 把新位置拼到 cache 末尾# 文档中标注为应避免的写法 cache mx.zeros((1, 0, d)) for x in steps: cache mx.concatenate([cache, x], axis1) mx.eval(cache)文档指出concatenate有两笔开销复制数据。concatenate每步都创建一个新数组并拷贝整个已有 cache。追加n个位置总共要复制n^2量级的元素。阻止 buffer 复用。MLX 会池化已释放的设备 buffer但只把 buffer 复用于大小相近的请求既不会把大块 buffer 拆小给小请求用也不会合并小块满足大请求。不断增长、每步尺寸都不同的 cache 使每次分配基本都直接来自 driver已释放的 buffer 反而闲置。这笔分配开销发生在 CPU 侧。文档提醒在 profile 里它表现为 kernel 之间的 GPU 空闲时间而不是某个慢 kernel这容易让人误以为模型本身才是瓶颈。主路径预分配 chunk 并用 slice_update 就地更新文档推荐的替代实现如下d是 cache 的最后一维大小x是每一步新生成的一个位置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))], axis1) cache mx.slice_update(cache, x, mx.array(offset), (1,)) offset 1 mx.eval(cache) keys cache[:, :offset]要点cache 以chunk个位置为单位增长。只有当写满当前 chunk 时才执行一次concatenate扩一个 chunk因此复制和分配都被摊薄了。每一步写入用mx.slice_update(cache, x, mx.array(offset), (1,))完成把x写到 cache 第 1 轴即offset所在轴索引为offset的切片位置offset记录已写入位置数。生成结束后用cache[:, :offset]取出真正有效的前缀作为 keys/values。文档还给出一个可读性更好的等价写法使用索引赋值cache[:, offset : offset 1, :] x它和slice_update二选一即可写入的仍是同一个预分配 cache。如何选择 chunk 大小文档给出的选择标准是一句话chunk 大小取 256 的倍数。它同时满足两个目的摊薄增长开销在 CUDA 上启用 fused cuDNN attention kernel。该 kernel 对单 token attention 有硬性条件key 和 value 数组必须是连续 cache 的切片cache 容量是 256 的倍数且已使用的上下文至少 256 个位置。不满足这些条件时文档明确说明其他 chunk 大小会静默地走更慢的路径不会报错。文档给出的实测数据以下数据来自文档是在 M4 Max 上对 20 个形状为[1, 4, N, 512]的bfloat16cache、每步追加一个位置的测量作为文档示例理解不是你必须复现出的固定数值上下文长度Concatenate预分配 就地更新5120.90 ms / step0.24 ms / step10241.11 ms / step0.21 ms / step40963.73 ms / step0.22 ms / step文档的结论预分配方案下随上下文长度增长单步时间基本保持恒定concatenation 则随长度增加越来越慢。这就是判断改造是否生效的依据——如果你的实现里 KV cache 的单步耗时随生成长度明显爬升对照文档描述优先怀疑每步发生了整块复制或新尺寸分配。参考本场景的完整来源文档KV cache 用法slice_update与concatenate的 API 说明见 Python ops 参考【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表