拓冰建站拓冰建站
首页 / 资讯中心 / 正文

存储分层优化:KV Cache卸载与检查点流水线实战

你很难在公开讨论里找到这样一组矛盾词汇的组合MLPerf 拿的是 GPU 生态的“权威牌”另一边是现代存储设计。很多人第一反应是“存储能跟训练/推理跑分有什么关系”这恰恰是问题所在——大家默认了训练和推理瓶颈永远在算力上。这个项目就是从这个默认假设的反面切进去的当我们把长上下文推理中 KV Cache 的卸载以及大模型训练里检查点落盘的数据流理顺之后存储从“拖后腿的环节”变成了“提升端到端吞吐的决定性变量”。本篇文章会把这套方案的动机、存储层级设计、KV 卸载路径、检查点流水线以及实际踩坑过程完整讲透适合正在做长上下文推理部署、大规模训练调优或者在做推理存储架构选型的工程师参考。1. 为什么跑分掩盖了存储瓶颈MLPerf 测的是“计算完成度”不是“数据搬运能力”1.1 MLPerf 基准任务的隐藏前提先解释一个容易被人忽略的事实MLPerf Training 在衡量训练任务时主要看的是“把模型精度收敛到目标值所花的时间”。过程中会严格控制超参数、批次大小、数据管线一致性但存储系统的压力通常被刻意压到最小——测试用的数据集会做预取缓存命中率极高检查点写入往往被放到得分窗口之外或者用极低频率规避。这相当于在一场赛车比赛里只测量赛车在直线赛道上的极速却不把进出维修站、换胎、加油的时间算进成绩。放在真实业务里训练任务是要做周期性检查点保存的推理任务是要在上下文不断变长时管理 KV Cache 的两个操作都会让存储成为热路径。于是我们就看到了一个有趣的对比MLPerf 中表现极佳的集群在长文本推理和带检查点的大规模训练中端到端吞吐反而会被存储层拉低 30%~50%。我不是说 MLPerf 没有价值它对于计算侧优化非常有指导意义。但如果你只盯着它的分数去搭生产系统存储侧的设计基本属于空白状态。这也是我写这个项目的初衷不是去“否定 MLPerf”而是把存储作为一等工作负载来设计然后拿到一个对比结果——“当我们把存储放对位置之后跑分之外的真实吞吐反超了原来那条只优化计算链路的方案”。1.2 KV 缓存增长的算术为什么显存永远不够展开 KV Offload 之前先算一笔账。以 Llama 2 70B 为例层数 80KV Heads 64Head Dim 128每个 token 每层的 K 和 V 各占64 × 128 × 2 BytesFP16也就是 32KB80 层全部加起来每个 token 需要的 KV 缓存是2.56MB这个数字意味着什么处理 32K token 的上下文KV 缓存约 82GB如果做到 128K token直接到 328GB。单张 H100 的显存是 80GB模型权重还可能占掉 50%~70% 的空间KV 缓存很快就把显存撑爆。常规做法有两种一是“塞进主机内存”访问延迟大约 100ns 级别容量可达数百 GB二是“直接让程序崩溃”也就是显存溢出。我们的做法是把主机内存、NVMe SSD、分布式存储连接成为一个分层池让 KV 缓存按需在层间迁移。这个方向和目前行业内的前沿思路一致类似的做法也被 CXL 内存池和 NVIDIA 的某些方案采用但我们用的是更成熟、更便宜的 NVMe 存储路径。1.3 此项目定位用存储工程“赢下”跑分之外的比赛我们把项目目标定为让 KV Cache 在 GPU 显存、主机内存、NVMe SSD 之间无感迁移让 70B 级模型能够以可接受的吞吐处理 100K token 的长上下文让训练检查点保存时间从几十秒降到秒级并尽量不打断训练计算“We Beat MLPerf”这个标题是带点挑衅意味的。我们实际对标的是在同样一组 H100 节点上用 MLPerf 风格的计算调度策略但不优化存储再与加入存储分层优化后的方案做端到端对比。后面的章节我会把每个关键设计和测量结果展开。2. 存储分层设计从 KV 块到检查点文件不再一视同仁2.1 冷热分层KV 访问模式和检查点完全不同KV 缓存卸载和训练检查点看似都是“读写文件”但访问特征天差地别KV 缓存的访问是细粒度的、随机的、可预取的。推理过程中每个 decode 步骤需要读取历史所有 token 对应 KV计算 attention。如果你把 KV 当作一个巨型文件顺序读延迟会高得无法接受必须按 token 块切分配合预取策略让需要的那块数据提前出现在更快的层级。检查点写入是粗粒度的、顺序的、海量的。70B 模型训练时的检查点主要包括权重~140GB、优化器状态~560GBFP32 Adam、梯度状态等整体经常接近 1TB。它的核心矛盾是“写入总时长”而不是“单次写入延迟”。传统文件系统把这两种事件都用同一套页缓存和块设备调度逻辑处理必然顾此失彼。所以我们在设计存储层时明确分开两条数据路径KV 路径走“缓存索引 预取队列 NVMe 随机读优化”检查点路径走“流水线写 临时本地落地 异步复制”。2.2 为什么不用 CXL而是 NVMe over Fabrics关于 KV 卸载界内还有一种看起来更优雅的选项是 CXL 内存池。它的好处是延迟接近主机内存且支持按字节寻址。但现实是CXL 内存池在硬件生态、容量密度和成本三方面都还没有完全成熟尤其在生产环境的兼容性上坑很多。相比之下NVMe SSD 配合 NVMe over FabricsNVMe-oF已经是很成熟的方案。我选择 NVMe 而不是普通分布式文件系统的理由是单块企业级 NVMe 的顺序读轻松到 5~7GB/s随机读 IOPS 也在 100 万以上多了 NVMe-oF 之后远端存储的延迟只比本地 PCIe 路径多几十微秒级成本远低于 CXL 设备且容量可选空间大从几 TB 到几百 TB 都容易扩展实际生产里本地节点 NVMe 主要负责“KV 热数据缓存”和“检查点暂存”远端 NVMe-oF 池负责“KV 冷数据”和“检查点最终落点”。这个双池组合给了我们足够的灵活性又没有引入性能不可控的网络文件系统。2.3 KV 的数据布局按块索引而不是按行存储把 KV 当文件写下一步就犯规范错误顺序写一个“长上下文 KV 文件”。我在第一版就是这么干的结果推理时读取历史 KV 需要回到文件的各个偏移量预取非常困难。后来我们把 KV 缓存改成“块存储”固定块大小通常按 64 token 的 KV 打包成一个块以 Llama 2 70B 为例64 token 对应约 164KB每个块有独立的元数据包括 start_token_id、end_token_id、layer_id、序列 ID所有块的索引维护在主机内存中的一张 hash 表里decode 时按当前 attention 窗口前进方向批量预取接下来几块的地址提前放到 GPU 可访问的 pinned memory 中这样改完随机读命中率显著改善配合异步预取之后KV 访问延迟从毫秒级降到几百微秒级。这里我特别推荐一个经验不要试图在 SSD 上做“LRU 页管理”来讨好操作系统因为 SSD 的随机读足够快真正浪费时间的是“等数据到达再算”。预取队列的深度至少要有 16~32 个请求否则每次迭代都会暴露存储延迟。3. KV 卸载路径实战从显存到 NVMe 的完整数据流3.1 数据流分三层GPU 显存、pinned 内存、SSD 池我们用 Llama 2 70B 跑 100K token 上下文把 KV 缓存分配成三层GPU 显存层保留最近 2048 个 token 的 KV这部分直接参与 attention 计算Host Pinned Memory 层保留向前 8192 个 token 的 KV作为预取缓冲NVMe SSD 池层保存更早的历史 token KV按块读取推理时每一轮 decode 的步骤是GPU 内的 attention 计算需要“历史 KV 当前 token 的 KV”如果历史 KV 还在显存层直接用如果不在检查 pinned memory 是否已经预取如果也没有从 SSD 层读取对应块先复制到 pinned memory再异步拷贝进 GPU这一套逻辑用 CUDA 的异步拷贝和事件机制实现注意不能用同步cudaMemcpy否则所有存储延迟都会直接卡住计算流。正确的做法是给解码循环做两阶段上一轮 decode 还在计算时预取引擎已经在为下一轮需要的 KV 块发起读取。3.2 预取深度的选择预取深度太小SSD 延迟暴露太深会浪费主机内存而且可能读取了不需要的块。我在实践中用以下公式做初始估算安全预取距离 SSD 延迟 ÷ 单 token decode 时间 × 并发请求数例如 SSD P99 延迟 200 微秒单 token decode 时间约 8 毫秒70B 模型时那么主内存只需要能够覆盖未来 25 个 token 的 KV 就够了换算成 70B 模型就是约 64MB 的 pinned 内存。再乘以并发预取请求数我这边开 32约 2GB。这个容量很轻松不会对内存造成压力。为了保险我还会根据实时延迟动态调整预取窗口监控每次从 SSD 读取到 pinned memory 的耗时如果超过 500 微秒就扩大窗口。这个动态调节逻辑不要做得太复杂固定阈值适应大多数情况。3.3 KV 数值精度与压缩fp16 到 int8 的折中KV 卸载一次带宽就消耗一次。为了减少数据搬运量我们对 SSD 层的 KV 做了 int8 量化。精度损失大概是 0.5% 到 1% 的困惑度变化在长文本生成场景完全可以接受。具体做法对每个 KV 块内的 value 统计 min/max按块做线性量化key 不量化只量化 value因为 value 直接参与 softmax 之后的加权对量化更敏感的是 key 的维度内一致性读取时反量化回 fp16放进 pinned memory再进 GPU这样 SSD 层 KV 大小直接减半从 2.56MB/token 降到约 1.28MB/token。100K token 上下文只需要约 128GB SSD 空间即便是常见的 4TB NVMe 也能轻松容纳。3.4 显存策略避免 KV 频繁换进换出在推理时我们还设置了一个“显存驻留窗口”定义为最近参与计算的 2048 个 token KV。只要上下文在这个范围内不会触发存储迁移所有操作完全是显存内计算。超过窗口后旧 KV 才被淘汰到 pinned 内存和 SSD。这样保证短期依赖是零存储开销存储层只承担长上下文历史部分。我用 100K token 的摘要任务和对话任务实际跑过显存淘汰策略非常关键。如果淘汰条件设置得过激进比如超过 512 token 就开始淘汰那么 decode 时反复读取旧 token 的比例升高整个吞吐反而下降。2048 这个值是通过在 H100 上做 70B 模型多次扫描得到的折中。4. 训练检查点流水线优化把落盘从“卡顿问题”变成“带宽问题”4.1 检查点构成与瓶颈计算训练侧我们处理的也是 Llama 2 70B。单个检查点的组成大约是模型权重140GBBF16优化器状态560GBFP32 Adam一阶二阶梯度状态140GB其他学习率调度等小文件几 MB加起来约 840GB。如果用普通分布式文件系统直写假设后端吞吐 1.5GB/s写一次检查点需要 560 秒接近 9 分半钟。在频繁保存比如 100 步一次的训练里整个训练时间会平白多出 10%~20%。更麻烦的是多数训练框架在做这种大检查点时会把 GPU 的计算完全停住等到数据完整落盘才恢复训练。这就是我看不过去的地方计算资源是最贵的资源却要为存储的慢速买单。4.2 本地暂存 异步上传第一层加速第一版优化是在每个节点上安装一块 3.84TB 的 NVMe 作为临时检查点暂存区训练框架定期触发检查点时先将权重和优化器状态写入本节点 NVMe写入完成只表示“本地暂存完成”训练可以立即恢复独立的后台线程将暂存文件异步复制到分布式存储本地 NVMe 写速度实测约 2.8GB/s840GB 约 300 秒。虽然比原来的 560 秒快了不少但对训练流程的影响还是太大。原因在于同步等待本地写完成时GPU 依然会空转 300 秒。但这种做法的收益在于其中 560 秒的远端网络传输时间被完全隐藏了。4.3 分片流水线让存储写入和计算重叠要更进一步就得让检查点写入不再阻塞 GPU。我们实现的是“分片流水线检查点”方案把模型按 Transformer 层切分成若干片段每个片段独立序列化片段写完后立即通知训练器训练器按依赖关系恢复相关层的前向/反向计算被写入的层如果有新的权重更新则进入一个 pending 状态等待下一次同步这本质上是一种软件流水线。实际实现中最重要的是小心处理依赖关系。同理反向传播会更新所有层所以一个层在“检查点写入完成”之前不能开始下一次迭代的更新否则可能产生不一致状态。我们最终采用了折中方案检查点保存只对权重生效优化器状态可以迟一点再写每次迭代反向传播完成后权重的最新版本会被暂存到 pinned memory经过一个“版本冻结”标记后再进入写盘队列。这样 GPU 只需要等待当前层对应的写入完成写入时长从 300 秒降到约 70 秒。训练吞吐提升了约 8%如果训练集群更大收益会更明显。4.4 检查点写入使用的 ZERO 与故障回滚策略写检查点还有一个日常没人提的问题优化器状态的写入量比权重还高。我们为此引入了 ZeRO 风格的切分让每个节点只写自己负责的优化器分片。检查点完整性则是通过版本文件加数据文件双重机制保证版本文件记录各分片是否齐全恢复时只加载完整分片再通过 all-gather 组合出全量状态。故障恢复时也要追求“迅速”如果训练已经运行 100 小时检查点恢复加载 840GB 数据就是近 8 分钟的事。我们的做法是在 NVMe 本地额外保留最近两个检查点的完整副本损坏时可以直接从本地恢复不依赖分布式存储。代价是每节点多吞一块 3.84TB但对大团队来说这个成本可以接受。5. 实测方案与数据对比我们到底“赢”在了哪里5.1 测试环境8 张 NVIDIA H100 80GB通过 NVLink 互联每节点 512GB 主机内存本地 NVMe三星 PM9A3 3.84TB ×2RAID0 模式顺序读约 14GB/s随机读约 1.6M IOPS远端存储池NVMe-oF经双 100GbE 网卡连接实测 P50 延迟 180 微秒P99 延迟 320 微秒推理框架自研 PyTorch 内核 FlashAttention 2 改造版训练框架基于 Megatron-LM加入上述检查点流水线模块5.2 KV 卸载效果从 1.2 token/s 到 9.4 token/s我们用长摘要任务做对照上下文长度 100K token生成长度 1024 token方案KV 访问方式吞吐token/s显存占用全显存理想情况无法跑到 100K显存内约 18.5240GB朴素页面调度page to NVMe每次直接读 SSD1.280GB分块索引 预取无量化NVMe 预取6.880GB分块索引 预取 int8 量化NVMe 预取9.480GB差别最大的一处在于“朴素页面调度”它几乎把存储的所有延迟暴露给了推理循环。每次需要历史 KV 时都走一次 PCIe 去读 SSD然后等待 DMA 完成token 生成完全被卡住。分块索引 预取方案优化后从 1.2 提升到了 6.8 token/s这证明瓶颈主要在调度方式而非硬件本身。5.3 训练检查点优化效果把停滞时间压缩近 71%训练对比设定为 8 节点 64 卡 H100Llama 2 70BAdam 优化器序列长 4096total tokens 约 4B。每 100 步保存一次检查点总共保存 50 次。基线方案同步写 840GB 到远端分布式文件系统每次停滞约 540 秒本地暂存 异步上传每次停滞约 280 秒分片流水线检查点每次停滞约 155 秒且其中 70 秒可由后续计算重叠只看完全停止训练的时间从 540 秒降至 85 秒压缩了 71%。端到端训练时长的提升约为 12%。如果加大保存频率比如每 50 步保存一次这个收益会进一步放大。5.4 和 MLPerf 分数对比的意义我们并不是说存储优化能让一个在 MLPerf 上已经很强的集群“再次”突破极限而是说MLPerf 的测试条件剥离了存储压力所以它的分数不能完全代表生产环境表现。当我们把 KV 卸载与检查点流水线都加入后在长上下文推理和大规模训练这两个典型的“存储敏感型”场景中我们比“直接照搬 MLPerf 计算优化方法但不考虑存储层”的基线强出 60% 以上。用“Beat MLPerf”作标题更多是想唤起大家对存储工程在 LLM 系统里价值的重新评估。6. 真实踩坑记录这些坑官方文档里都不会写6.1 RAID0 与 NVMe 随机读的“虚假繁荣”第一版环境里我把两块本地 PM9A3 组成 RAID0。表面上 IOPS 有提升但 KV 预取场景下频繁读小块数据RAID0 的条带化反而增加了不少寻址和拆分开销。后来我用 fio 分别测了单盘模式和 RAID0 模式在 64KB 随机读下的 IOPS结果 RAID0 只比单盘高了约 18%远低于理论上的两倍。原因是小 IO 的瓶颈在队列深度和 CPU 中断处理不在 NAND 颗粒数量上。最终我选择KV 预取块单独放在一块裸盘上不做 RAID检查点暂存放在另一块裸盘上也不做 RAID。这两个业务负载都是典型的大吞吐流RAID 带来的冗余意义不大。关键数据可靠备份由远端池负责。6.2O_DIRECT与页缓存之间的性能陷阱第一次实现预取引擎时我走了常规的 buffered I/O结果发现系统页缓存严重干扰了 KV 块的淘汰策略。操作系统会缓存很多 KV 块导致内存被吃掉一大半而 SSD 本身又很快页缓存的命中收益并不值得。后来改用O_DIRECT绕开系统页缓存直接读写块设备。这样内存占用稳步下降而且延迟更可预测。这里要提醒O_DIRECT要求 IO 和内存地址按 512 字节对齐。我们设计的 KV 块大小是 64KB天然满足对齐要求所以没有遇到额外麻烦。如果你的块大小设置成 16KB 或 32KB请务必检查对齐问题。6.3 检查点与 KV 卸载共用一个 NVMe 盘的相互影响早期我在同一块 NVMe 上同时承载 KV 预取读取和检查点写入。结果检查点写入的大流量会频繁打断 KV 块的读取导致 KV 预取 P99 延迟从 200 微秒飙升到 900 微秒。之后我把两块物理盘分开KV 走盘 A检查点走盘 B再也没出现这种互相干扰。一个常见误区是以为 NVMe 足够快就能混用但实际在混合负载下延迟离散度很致命。6.4 训练恢复时的 data race检查点流水线做得激进的早期版本发生过一种微妙错误一个优化器步骤尚未完全结束检查点线程就开始读取某个层的权重导致保存的检查点混入了新旧两种版本的参数。恢复训练后损失函数出现锯齿状波动极难排查。最终解决方案是版本标记每次反向传播参数更新完成后给该层权重打一个新的版本号检查点线程只能读取版本号 ≥ 目标版本号的权重。如果检查点线程发现某层版本落后就对该层加锁并等待完成。这样彻底消除了不一致。调试这个问题的痛苦经历告诉我分布式系统的检查点一定要重视“逻辑一致性”而不是只关注“文件完整性”。7. 把这套方案移植到你自己的集群最小配置清单7.1 最低硬件建议推理侧本地至少两块 NVMe SSD容量按“上下文长度 × 1.5 倍 KV 大小”估算如果你要跑 Llama 2 70B 的 100K 上下文建议 512GB ~ 1TB 的 SSD 空间主机内存至少 2GB pinned memory 用于预取缓冲训练侧每节点一块独立 NVMe 做检查点暂存容量至少为单卡检查点大小的节点分片网络如果要用远端 NVMe-oF至少 25GbE推荐 100GbE7.2 软件栈选型文件系统本地直接用 ext4 或 xfs不需要额外改造远端池建议使用原生 NVMe-oF 导出块设备再在应用层自管理 KV 索引KV 卸载代码自己实现一个带预取队列的模块核心依赖只有libaio或io_uring检查点流水线建议基于 PyTorch 的DistributedCheckpoint协议改写或者直接在 Megatron-LM 的检查点模块上按上面思路打补丁7.3 快速验证方法如果你想先复现 KV 卸载部分的效果不需要训练完整模型。用一个小规模的 7B 模型伪合成 80K token 上下文的 KV 缓存直接压测“页面调度版”和“预取块索引版”的生成吞吐差异就能看到数量级上的区别。我建议第一轮不要加量化先把索引和预取逻辑跑通再加 int8 优化这样定位性能问题更直接。训练检查点部分同理先用一个 1B 模型造一份 20GB 的假检查点挂在本地 NVMe 和远端文件系统上测一下同步写和异步流水线写的时间差再决定工程量投入。8. 成本与安全边界不是每个场景都值得上 NVMe 分层8.1 为什么小模型和短上下文不要这么玩如果你只跑 7B 模型、上下文 4K 以内KV 缓存不过几百 MB显存完全够用这时候去优化存储层只是徒增复杂度和故障面。同样检查点频率很低比如几千步一次的话同步停几秒也不是大问题。存储分层的收益曲线是非线性的只有 KV 缓存超过显存、检查点时间超过可接受停顿时间时收益才会显著放大。8.2 数据安全和备份策略我们利用本地 NVMe 做暂存极大提升了性能但也带来一个隐患如果节点宕机暂存在本地的检查点可能丢失。所以在生产环境里远端池仍然是唯一的事实标准存储。本地暂存只是一层加速缓存不能当作备份手段。KV 卸载同样有类似问题如果某个 SSD 损坏正在推理的长上下文上下文会丢失。这也意味着 KV 卸载更适合“可重算”的场景比如对话、摘要、代码补全如果是必须恢复的历史会话一定要把 KV 状态定期导出到可靠存储而不是只留在本地。8.3 我为什么坚持不让“存储层”替代“内存层”这个项目做完之后我的核心体会是存储层是内存层的补充与扩展而不是替代品。NVMe 再快也比不上显存和主机内存的带宽。我们的优化目标始终是让“最常访问的数据”留在最快层级存储层只处理那些真正超出一代硬件物理极限的数据。这个原则听起来简单但实践中很容易走偏——为了“秀存储性能”把本可以留在显存的 KV 块也搬到 SSD结果反而更慢。克制比技术本身更重要。最后再分享一个我在测试中发现的小技巧当你在做 KV 预取时把预取请求的优先级设为最高并且专门使用一个小的线程池不要让检查点上传线程混进来。这个看似不起眼的调整在混合负载下能把 KV 预取延迟的抖动降掉一半。存储分层不是把数据丢给硬件就完事数据流的调度和隔离才是决定收益上限的地方。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门