Prhub

#22493 Add MambaPool kvcache offloading during retraction

原始 PR 作者 hlu1 合并时间 2026-04-22 08:51 文件变更 5 提交数 3 评论 14 代码增减 +193 / -16

执行摘要

在请求回退期间添加 MambaPool 的 KV 缓存卸载,确保 Mamba 混合模型的完整状态保存和恢复。

根据PR body描述,Mamba混合模型(如Qwen3.5-397B-A17B)将SSM状态(conv和temporal缓冲区)存储在MambaPool中,与注意力KV缓存分离。在请求回退期间,只有KV缓存被保存到CPU,而Mamba状态被丢弃,导致回退请求的生成损坏。

该PR值得精读,重点关注HybridLinearKVPool的设计,学习如何通过可选参数扩展缓存池以支持多种状态类型的联合卸载,以及CPU-GPU数据传输的同步机制。

讨论亮点

Review中主要讨论了kv_cache_cpu的内容澄清:

  • yizhang2077询问kv_cache_cpu是否同时包含KV缓存和Mamba状态。
  • hzh0425进一步质疑Mamba状态存储位置。
  • hlu1澄清实现细节,指出HybridLinearKVPool.get_cpu_copy会复制两者到CPU,并在后续提交中添加注释以明确。

实现拆解

  1. MambaPool CPU快照方法:在python/sglang/srt/mem_cache/memory_pool.py的MambaPool类中添加get_cpu_copyload_cpu_copy方法,使用torch.cuda.synchronize()确保数据一致性,将conv和temporal状态复制到CPU并恢复。
  2. HybridLinearKVPool扩展:在同一文件中扩展HybridLinearKVPool的get_cpu_copyload_cpu_copy方法,通过可选mamba_indices参数联合处理KV缓存和Mamba状态的卸载,对非混合模型保持无操作。
  3. 请求状态卸载集成:修改python/sglang/srt/managers/schedule_batch.pyReq.offload_kv_cacheload_kv_cache方法,传递mamba_pool_idx参数以触发完整状态卸载,并添加注释澄清kv_cache_cpu存储KV和Mamba状态元组。
  4. 调度器日志增强:更新python/sglang/srt/managers/scheduler.pyupdate_running_batch方法,添加#mamba_num_gained日志以追踪回退期间Mamba池状态变化。
  5. 测试配套:在test/registered/unit/mem_cache/test_mamba_unittest.py中添加两个单元测试test_mamba_pool_cpu_offloadtest_hybrid_kv_pool_cpu_offload,验证CPU卸载往返的正确性。
文件 模块 状态 重要度
python/sglang/srt/mem_cache/memory_pool.py 内存缓存池 modified 7.64
test/registered/unit/mem_cache/test_mamba_unittest.py 单元测试 modified 6.14
python/sglang/srt/mem_cache/allocator.py 分配器 modified 5.86
python/sglang/srt/managers/scheduler.py 调度器 modified 5.35
python/sglang/srt/managers/schedule_batch.py 调度批次 modified 5.51

关键符号

get_cpu_copy load_cpu_copy

关键源码片段

python/sglang/srt/mem_cache/memory_pool.py core-logic

实现 MambaPool 和 HybridLinearKVPool 的 CPU 卸载核心逻辑,是功能变更的主要文件。

class MambaPool:
    # ... 其他方法省略 ...
​
    def get_cpu_copy(self, indices, **kwargs):
        # 同步 CUDA 以确保数据一致性,防止异步操作导致状态不一致
        torch.cuda.synchronize()
        # 将 Mamba 缓存中的 conv 状态(列表形式)复制到 CPU,使用非阻塞传输提升性能
        conv_cpu = [
            conv[:, indices].to("cpu", non_blocking=True)
            for conv in self.mamba_cache.conv
        ]
        # 将 temporal 状态复制到 CPU
        temporal_cpu = self.mamba_cache.temporal[:, indices].to(
            "cpu", non_blocking=True
        )
        torch.cuda.synchronize()
        return conv_cpu, temporal_cpu # 返回 CPU 上的 Mamba 状态元组
​
    def load_cpu_copy(self, mamba_cache_cpu, indices, **kwargs):
        conv_cpu, temporal_cpu = mamba_cache_cpu
        torch.cuda.synchronize()
        # 将 CPU 上的 conv 状态恢复至 GPU,使用非阻塞传输
        for i, conv in enumerate(self.mamba_cache.conv):
            conv[:, indices] = conv_cpu[i].to(conv.device, non_blocking=True)
        # 恢复 temporal 状态
        self.mamba_cache.temporal[:, indices] = temporal_cpu.to(
            self.mamba_cache.temporal.device, non_blocking=True
        )
        torch.cuda.synchronize()

评论区精华

kv_cache_cpu 内容澄清 正确性

yizhang2077 和 hzh0425 在 schedule_batch.py 中询问 kv_cache_cpu 是否同时存储 KV 缓存和 Mamba 状态,担心数据存储逻辑不清晰。

结论:hlu1 澄清 HybridLinearKVPool.get_cpu_copy 会复制两者到 CPU,并在后续提交中添加注释以明确实现细节。 · 已解决

风险与影响

技术风险包括:

  • 状态不一致风险:如果mamba_indices传递错误或处理不当,可能导致Mamba状态未正确保存或恢复,影响生成准确性。
  • 性能影响:CPU卸载涉及GPU-CPU数据传输和同步,可能增加延迟,尤其是在高频回退场景下。
  • 兼容性问题:修改了get_cpu_copyload_cpu_copy接口为**kwargs,需确保所有调用方兼容,但已通过默认参数保持向后兼容。

对使用Mamba混合模型的用户,修复了请求回退时生成损坏的问题,提升了系统健壮性和推理可靠性;对非混合模型无影响,因为变更通过可选参数实现。影响范围限于涉及缓存卸载和回退调度的模块,如内存池、调度器和请求管理。

状态一致性风险 核心路径变更

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论