尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

SGLang Hybrid Linear attention Mamba KV cache管理

SGLang Hybrid Linear attention Mamba KV cache管理 RefKimi K3 Tech Blog: Open Frontier Intelligence当 Prefix Cache 遇见 KDAMooncake 如何 Day-0 支持 Kimi K3https://pytorch.org/blog/hybrid-models-meet-sglang-more-than-full-attention/Mamba模型简介mamba-ish可以理解为“具有 Mamba 类似状态管理特征的模型”。其中-ish是英语后缀表示“类似……的”。它不是一个严格的论文模型名称而是 SGLang 内部使用的工程分类。采用Linear和dense layer 3:1混合的模型当前有Qwen 3.6和kimi k3等。Qwen 3.6是GQA Linear而kimi-k3是DeepSeek MLA Linearsglang内部采用hybrid mamba模式来管理混合full和linear attention的kv cache.Dense/full-attention 层KV cache 按 token 追加显存随序列长度增长。Mamba/Gated-DeltaNet 层每层、每个请求只有一份conv state recurrent/SSM state。每处理一个 token 都会递推更新但属于原地覆盖不会为每个 token 保存一份。核心特征普通 Transformer attention 保存每个历史 token 的 KVtoken 1 → K₁, V₁ token 2 → K₂, V₂ ... token T → Kₜ, Vₜ缓存随着序列长度增长KV cache memory ∝ sequence lengthMamba、SSM 和部分线性注意力模型则维护递归状态state_t update(state_{t-1}, input_t)处理完一个 token 后历史信息被压缩进固定形状的状态中不需要为每个历史 token 保存一份 KVstate memory ≈ 固定大小 / requestSGLang 把需要这种“每请求一个递归状态 slot”管理方式的模型统称为mamba-ish。状态通常包含什么SGLang 中一个 Mamba-style slot 通常包含两部分conv state temporal/SSM stateConv state保存短卷积所需的最近几个输入[channels, conv_kernel - 1]它类似一个滑动窗口。Temporal/SSM state保存长期递归状态。Mamba2 中可能是标准 SSM stateGDN/KDA 中通常是类似下面的线性注意力状态矩阵[HV, V, K]虽然 GDN/KDA 从算法命名上不叫 Mamba但它们的推理状态具有相同的工程性质每个请求需要一个持久状态每生成一个 token 更新一次状态不是按 token 索引的 KV cacheprefix cache、复制、回滚和 speculative commit 都需要特殊处理。因此它们也被纳入mamba-ish。Mamba KV cache保存逻辑核心代码python\sglang\srt\mem_cache\mamba_radix_cache.pypython\sglang\srt\mem_cache\unified_cache_components\mamba_component.py对于dense layer的kv cache是按token存储的每个token存储一个固定大小的key和value cache对于MLA只需要存一个单独的hidden的cache。而对于linear attention的mamba state部分是在一个stage上循环累加的理论上一个请求最少只需要存储一个state张量。linear attention的mamba state不能像dense layer的kv cache那样逐token甚至page存储因为单个mamba state的容量很大存储太密集导致存储空间需求巨大而存储太稀疏可能导致命中率降低。sglang的mamba sate实际存储逻辑prefill部分每次chunked prefill完成后存储一次。例如输入13600 chunked prefill6144总共进行3次prefill计算存3个state。decode部分当(输入长度输出长度)%mamba_track_interval时更新一次mamba state但是只offload最终的那一个mamba state。也就是decode部分如果输出很短不会产生新的mamba state但是通常足够长会产生一个mamba state存储。因此mamba state存储数量最大为(in_len chunked_prefill_size -1)/chunked_prefill_size 1个状态数量或者说in_len // chunked_prefill_size 2。例如输入13600, 输出512 chunked prefill6144总共进行3次prefill计算存3个state然后decode完成存一次前提是decode相比prefill多一个page的token总共4个mamba state。Mamba state驱逐逻辑LRU驱逐Kimi-k3的kv cache大小与传统 Attention 不同KDA 并不会长期保存每个历史 token 的 Key 和 Value而是将历史信息不断递推recurrent到一个固定大小的状态中。模型会为每个 channel 学习不同的衰减系数并利用 Delta Correction 控制新信息写入状态的强度从而在有限状态容量下保留尽可能丰富的历史信息。同时为了兼顾局部建模能力KDA 还维护一个固定长度的 Convolution Window用于保存最近几个 token 的局部信息。因此对于每个 KDA 层而言真正需要持续维护的历史并不是一长串 KV而是两部分状态Temporal State递推更新的历史状态Convolution Window最近几个 token 的局部窗口。随着新的 token 到来这两部分状态都会不断原地更新而不会像传统 KV Cache 一样持续增长。这种设计最大的优势是推理过程中需要访问的缓存大小不再随上下文长度线性增长。即使面对百万级上下文KDA 层需要维护的递推状态依然保持固定规模大幅降低了长上下文推理的显存压力。kimi-k3 linear: dense3:1 93层 23模块x (3 linear 1 MLA) 最后一层 MLA也就是69 KDA linear 24 MLA dense层。Dense层的kv cache大小kimi-k3 dense层采用MLA因此kv cache大小与deepseek v3.1一致每一层576个元素。BF16 kv cache每个token的存储大小为kv_lora_rank qk_rope_head_dim 512 64 57624layer*576*2 (bf16) 27 KB。FP8直接所有token直接FP8量化没有像deepseek v3.2那样部分BF16部分FP8因此kv cache大小直接减半为每个token 13.5 KB。对于Dense层在TP并行非DP attention/DCP的情况下同一个请求每个GPU的KV cache是一模一样的。Linear的kv cache大小具体计算逻辑模型配置为KDA heads96head dimension128short-conv kernel4所以每个 TP rank 上local_heads 96 / TP # head per GPU K V 128 # head_dim conv_history kernel_size - 1 3 KDA_layers 69TP8的场景每个GPU的head数为96/8 12。KDA 不需要为每个历史 token 保存 K/V。它把历史压缩进一个固定大小的矩阵状态S: [local_heads, V, K] [12, 128, 128]同时KDA 输入前有一个 kernel size 为 4 的短卷积所以还要保存最近4-13个 Q/K/V 投影输入。因此每个状态 slot、每个 KDA 层有两部分conv state: [3, Q_local K_local V_local] SSM state: [head_num/TP, q_head_dim, v_head_dim]其中每个 Q/K/V 的本地宽度为12 heads × 128 1536所以conv state shape [3, 1536 × 3] [3, 4608] SSM state shape [12, 128, 128]shape分配在mamba2_cache_params函数中调用KimiLinearStateShape.create初始化。因此TP8并行时每个GPU的kv cache大小SSM大小为96/8*128*128*69(layer)*2 (bf16) 25.9MB。conv state大小为3*1536*3*69*2 1.82MB总的kv需要所有GPU加起来TP 8总和为(25.91.82)*8 221.8MB.Speculative decoding--enable-linear-replayssm-spec分配推测解码相关的state kv cache。KDA 的 recurrent update 可以概括为S S · Diag(alpha) d · kᵀ d beta · (v - (S · Diag(alpha)) · k)KDA 的alpha是逐 K-channel 的向量而不是每个 head 一个标量。普通 speculative target verify 如果一次验证 D 个候选 token为了最后只提交 accepted prefix通常需要保存每一步的完整 SSM state[num_layers, requests, D, heads, V, K]Kimi-K3 的[heads,V,K]状态很大对每个 draft token 保存一份成本非常高。ReplaySSM spec 改成verify 时仍然算出每一步输出不保存每一步完整的[V,K]state只保存生成这个 state 所需的轻量输入记录acceptance 结束后仅把被接受的前缀按原 recurrent 顺序重新播放将重放结果写回持久化 SSM checkpoint。KDA 的实现是“每次 commit 都 exact-fold”不是普通 decode ReplaySSM 所说的“每 L 步才 flush”。相应说明在 [kda_replayssm_spec_decode.py (line 11)](/D:/codes/open_engine/sglang/sglang_kimi_k3/python/sglang/kernels/ops/attention/fla/kda_replayssm_spec_decode.py:11)。--linear-replayssm-cache-len没有显式设置因此使用默认并发量16。所有 ring 都按 69 层、81 个 slot、16 个位置分配。BufferShapedtype含义d[69,81,12,16,128]BF16修正后的 delta/value 向量k[69,81,12,16,128]BF16归一化/缩放后的 keyg[69,81,12,16,128]FP32KDA 的逐 K-channel log-decay gaterawv[69,81,12,16,128]BF16exact-fold 使用的原始 value 输入rawk[69,81,12,16,128]BF16exact-fold 使用的归一化前 keybeta[69,81,12,16]FP32每个 head、每一步的 delta update 系数分配代码在 [memory_pool.py (line 591)](/D:/codes/open_engine/sglang/sglang_kimi_k3/python/sglang/srt/mem_cache/memory_pool.py:591)。d / kspec 模式下它们采用 conv/activation dtype即 BF1669 × 81 × 12 × 16 × 128 × 2 274,710,528 bytes 0.255844 GiB → 各显示 0.256GBgKDA 的 gate 是逐 K-channel 的所以包含最后一个 128 维而且强制使用 FP3269 × 81 × 12 × 16 × 128 × 4 549,421,056 bytes 0.511688 GiB → 显示 0.512GB这也直接证明日志虽然写着 “GDN”实际 shape 是 KDA真正 GDN 的g是每 head 一个标量shape 为[69,81,12,16]KDA 的g是 128 维向量所以恰好大 128 倍rawv / rawk两者 shape 和d/k一样都是 BF16各 0.255844 GiB → 各显示 0.256GBverify kernel 写入的是尚未做 delta correction 的v尚未做 L2 normalization 的kkernel 内实际形成的 FP32gsigmoid(b)后的 FP32 beta对应写入逻辑在 [fused_sigmoid_gating_recurrent.py (line 210)](/D:/codes/open_engine/sglang/sglang_kimi_k3/python/sglang/kernels/ops/attention/fla/fused_sigmoid_gating_recurrent.py:210)。beta69 × 81 × 12 × 16 × 4 4,292,352 bytes 0.003998 GiB → 显示 0.004GBRing 总大小d k g rawv rawk beta 1.539063 GiB每个 slot 跨 69 层20,401,920 bytes ≈ 19.457 MiB需要指出在 KDA spec 路径中真正用于 commit exact-fold 的主要是rawv rawk g betad/k主要是 GDN/普通 ReplaySSM reconstruction 需要的。当前统一内存池仍为 KDA 分配它们源码注释也明确说它们对 KDA fold 看起来是“dead weight”但暂时保留以避免改变 decode dispatch 行为。因此这里约有d k 0.511688 GiB/卡属于当前实现的额外开销。实际分配样例# no spark spec kv_cache_dtype bfloat16 mamba_ssm_dtypefloat32 Mamba Cache is allocated. max_mamba_cache_size: 51, conv_state size: 0.09GB, ssm_state size: 2.63GB GDN ReplaySSM ring buffers allocated (L16): d0.164GB, k0.164GB, g0.328GB rawv0.164GB, rawk0.164GB, beta0.003GB KV Cache is allocated. dtype: torch.bfloat16, #tokens: 974656, KV size: 25.10 GB kv_cache_dtypefp8_e4m3 mamba_ssm_dtypebfloat16 Mamba Cache is allocated. max_mamba_cache_size: 80, conv_state size: 0.14GB, ssm_state size: 2.05GB GDN ReplaySSM ring buffers allocated (L16): d0.256GB, k0.256GB, g0.512GB rawv0.256GB, rawk0.256GB, beta0.004GB KV Cache is allocated. dtype: torch.float8_e4m3fn, #tokens: 1947648, KV size: 25.08 GB实际分配的slot为size 1两种Attention KV大小设置上面介绍了mamba state存储数量每个请求需要大约为prefill // chunked_prefill_size 2个mamba state状态数量。单个token的dense kv cache和单个mamba state的内存占用大小是确定的。这导致一个后果不同的请求长度需要设置不同的mamba-full-memory-ratio短输入需要设置比较大的值而长输入需要设置小的值。设置不合理会导致推理的并发量受限于主kv或者是Mamba kv例如下面这个例子主kv还有很大的空闲但是mamba state kv已经用满了导致并发被限制# in 16k out 3k Decode batch, #running-req: 16, #full token: 206080, full token usage: 0.11, mamba num: 64, mamba usage: 0.80max_running_requests is capped to 16 by the mamba state cache (max_mamba_cache_size80, 5 state slots per request). To raise it: increase --mamba-full-memory-ratio or --max-mamba-cache-size, or halve the state size with --mamba-ssm-dtype bfloat16.resolve_max_num_reqs里面使用分配的总的mamba state除以_calculate_mamba_ratio()计算的每个请求预留的mamba state来计算最大的并发数量例如总共分配80个slot每个请求预留5个 5 基础安全容量 3 overlap ping-pong buffer 2那么只能并发16。这时候要提升并发需要增大mamba内存分配比例或者通过SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK1设置把预留数量降低。相关参数设置--max-mamba-cache-size人工指定整个 Mamba 状态池最多有多少物理槽位。这个参数什么时候有用因为系统需要每个请求预留4或者5个slot的mamba state因此最小的slot数量分配就是期望的decode并发数乘以预留slot数量。--mamba-full-memory-ratio设置sglang启动时 Mamba 状态与 Full KV 的显存预算比例通过full和mamba的kv cache比例来自动计算mamba state槽位。sglang官方的计算器Kimi-K3 - SGLang Documentation这个计算公式也有一些缺陷没有考虑sglang的chunked prefill存储逻辑因为这个比例计算跟chunked prefill size有关如果通过--mamba-max-states-per-path设置了每个请求最大的mamba state数量这个公式也需要修改。当前每个请求预留了4-5个slot设置SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK1时为4否则为5还需要根据这个值和最大并发量来确定最优ratio比例。只开TP不开启DCP时mamba:dense_kv因为mamba部分按GPU进行head切分因此大小/TP数量因此比例更低而开启DCP时没有这个冗余这个mem ratio要乘以TP数。--mamba-max-states-per-path每条 Radix 路径保留多少个历史 Mamba 检查点。整个请求的Dense KV cache在运行未结束前是不能进行驱逐的所有token的kv cache必须保持但是运行中请求的mamba state只需要保留一份可以进行驱逐。这种情况的mamba kv cache是从根部往尾部驱逐而不是从尾部往根部驱逐。而请求之间的前缀匹配进行驱逐的时候两者都应该从尾部往根部驱逐。从上面介绍可以看到请求越长分配的mamba state越多但是前面部分的mamba state可能对缓存命中率的贡献并不是很大但是却占用大量存储。--mamba-max-states-per-path可以减少长会话不断延伸时积累的历史状态例如--mamba-max-states-per-path 3这会让历史路径释放更多槽位给新请求使用但代价是请求从较浅前缀分叉时可能找不到对应的 Mamba 状态需要从更早的状态重新计算Mamba prefix-cache 命中效果可能下降若存在 HiCache host backupGPU 状态被删除后 host 副本仍保留但命中时需要重新加载。overlap schedule时至少每个请求需要2个slot。新请求进入时如果空闲 slot 不足SGLang 会自动从 Radix Cache 中LRU 驱逐“未锁定、可驱逐”的历史 Mamba checkpoint活跃请求正在使用或被锁定的状态不会被驱逐。因此这个--mamba-max-states-per-path通常并不需要设置。--mamba-track-interval控制输出部分的mamba state更新逻辑当前默认256。例如 prompt 长度为1000、interval 为256decode checkpoint 会在总长度1024、1280、1536……也就是分别 decode 约24、280、536……个已处理 token 后触发而不是固定在 decode 输出长度256、512……时触发。最终 HiCache 通常只 offload 最新边界对应的一个 Mamba state slot。值越小缓存粒度更细、前缀命中后需要重算的 token 更少但状态保存更频繁显存和执行开销可能增加。值越大保存开销更低但缓存粒度更粗前缀复用效果可能降低。mamba_track_interval的核心存储链路分为四步判断是否到达存储边界在 [schedule_batch.py (line 2999)](/D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/managers/schedule_batch.py:2999)mamba_track_interval get_exec().mamba.mamba_track_interval self.mamba_track_mask ( self.seq_lens_cpu % mamba_track_interval 0 )只有序列长度为 interval 整数倍的请求mamba_track_mask才为True。确定快照目标槽位在 [schedule_batch.py (line 1796)](/D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/managers/schedule_batch.py:1796) 的set_mamba_track_indices_from_reqs()中根据请求的 ping-pong buffer 生成batch.mamba_track_indices它表示当前 Mamba 状态应该写入 Mamba state pool 的哪个槽位。真正复制 Mamba 状态主要入口在 [hybrid_linear_attn_backend.py (line 711)](/D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py:711)track_mamba_states_if_needed( conv_states, ssm_states, cache_indices, # 当前运行状态 forward_batch.mamba_track_mask, # 是否到达 interval self.forward_metadata.mamba_track_indices, # 快照目标槽位 ... )真正执行复制的是 Triton kernel[mamba_state_scatter_triton.py (line 13)](/D:/codes/open_engine/sglang/sglang_github/python/sglang/kernels/ops/mamba/mamba_state_scatter_triton.py:13)其核心逻辑等价于if mamba_track_mask[i]: conv_states[track_slot] conv_states[active_slot] ssm_states[track_slot] ssm_states[active_slot]也就是说存储的是两部分convolution stateSSM/recurrent state它们存进 Mamba state pool 的额外 tracking slot而不是普通 token KV Cache。更新 checkpoint 元数据并插入 Radix Cacheforward 完成后[batch_result_processor.py (line 1074)](/D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/managers/scheduler_components/batch_result_processor.py:1074) 会记录req.mamba_last_track_seqlen track_seqlen非 lazy 策略还会切换 ping-pong 槽位req.mamba_next_track_idx other_idx请求完成或中途缓存时[mamba_radix_cache.py (line 544)](/D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/mem_cache/mamba_radix_cache.py:544) 使用cache_len req.mamba_last_track_seqlen mamba_value src_active.clone() self.insert(...)把 tracking slot 对应的状态与 token 前缀一起挂到 Radix Tree 节点上。整体流程是seq_len 到达 interval 边界 → mamba_track_maskTrue → 把当前 conv/SSM 状态复制到 ping-pong tracking slot → 记录 mamba_last_track_seqlen → 请求缓存时将该 slot 插入 Radix Cache--mamba-radix-cache-strategy可选项MAMBA_RADIX_CACHE_STRATEGY_CHOICES [ auto, no_buffer, extra_buffer, extra_buffer_lazy, ]默认为auto也就是extra_buffer模式。--mamba-radix-cache-strategy决定混合 Attention Mamba/GDN/KDA 模型如何保存线性注意力的循环状态。它主要在三件事之间权衡是否启用 overlap scheduler。能否缓存 Radix Tree 分叉点上的 Mamba 状态。每个运行中请求要预留多少 Mamba state slot从而影响最大并发。四种策略对比策略Overlap scheduler分叉点状态缓存每请求容量预留适用场景auto自动决定自动决定取决于解析结果通常首选no_buffer不支持未实现3 slots显存紧张、兼容性优先、ReplaySSMextra_buffer支持支持overlap 开启时 5 slots吞吐优先、稳定生产配置extra_buffer_lazy必须开启支持4 slotsMamba state 容量成为瓶颈时auto开启 overlap schedule和page_size1默认设置为extra_buffer。no_buffer不支持overlap schedule只支持page_size1。extra_buffer普通extra_buffer在 overlap 开启时为每个请求预先分配两个 track slottrack slot ACPU/Radix Cache 可以安全读取的旧快照 track slot BGPU forward 正在写入的新快照下一轮两者交换也就是 ping-pong第 t 轮 读取 A写入 B 第 t1 轮读取 B写入 A之所以需要两个是因为 overlap scheduler 允许CPU处理上一轮结果、更新 Radix Cache GPU同时执行下一轮 forward只用一个快照槽时CPU 读取状态和 GPU 覆盖状态可能发生竞争。代码直接定义self.mamba_ping_pong_track_buffer_size ( 2 if enable_overlap_schedule else 1 )见 [memory_pool.py (line 1142)](D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/mem_cache/memory_pool.py:1142)。普通extra_buffer会在请求进入时一次性申请全部两个 slot见 [memory_pool.py (line 1372)](D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/mem_cache/memory_pool.py:1372)。因此最终容量系数是base 3 overlap ping-pong 2 --------------------- 总计 5对应代码MAMBA_CACHE_SIZE_MAX_RUNNING_REQUESTS_RATIO 3 MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP 2见 [kv_cache_configurator.py (line 108)](D:/codes/open_engine/sglang/sglang_github/python/sglang/srt/mem_cache/kv_cache_configurator.py:108)。extra_buffer_lazyextra_buffer_lazy unsupported under PD disaggregation;SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK1frees one more slot per request (experimental, under validation)Unified Memory Pool实现路径python\sglang\srt\mem_cache\unified_memory_pool.py当前问题与多项技术不兼容--enable-unified-memory is not yet compatible with PD disaggregation.--enable-unified-memory is not yet compatible with speculative decoding.--enable-unified-memory is not yet compatible with hierarchical host-tiered KV cache--enable-unified-memory is not yet compatible with decode context parallelism (--dcp-size 1)针对设置不合理的--mamba-full-memory-ratio会导致无法自适应不同的业务场景问题sglang社区提出了主kv和mamba state共用一份kv cache的方案--enable-unified-memoryReplace the statically-partitioned hybrid-model pools (full-attn KV SWA/Mamba state) with one byte buffer split dynamically between sub-pools. Requires the Triton attention / linear-attn / Mamba backends; not yet compatible with PD disaggregation or speculative decoding.它让 Full KV Cache 和 Mamba Cache 共用同一块显存两边从相反方向动态增长长请求多更多显存用于 Full KV token。短请求多空闲的 Full KV 显存可动态转成更多 Mamba 槽位。mamba_full_memory_ratio只参与确定启动时总预算不再固定运行时的两边分界。开启enable_unified_memory后低地址 高地址 | Mamba/SWA → → → 动态空闲区 ← ← ← Full KV | 边界随运行负载变化Full KV 从高地址向下增长Mamba/SWA 从低地址向上增长直到两边相遇。底层物理存储系统只分配一块 GPU 字节缓冲区self._raw torch.empty(total_bytes, dtypetorch.uint8, devicedevice)然后在它上面构造不同的 Tensor viewMHAK/V viewMLA每层 dense viewMambaconv state 和 temporal state viewSWAFull KV 和 SWA KV view这些 view 指向同一块物理显存但分配器保证两边实际占用的字节区域不重叠。实现见class UnifiedKVPool。
返回列表