SGLang之MambaRadixCache机制

鱿鱼圈 Lv4

本文详细解析 SGLang 中 MambaRadixCache 的内存空间布局、KV Cache 与 SSM Cache 的缓存机制原理、Prefill 和 Decode 的完整流程,并通过端到端的例子串联所有概念。

1. 内存空间布局

1.1 三大内存池

MambaRadixCache 管理三块独立的 GPU 显存池:每 256 token 做一次 track(不是为了 tree 缓存,而是为了):

Chunked prefill 被抢占: scheduler 打断后,需要从 checkpoint 恢复

Speculative decoding 回滚: draft 被 reject 后回退到最近 checkpoint

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
┌──────────────────────────────────────────────────────────────────────────────┐
│ GPU 显存 │
│ │
│ ┌────────────────────────────────┐ ┌────────────────────────────────────┐ │
│ │ KV Pool (MHATokenToKVPool) │ │ SSM Pool (MambaPool) │ │
│ │ │ │ │ │
│ │ k_buffer[layer][slot][H][D] │ │ conv[layer][slot][D][K-1] │ │
│ │ v_buffer[layer][slot][H][D] │ │ temporal[layer][slot][HV][K][V] │ │
│ │ │ │ │ │
│ │ 共 kv_size 个 slot │ │ 共 mamba_size 个 slot │ │
│ │ 按 token 粒度分配 │ │ 按 状态(checkpoint) 粒度分配 │ │
│ └────────────────────────────────┘ └────────────────────────────────────┘ │
│ │
│ ┌──────────────────────────────────────────────────────────────────────────┐│
│ │ ReqToTokenPool (映射表) ││
│ │ req_to_token[req_pool_idx, pos] → KV Pool slot index ││
│ │ req_index_to_mamba_index_mapping[req_pool_idx] → SSM Pool slot index ││
│ └──────────────────────────────────────────────────────────────────────────┘│
└──────────────────────────────────────────────────────────────────────────────┘

1.2 KV Pool 空间

1
2
3
# memory_pool.py: MHATokenToKVPool._create_buffers
k_buffer = [torch.zeros((size, head_num, head_dim)) for _ in range(layer_num)]
v_buffer = [torch.zeros((size, head_num, v_head_dim)) for _ in range(layer_num)]
  • 分配粒度: 1 slot = 1 个 token 在所有 Full Attention 层的 K 和 V
  • 容量: size 个 slot(例如 100,000)
  • 一个 request 占用: seq_len 个 slot(随序列长度线性增长)
  • 管理方式: TokenToKVPoolAllocator 维护 free list

1.3 SSM Pool 空间

1
2
3
# memory_pool.py: MambaPool.__init__
conv_state = [torch.zeros((num_mamba_layers, size+1, D, K-1))] # 卷积状态
temporal = torch.zeros((num_mamba_layers, size+1, HV, K, V)) # SSM 递归状态
  • 分配粒度: 1 slot = 所有 Mamba 层的完整隐状态(conv + temporal)
  • 容量: mamba_size 个 slot(例如 2,000)
  • 单个 slot 大小: 固定,与序列长度无关(例如 ~112MB for Kimi-K2)
  • 一个 request 占用: 1 个工作 slot + 2 个 ping-pong track slot = 共 3 个
  • 管理方式: MambaPool 维护 free_slots tensor

1.4 ReqToTokenPool 映射表

1
2
# memory_pool.py: ReqToTokenPool
req_to_token = torch.zeros((req_pool_size, max_context_len), dtype=torch.int32)

每个 request 分配一行,记录:

  • req_to_token[req_pool_idx, pos] = token 在 KV Pool 中的 slot index
  • req_index_to_mamba_index_mapping[req_pool_idx] = SSM Pool 的 slot index

1.5 空间对比

维度 KV Pool SSM Pool
每 slot 存什么 1 token 的 K、V(所有 full-attn 层) 完整 SSM state(所有 mamba 层的 conv + temporal)
slot 大小 小(~KB 级) 大(~百 MB 级)
每 request 占用 O(seq_len) 个 固定 1~3 个
总容量决定 能同时服务多少 token 能缓存多少个前缀 checkpoint
淘汰 LRU full_lru_list mamba_lru_list(独立)

2. Radix Tree 数据结构

2.1 TreeNode 定义

1
2
3
4
5
6
7
8
class TreeNode:
key: RadixKey # 该节点对应的 token 子序列
value: Optional[torch.Tensor] # KV Pool indices(每 token 一个)
mamba_value: Optional[torch.Tensor] # SSM Pool index(就 1 个整数)
children: dict # 子节点
parent: TreeNode # 父节点
full_lock_ref: int # KV 引用计数(被多少 request 共享)
mamba_lock_ref: int # SSM 引用计数

关键点:value mamba_value 存的都是 Pool 的 slot index(指针),不是数据本身。

2.2 Radix Tree 如何存储前缀

1
2
3
4
5
6
7
8
9
10
11
示例:3 个 request 分别处理过以下 prompt:
R1: "You are a helpful assistant."
R2: "You are a helpful assistant. Write code."
R3: "You are a coding expert."

Tree 结构:
root
├── ["You are a "]
│ ├── ["helpful assistant."] → KV=[10..39], SSM=[slot 5]
│ │ └── [" Write code."] → KV=[40..50], SSM=[slot 8]
│ └── ["coding expert."] → KV=[51..65], SSM=[slot 12]

2.3 Tombstone 机制

当 SSM Pool 不够时,mamba_lru_list 淘汰最久未用的节点的 SSM state:

1
2
淘汰前: node.mamba_value = tensor([5])  ← SSM 有效
淘汰后: node.mamba_value = None ← Tombstone(KV 还在,SSM 没了)

Tombstone 节点:

  • KV 仍可用(Full Attention 前缀匹配仍有效)
  • 但不能作为 Mamba 的有效匹配点(没有 SSM state 就无法从此继续)

3. 缓存机制原理

3.1 前缀共享——不 copy 数据,只共享 index

核心设计:多个 request 可以通过 radix tree 共享同一批 KV Pool slot,而不是各自持有副本。

1
2
3
4
5
6
7
8
Request A 和 B 共享前缀 "You are a helpful assistant.":

KV Pool: [... slot 10 ...] [... slot 11 ...] ... [... slot 39 ...]
↑ ↑ ↑
Tree node.value: [10, 11, ..., 39] ← 只有一份
↑ ↑ ↑
req_to_token[A]: [10, 11, ..., 39, 新token...] ← 共享前缀
req_to_token[B]: [10, 11, ..., 39, 新token...] ← 共享前缀

通过 lock_ref 保护共享节点不被淘汰。

3.2 SSM Copy-on-Write (COW)

SSM state 不能直接共享(因为后续 forward 会修改它),所以用 COW:

  1. match_prefix 命中带 mamba_value 的节点
  2. 分配一个新的 SSM Pool slot 给 request
  3. 记录 req.mamba_cow_src_index = node.mamba_value(延迟 copy)
  4. Forward 开始时执行真正的 copy:
1
2
3
# MambaPool.copy_from:
mamba_cache.conv[:, dst_slot] = mamba_cache.conv[:, src_slot]
mamba_cache.temporal[:, dst_slot] = mamba_cache.temporal[:, src_slot]

之后 request 在自己的 slot 上独立更新,tree 中的原始 slot 不受影响。

3.3 双 LRU 独立淘汰

1
2
3
4
5
6
7
full_lru_list:  管理 KV Pool slot 的淘汰
mamba_lru_list: 管理 SSM Pool slot 的淘汰

两者独立运行:
- 内存紧张时可以只淘汰 SSM(节点变 tombstone),KV 保留
- 也可以只淘汰 KV(节点从 tree 删除)
- 淘汰 SSM 比淘汰 KV 更常见(SSM slot 更少更贵)

3.4 Ping-Pong Track Buffer

解决 overlap schedule 下的读写竞争问题:

1
2
3
4
5
每个 request (extra_buffer 模式) 的 SSM 资源:
├── mamba_pool_idx (1 slot) → forward 的工作 state
└── ping_pong_track_buffer (2 slots) → checkpoint 暂存
├── slot A (ping)
└── slot B (pong)

Overlap schedule 下,batch N+1 的 forward 和 batch N 的 cache 操作并行执行:

  • Forward 写入 slot A
  • 同时 cache_unfinished_req 读取 slot B(上次写完的)
  • 读写分离,互不干扰

每隔 mamba_track_interval(默认 256)token 做一次 checkpoint,交替使用 A/B。

4. Prefill 流程

4.1 非 Chunked Prefill(短 prompt 一次处理完)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
Request: "You are a helpful assistant."

1. match_prefix:
└── 从 tree root 查找最长匹配前缀
└── 假设找到 "You are a " 有 SSM state
└── 返回: prefix_indices=[10..19], last_node=匹配节点
└── COW: alloc 新 SSM slot, 记录 cow_src_index

2. alloc:
└── 分配 KV Pool slot 给剩余 token "helpful assistant."
└── req_to_token[req, 0:10] = [10..19] (来自 tree)
└── req_to_token[req, 10:30] = [新分配的 slot]

3. forward_extend (prefill):
└── _execute_deferred_mamba_cow_and_clear:
└── copy SSM state 从 tree 的 slot 到 request 的新 slot
└── Full Attention: 只计算 "helpful assistant." 的 KV(前缀的 KV 直接从 pool 读)
└── Mamba: 从 COW 的 state 继续递归计算
└── Track: 在 chunk 对齐位置保存 SSM checkpoint 到 ping-pong buffer

4. cache_finished_req:
└── 把 ping-pong buffer 中最后的 SSM slot "捐赠"给 tree
└── insert(key=完整token序列, value=KV indices, mamba_value=SSM slot)
└── Tree 创建新节点(或复用已有节点)
└── dec_lock_ref: 解锁旧节点

4.2 Chunked Prefill(长 prompt 分多次处理)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
Request: "ABCDEFGH" (8 tokens, chunk_size=4)

═══ Chunk 1: [A,B,C,D] ═══

1. match_prefix → 假设未命中,从 root 开始
prefix_indices=[], cache_protected_len=0

2. alloc KV slot [10,11,12,13] for tokens [A,B,C,D]
req_to_token = [10, 11, 12, 13, ?, ?, ?, ?]

3. forward_extend:
Mamba: 从零开始递归计算
Track: 保存 SSM state 到 ping-pong buffer (假设 slot A = pool slot 100)

4. cache_unfinished_req:
token_ids = [A,B,C,D] (fill_ids,累积所有已处理的 token)
kv_indices = [10,11,12,13]

insert(key=[A,B,C,D], value=[10,11,12,13], mamba_value=slot 100,
prev_prefix_len=0)
→ 创建 tree 节点: key=[A,B,C,D], value=[10,11,12,13], mamba_value=[100]

match_prefix 回来 → new_indices=[10,11,12,13]
req_to_token[req, 0:4] = [10,11,12,13] (写回 tree 的 indices)
cache_protected_len = 4
lock 新节点

═══ Chunk 2: [E,F,G,H] ═══

1. match_prefix → 命中 [A,B,C,D] 节点(有 SSM state)
prefix_indices=[10,11,12,13], 不需要 COW(自己的 state 还在工作 slot 里)

2. alloc KV slot [20,21,22,23] for tokens [E,F,G,H]
req_to_token = [10, 11, 12, 13, 20, 21, 22, 23]

3. forward_extend:
Full Attention: 用 [10..13] 的 KV 做 prefix, 只计算 [E,F,G,H]
Mamba: 从 checkpoint 继续
Track: 保存新 SSM state 到 ping-pong buffer (假设 slot B = pool slot 101)

4. cache_unfinished_req (或 cache_finished_req):
token_ids = [A,B,C,D,E,F,G,H] ← fill_ids 累积了全部
kv_indices = [10,11,12,13,20,21,22,23]

insert(key=[A,B,C,D,E,F,G,H], value=[10,11,12,13,20,21,22,23],
mamba_value=slot 101, prev_prefix_len=4)

while 循环匹配 [A,B,C,D]:
prev_prefix_len(4) < total(0) + prefix_len(4) → 4 < 4 → False
不 free [10,11,12,13](这些是 lock 的共享 slot)

剩余 key=[E,F,G,H] → 创建新节点:
node.value=[20,21,22,23], node.mamba_value=[101]

Tree: root → [A,B,C,D](KV:[10..13], SSM:[100])
└── [E,F,G,H](KV:[20..23], SSM:[101])

4.3 free 重复 KV 的场景

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
接上面,假设在 Chunk 1 和 Chunk 2 之间,Request B 已经 cache 了 [A,B,C,D,E,F]:

Tree: root → [A,B,C,D](KV:[10..13], SSM:[100])
└── [E,F](KV:[30,31], SSM:[105]) ← Request B 留下的

现在 Request A 的 Chunk 2 完成:
insert(key=[A,B,C,D,E,F,G,H], value=[10,11,12,13,20,21,22,23],
prev_prefix_len=4)

while 循环第一轮: 匹配 [A,B,C,D], prefix_len=4
4 < 0+4 → False, 不 free

while 循环第二轮: 匹配 [E,F], prefix_len=2
4 < 4+2 → True!
start = max(0, 4-4) = 0
free(value[0:2]) → free([20, 21]) ← A 自己算的 E,F 的 KV slot

为什么 free? 因为 tree 已有 [E,F] 的 KV (slot [30,31]),
A 自己算的 [20,21] 是多余副本。

剩余 key=[G,H] → 创建新节点:
node.value=[22,23], node.mamba_value=[101]

后续 match_prefix 回来会把 req_to_token 更新为:
req_to_token = [10, 11, 12, 13, 30, 31, 22, 23]
└─ tree 的 ─┘ └ tree ┘ └─ 新 ─┘

5. Decode 流程

5.1 单 token decode

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
Request 在 decode 阶段,每次生成 1 个 token:

1. scheduler 从 running batch 中取出 request

2. forward_decode:
Full Attention: 用 req_to_token 中所有历史 KV slot 做 attention
新 token 的 KV 写入新分配的 slot
Mamba: 读 mamba_pool_idx 对应的 state,更新后写回(in-place)

3. 每隔 mamba_track_interval (256) token:
把当前 SSM state copy 到 ping-pong buffer
翻转 mamba_next_track_idx (0→1 或 1→0)

4. 生成 EOS → cache_finished_req:
从 ping-pong buffer 取最后一个 checkpoint
insert 到 radix tree

5.2 为什么 decode 中间不存 SSM 到 tree

  • Decode 产生的每个 token 是该 request 独有的采样结果
  • 不会有其他 request 产生完全相同的 token 序列
  • 所以中间 SSM state 永远不会被其他 request 命中
  • 存了也是浪费 SSM Pool 空间

唯一有价值的是 最终 state(用于多轮对话的下一轮复用)。

5.3 Track Interval 的权衡

每 256 token 做一次 track(不是为了 tree 缓存,而是为了):

  • Chunked prefill 被抢占: scheduler 打断后,需要从 checkpoint 恢复
  • Speculative decoding 回滚: draft 被 reject 后回退到最近 checkpoint
Track 频率 Decode 性能 被抢占时重算量
每 token 慢(每次多一次 memcpy) 0
每 256 token 最多 255 token
不 track 最快 全部

6. 完整端到端例子

场景设定

1
2
3
模型: Kimi-K2 (KDA hybrid, 28 full-attn layers + 32 mamba layers)
Request 1: "You are a helpful assistant." (30 tokens)
Request 2: "You are a helpful assistant. Write Python code." (45 tokens)

6.1 Server 启动 + Warmup

1
2
3
4
5
6
7
8
9
10
11
12
启动参数: --warmups kda_cache --warmup-prompts-file prompts.json
prompts.json: ["You are a helpful assistant."]

1. kda_cache warmup:
构造 GenerateReqInput(text="You are a helpful assistant.", max_new_tokens=1)
走完整 pipeline → prefill 30 tokens → 生成 1 token → 结束

2. cache_finished_req:
insert(key=[tok0..tok30], value=[KV slot 1..31], mamba_value=[SSM slot 1])

3. Tree 状态:
root → ["You are a helpful assistant. {1_token}"](KV:[1..31], SSM:[slot 1])

6.2 Request 1 到来

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
prompt = "You are a helpful assistant. Tell me a joke."
tokens = [tok0..tok29, tok30..tok39] (40 tokens)

1. match_prefix:
沿 tree 匹配 → 命中 "You are a helpful assistant." (30 tokens)
last_node 有 mamba_value=[slot 1]

COW:
alloc 新 SSM slot 50 给 request 1
req.mamba_cow_src_index = [slot 1]
req.mamba_pool_idx = slot 50

返回 prefix_indices = [1, 2, ..., 30] (tree 中的 KV indices)

2. alloc:
分配 KV slot [100,101,...,109] 给剩余 10 token "Tell me a joke."
req_to_token = [1,2,...,30, 100,101,...,109]
└── 共享 tree ──┘ └── 新分配 ──┘

3. forward_extend:
① _execute_deferred_mamba_cow_and_clear:
mamba_pool.copy_from(src=[slot 1], dst=[slot 50])
→ SSM state 从 tree 的 slot 1 copy 到 request 的 slot 50

② Full Attention:
prefix token [0:30]: 直接读 KV Pool slot [1..30](已有,不计算)
新 token [30:40]: 计算 KV,存入 slot [100..109]

③ Mamba:
从 slot 50 的 state 开始(= "You are a helpful assistant." 处理完的 state)
递归计算 " Tell me a joke." 的 10 个 token
slot 50 的 state 被更新(现在包含全部 40 token 的信息)

4. Decode (假设生成 5 个 token):
每个 token: Full Attention attend 所有历史 + Mamba 更新 slot 50
KV Pool 再分配 5 个 slot

不触发 track(5 < 256)

5. EOS → cache_finished_req:
token_ids = [tok0..tok44] (40 prompt + 5 generated)
kv_indices = [1,...,30, 100,...,109, 新slot×5]
mamba_value = [slot 50] (最终 SSM state)

insert:
while 循环匹配已有的 "You are a helpful assistant." (30 tokens)
→ 不 free(prev_prefix_len=30,全在保护区)
free(重复部分的KV) 如果有
创建新节点 " Tell me a joke.{5_tokens}"

Tree:
root → ["You are a helpful assistant."](KV:[1..30], SSM:[slot 1])
├── [" {warmup的1token}"](KV:[31], SSM:[slot X])
└── [" Tell me a joke.{5tok}"](KV:[100..114], SSM:[slot 50])

6.3 Request 2 到来(复用前缀)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
prompt = "You are a helpful assistant. Write Python code."
tokens = [tok0..tok29, tok30..tok44] (45 tokens)

1. match_prefix:
匹配 "You are a helpful assistant." → 命中! last_node 有 SSM:[slot 1]

COW: alloc slot 60, 记录 cow_src=[slot 1]
返回 prefix_indices = [1..30]

2. alloc:
分配 KV slot [200..214] 给 " Write Python code." (15 tokens)
req_to_token = [1,...,30, 200,...,214]

3. forward_extend:
① COW: copy SSM slot 1 → slot 60
② Full Attention: 只计算 15 个新 token
③ Mamba: 从 slot 60 继续递归 15 个 token

对比没有 warmup:
需要从头计算 45 个 token 的 Mamba 递归
有 warmup:
只需计算 15 个 token 的 Mamba 递归(节省 30 token 的串行计算)

4. ... (后续 decode + cache_finished_req 同理)

6.4 复用限制:长无法复用于短

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
假设 tree 中只有:
root → ["ABCDEFGH"](KV:[1..8], SSM:[slot 1]) (来自一个 8 token request)

新 request: "ABCD" (4 tokens)

match_prefix:
while 循环进入 "ABCDEFGH" 节点
prefix_len = 4 < len(node.key)=8
→ _split_node:
root → ["ABCD"](KV:[1..4], mamba_value=None!) ← 中间节点,没有 SSM!
└── ["EFGH"](KV:[5..8], SSM:[slot 1])

best_last_node = root(因为 "ABCD" 节点没有 mamba_value)
返回 prefix_indices = [] ← 无法复用!

即使 KV 匹配了 4 个 token,但没有对应的 SSM state
→ request 必须从头计算 Mamba(Full Attention 仍可复用 KV)

7. 关键设计总结

7.1 为什么 KV 可以共享但 SSM 需要 COW

KV Cache SSM State
读/写模式 Attention 只读历史 KV Forward 每个 token 都修改 state
能否共享 可以(只读) 不行(会被修改)
方案 直接共享 Pool index Copy-on-Write

7.2 Radix Tree 存指针不存数据

1
2
3
4
5
Tree node.value = [10, 11, 12]     ← 3 个 int,指向 KV Pool
Tree node.mamba_value = [5] ← 1 个 int,指向 SSM Pool

实际数据始终在 Pool 的 GPU 显存里,不动。
insert/evict 只是转移"谁拥有这个 slot"的所有权。

7.3 Lock 引用计数

1
2
3
4
5
6
inc_lock_ref: request 开始使用 tree 节点 → lock_ref++,节点不可淘汰
dec_lock_ref: request 释放 tree 节点 → lock_ref--
lock_ref == 0: 节点可被 LRU 淘汰

不变量: 如果 mamba_lock_ref > 0,则 full_lock_ref 必定 > 0
(SSM 依赖 KV,有 SSM 引用就必须保护 KV)

7.4 Warmup 的价值

1
2
3
4
5
不用 warmup:  第一个请求必须从头计算 Mamba(串行,慢)
用了 warmup: 第一个请求就能 COW 缓存的 SSM state(快)

最佳实践: warmup 最短的公共前缀(system prompt)
所有以此开头的请求都能命中

8. KDA Speculative Decoding 机制

8.1 整体架构

KDA 模型的 speculative decoding 使用 FROZEN_KV_MTP 算法:

1
2
3
4
5
6
7
8
9
10
11
12
Target Model (完整 KDA LLM):
Layer 0: Full Attention → 产生 KV Cache
Layer 1: Mamba (SSM) → 产生 SSM State
Layer 2: Full Attention
Layer 3: Mamba
...
LM Head → logits

Draft Model (轻量级 MTP 头):
共享 target 的 embedding 和 lm_head
几层 Transformer block(注意力只读 target 已有的 KV)
不含 Mamba 层,不自己维护 KV

8.2 Frozen KV 设计

“Frozen”(冻结)指 draft 推理全程中,target 的 KV 序列不增长、不变化:

  • Draft 模型不写 KV Pool
  • Draft 每步的 RoPE position 固定为 seq_lens - 1
  • 通过上下文管理器临时将 attention backend 的 KV pool 切换为 target 的(只读)
1
2
3
4
5
6
7
8
9
# frozen_kv_mtp_utils.py
@contextmanager
def target_kv_pool_view(forward_batch, kv_context, draft_attn_backend):
saved_pool = draft_attn_backend.token_to_kv_pool
draft_attn_backend.token_to_kv_pool = kv_context.target_token_to_kv_pool
try:
yield
finally:
draft_attn_backend.token_to_kv_pool = saved_pool

8.3 Verify 阶段的 Mamba 处理

Verify 时 target 模型对 EAGLE tree 所有 draft tokens 做一次 batched forward。 所有 draft tokens 从 layer 0 到 layer 31 并行走一遍(层间串行,层内 batch 并行)。

对 Mamba 层的特殊处理:

  • disable_state_update=True:不直接修改主 SSM state
  • 每个 draft position 的递归结果写入 intermediate_states_buffer
  • 通过 retrieve_parent_token 数组实现 EAGLE tree 的 parent→child SSM 继承
1
2
3
4
5
6
7
8
9
# mamba.py:651 (verify 路径)
selective_state_update(
ssm_state,
...,
disable_state_update=True,
intermediate_states_buffer=layer_cache.intermediate_ssm,
cache_steps=draft_token_num,
retrieve_parent_token=metadata.retrieve_parent_token,
)

8.4 Intermediate Buffer 布局与回滚

1
2
3
4
5
6
7
8
# memory_pool.py:301 — 额外分配的显存
intermediate_ssm_state_cache = torch.zeros(
num_mamba_layers, # 32
spec_state_size + 1, # max_running_requests + 1
draft_token_num, # 63 (topk=4,steps=3 的树节点数)
nheads * head_dim,
ssm_state_size,
)

Verify 后通过 fused_mamba_state_scatter_with_mask O(1) 回滚:

1
2
3
4
5
6
7
8
# eagle_worker.py:1087
fused_mamba_state_scatter_with_mask(
ssm_states, # 主 SSM state pool
intermediate_state_cache, # intermediate buffer
state_indices_tensor, # request → pool slot
last_correct_step_indices, # 每个 req 最后正确的 step
)
# 效果: main_state[layer, slot] = buffer[layer, req, accepted_step]

8.5 当前方案的内存开销

Intermediate buffer 是主 SSM state 的 draft_token_num 倍:

topk steps draft_token_num 额外内存相对主 state
1 3 3
2 3 7
4 3 63 63×

每层的 intermediate buffer 独立保留(不可释放),因为 verify 判定后要从所有层的 buffer 中 scatter 正确 step 的值。

8.6 与 RadixCache 的集成

Spec decode 全程中 Radix Tree 不被修改。所有操作在 Pool 级别闭环:

  1. Verify 后:scatter 正确 SSM state 回主 pool
  2. 如果跨越 track interval(256):scatter 到 ping-pong buffer
  3. Request 结束:cache_finished_req 从 ping-pong buffer 取 SSM 存入 tree
1
2
3
4
Verify → scatter 正确 step → 主 state 更新
→ 如果跨 256 → ping-pong buffer 也更新

Request 结束 → cache_finished_req → ping-pong → radix tree 节点

9. 延迟更新方案(优化 Intermediate Buffer 内存)

9.1 问题动机

当前方案的 intermediate buffer 常驻显存:[32 layers × N × T × state_size]。 当 topk=4, T=63 时,额外显存可达数 GB,非常昂贵。根本原因是 verify 后需要从每一层的 buffer 中 scatter,所以所有层的 buffer 都不能释放。

9.2 方案核心思想

Verify 时不更新主 SSM state,verify 结束后也不立即更新,而是等到下一轮开始时,用被接受 tokens 的"输入激活值"把 SSM 重新递推到正确状态。

1
2
3
4
5
6
7
8
9
10
11
Round t:
1. Replay: 用上一轮保存的激活值,S(t-1) → S(t)
2. Draft: 从 S(t) 开始 seed
3. Verify forward: Mamba 层临时算各节点 S(用 1 层复用 buffer 得到 logits)
4. Verify 判定: 确定接受 [A, B, C]
5. 保存被接受 tokens 在每层的输入激活值
6. 主 state 仍为 S(t)(不更新)

Round t+1:
1. Replay: 用 Round t 保存的激活值,S(t) → S(t+1)
2. ...

9.3 什么是"输入激活值"

SSM 的递归公式为 S_new = A × S_old + B × x,其中:

  • A — 模型参数(永远在那里)
  • S_old — 主 pool 里的当前值(永远在那里)
  • B, x, dt — 线性投影算出的中间值(只在 forward 瞬间存在)
1
输入激活 = (hidden_states_d, dt, B, C)  ← Mamba 层做 SSM 递归需要的原料

Replay 时:只要有激活值 + 模型参数 + 当前 S_old,就能重算出 S_new:

1
2
3
4
# 下一轮开始时
for layer in range(32):
for token in accepted_tokens:
main_state[layer] = A[layer] * main_state[layer] + B_saved[layer,token] * x_saved[layer,token]

9.4 为什么只保留 accepted 的激活值就够了

时序上,verify 判定在保存之前

1
2
3
1. Verify forward → 得到 63 个 logits
2. 贪心判定 → 知道只有 [d0, d1, d3] 被接受
3. ★ 这时才保存激活值 ★ → 只保存 d0, d1, d3 的

从语义上,被拒绝的 tokens 等于从未发生过。request 的真实序列永远是一条线性链——只有 accepted 路径上的 tokens 对 S(t) 有贡献:

1
2
3
S(t-1) → process(d0) → S_a
S_a → process(d1) → S_b
S_b → process(d3) → S_c = S(t)

被拒绝的 tokens 的激活值对 S(t) 没有任何贡献,无需保存。

9.5 Verify 期间的 1 层复用 Buffer

虽然不存全量快照,verify forward 中 Mamba 层仍需临时计算 tree 各节点的 S(因为 output = C × S_new 需要 S 才能算出 logits)。

但这些临时 S 只在单层计算期间有用:

1
2
3
4
5
模型前向是逐层串行执行的:
Layer 0 → Layer 1 → Layer 2 → ... → Layer 31

每一层的输入是上一层的 output(hidden_states),不是上一层的 SSM state。
Layer 1 从自己的主 state 出发计算,不需要 Layer 0 的临时 S。

所以可以只分配 1 层大小的 buffer,逐层覆盖复用:

1
2
3
4
5
6
7
8
9
10
11
Mamba Layer 0 执行:
buffer 存 Layer 0 各节点的临时 S
所有节点算完 → 得到 output_0 → 传给下一层
★ buffer 中 Layer 0 的值此后无人再需要 ★

Mamba Layer 1 执行:
覆盖同一块 buffer → 存 Layer 1 各节点的临时 S
得到 output_1 → 传给下一层
★ buffer 又没用了 ★

...Layer 31 同理...

当前方案不能这样做:因为 verify 结束后需要从每层的 buffer 中 scatter,如果 Layer 0 的 buffer 被 Layer 1 覆盖了就取不回来了。延迟方案不做 scatter(改为 replay),所以可以覆盖。

9.6 EAGLE Tree 内部 parent→child 的 S 共享

单层 buffer 内部,tree 节点的依赖通过 retrieve_parent_token 数组解决:

1
2
3
4
5
6
7
8
9
10
11
12
Buffer layout (单层, 单 request): [T slots, state_size]

EAGLE tree: Buffer slot:
bonus (从主 state 读)
/ | \
d0 d1 d2 buffer[0], buffer[1], buffer[2]
/|\ /|\
d3 d4 d5 d6 d7 d8 buffer[3], buffer[4], ..., buffer[8]

retrieve_parent_token = [-1, -1, -1, 0, 0, 0, 1, 1, 1]
└ step 0 ┘ └── step 1 ──────┘
parent=主state parent=buffer[0/1]

Kernel 内部按 step 顺序处理:

1
2
3
Step 0: d0, d1, d2 从主 state 读 → 结果写入 buffer[0,1,2]
Step 1: d3 读 buffer[0](=d0 的 S), d6 读 buffer[1](=d1 的 S)
→ 结果写入 buffer[3..8]

9.7 内存对比

设 L=32 layers, N=32 requests, T=63 (topk=4), accepted=4, bf16

大小对比

1
2
3
SSM state 单份:    nheads × head_dim × dstate ≈ 500K elements
输入激活值单份: hidden_d + dt + B + C ≈ 10K elements
比值: 约 50× 更小

总内存对比

当前方案 延迟方案
常驻内存(verify 后) L×N×T×state = 32×32×63×500K×2B ≈ 数 GB L×accepted×activation = 32×4×10K×2B ≈ 2.5 MB
峰值内存(verify 期间) 同上(不释放) 1×N×T×state ≈ 当前的 1/32
缩减比(常驻) 基准 ~1000×
缩减比(峰值) 基准 ~32×

9.8 延迟代价

下一轮开始时的 replay 耗时:

1
2
3
4
5
6
replay = num_layers × num_accepted × 单次 SSM update kernel
= 32 × 4 × ~10μs
≈ 100-300 μs

对比 verify forward 本身: 5-20 ms
overhead: ~2%,可忽略

9.9 正确性保证

主 state 更新顺序

1
2
3
4
Round t-1 结束: 保存了 accepted 的激活值, S 仍为 S(t-2)
Round t 开始: replay S(t-2) → S(t-1) ✓ (此时 S 是正确的)
Round t draft: 从 S(t-1) 出发 seed ✓
Round t verify: 从 S(t-1) 出发计算各节点 logits ✓

前提条件:Draft 模型本身不含 Mamba 层(Frozen-KV MTP 满足,draft 只读 target KV)。

Accepted 路径线性链保证:Verify 贪心判定保证接受路径从根到叶连续——不可能"跳着接受"。所以逐个 replay accepted tokens 就能正确递推 S。

9.10 Track Interval 兼容

Replay 时顺便检查是否跨越 256 边界:

1
2
3
4
5
for i, token in enumerate(accepted_tokens):
main_state[layer] = A * main_state[layer] + B_saved[layer, i] * x_saved[layer, i]
current_seq_len += 1
if current_seq_len % 256 == 0:
ping_pong_buffer[track_slot] = main_state[layer].clone()

比当前方案(从 intermediate buffer scatter 特定 step)更直观。

9.11 完整时序对比

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
═══ 当前方案 ═══════════════════════════════════════════════════════

Round t:
[Draft] → [Verify forward: 63 节点, 32 层全部写入 intermediate buffer]
→ [Verify 判定: 接受 4 个]
→ [Scatter: 从 32 层 buffer 中取 step=k 写回主 state]
→ 主 state 立即 = S(t)

Round t+1:
[Draft: 从 S(t) 开始] → [Verify] → [Scatter] → S(t+1)

常驻内存: 32 × N × 63 × state_size (不释放)


═══ 延迟更新方案 ══════════════════════════════════════════════════

Round t:
[Replay: 用 Round t-1 的激活值, S(t-2)→S(t-1)]
[Draft: 从 S(t-1) 开始]
[Verify forward: 63 节点, 用 1 层复用 buffer 临时算 logits]
[Verify 判定: 接受 4 个]
[保存 4 个 token × 32 层的输入激活值]
主 state 仍为 S(t-1)

Round t+1:
[Replay: 用 Round t 的激活值, S(t-1)→S(t)]
[Draft: 从 S(t) 开始] → ...

常驻内存: 32 × 4 × activation_size (~MB 级)
峰值内存: 1 × N × 63 × state_size (verify 期间临时, 可释放)

9.12 方案优缺点总结

维度 当前方案(全量快照) 延迟更新方案
峰值内存 32 × N × T × state 1 × N × T × state(32× 缩减
常驻内存 同峰值 32 × accepted × activation(~1000× 缩减
Verify 后延迟 0(O(1) scatter) +100-300 μs(replay)
实现复杂度 现有 需改动:1) 缓存激活逻辑 2) buffer 复用 3) replay kernel
正确性 ✓(replay 后再 draft,顺序正确)
适用条件 通用 Draft 无 Mamba 层(Frozen-KV MTP 满足)
Track 兼容 需要从 buffer scatter 特定 step replay 过程中自然检查,更简洁

9.13 适用场景选择

  • topk=1, T=3:当前方案只需 3× 主 state,延迟方案收益有限
  • topk=2, T=7:延迟方案可节省约 90% intermediate buffer
  • topk=4, T=63:延迟方案价值最大,从 ~GB 级降到 ~MB 级
  • 内存紧张 + 需要大 batch:延迟方案释放的显存可换来更多并发 request
  • 标题: SGLang之MambaRadixCache机制
  • 作者: 鱿鱼圈
  • 创建于 : 2026-05-29 23:50:00
  • 更新于 : 2026-06-14 22:47:27
  • 链接: https://yuyanqi.com/2026/05/29/SGLang之MambaRadixCache机制解析/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论
目录
SGLang之MambaRadixCache机制