sparkle/mlx_streaming/mtp/drafter.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

189 lines
9.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""真实 MTP drafter把 Qwen3NextMTP 包成 mtp_generate 需要的 draft/sync 接口。"""
import mlx.core as mx
from mlx_streaming.mtp.qwen3_next_mtp import mtp_step
from mlx_streaming.mtp.kv_cache import _snapshot, _restore
class MTPDrafter:
"""把 Qwen3NextMTP 包成 mtp_generate 需要的 drafter 接口。"""
def __init__(self, mtp, lm_head):
self.mtp = mtp
self.lm_head = lm_head
self.embed_tokens = mtp.embed_tokens
def make_cache(self):
from mlx_lm.models import cache as kc
return [kc.KVCache()] # MTP 单层全注意力
def draft(self, H_last, x_ids, mtp_cache, K, topk: int = 0):
# 注意:每草稿一次 int(argmax) host 同步反而最快 —— 它让 MLX 把每步 draft 的
# 图保持很小、及时释放;试过"全程 GPU argmax、末尾一次同步"攒大惰性图,draft
# 慢 5×、端到端 24.9→16.3(A/B 证伪,见报告)。保持逐步同步。
# topk>0(探针用):额外返回每个位置 MTP 的 top-k 候选 id(降序),量树形展开的救回上界。
drafts, cands = [], []
h, cur = H_last, x_ids
for _ in range(K):
logits, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
lg = logits[0].reshape(-1)
if topk > 0:
order = [int(i) for i in mx.argsort(lg)[-topk:].tolist()][::-1] # 降序 top-k
d = order[0]
cands.append(order)
else:
d = int(mx.argmax(lg))
drafts.append(d)
h, cur = mh, mx.array([[d]])
if topk > 0:
return drafts, cands
return drafts
def draft_adaptive(self, H_last, x_ids, mtp_cache, depth_max, tau):
"""置信度门控动态深度:逐位贪婪抽,累计置信度 C=∏p_i 跌破 tau 即停,最多 depth_max 位。
返回可变长度(1..depth_max)的草稿链。C 用每位 top-1 softmax 概率连乘;低置信提前收敛
(本步只 verify 到当前深度,省后续位置的专家加载),高置信抽满到 depth_max。始终至少抽 1 位
(depth=1 即退化为普通单 token 解码:verify 只喂 [x]、恒接受模型真值 1 个)。
"""
drafts = []
h, cur = H_last, x_ids
conf = 1.0
for i in range(depth_max):
logits, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
lg = logits[0].reshape(-1)
d = int(mx.argmax(lg))
drafts.append(d)
h, cur = mh, mx.array([[d]])
if i + 1 < depth_max: # 末位无需再判是否加深
conf *= float(mx.softmax(lg)[d])
if conf < tau:
break
return drafts
def draft_adaptive_tree(self, H_last, x_ids, mtp_cache, depth_max, tau):
"""置信度门控变长 + pos0 分支:返回 (chainA, chainB)。
chainA 同 draft_adaptive(累计置信度门控,变长 n=1..depth_max);chainB 是第1 位 top-2 分支,
续抽到与 chainA 同深度 n(仅 n>=2 时,浅步 depth=1 无位置可救 → chainB=None)。用快照隔离
A/B 分叉,B 不带 A 的递归污染。供合并路径:动态定深 + 对深链做 pos0 救回。
"""
logits1, mh1 = mtp_step(self.mtp, H_last, x_ids, self.lm_head, mtp_cache[0])
lg = logits1[0].reshape(-1)
top2 = [int(i) for i in mx.argsort(lg)[-2:].tolist()][::-1] # [d1a, d1b] 降序
d1a, d1b = top2[0], top2[1]
snap_after_x = _snapshot(mtp_cache)
# chainA:从 d1a 续抽,累计置信度跌破 tau 即停(与 draft_adaptive 同语义)。
chainA = [d1a]
conf = float(mx.softmax(lg)[d1a])
h, cur = mh1, mx.array([[d1a]])
while len(chainA) < depth_max and conf >= tau:
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
lg2 = lo[0].reshape(-1)
d = int(mx.argmax(lg2))
chainA.append(d)
conf *= float(mx.softmax(lg2)[d])
h, cur = mh, mx.array([[d]])
n = len(chainA)
chainB = None
if n >= 2: # 浅步无后续位置可救,不抽 B
_restore(mtp_cache, snap_after_x)
chainB = [d1b]
h, cur = mh1, mx.array([[d1b]])
for _ in range(n - 1):
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
d = int(mx.argmax(lo[0]))
chainB.append(d)
h, cur = mh, mx.array([[d]])
return chainA, chainB
def draft_tree(self, H_last, x_ids, mtp_cache, K, pos1=True):
"""最小树:第1(始终)、第2(pos1=True 时)草稿位置展开 top-2,返回三条链,各长 K。
- chainA = [d1a(top-1), d2a(top-1), d3a] 全 top-1,同 draft。
- chainB = [d1b(top-2), d2b, d3b] 第1 草稿位置 top-2 分支 → pos0 救回。
- chainC = [d1a, d2c(第2 位次选), d3c] 或 None 第2 草稿位置 top-2 分支 → pos1 救回。
各分叉点用 mtp_cache 快照隔离,保证任一分支不带其它分支的递归污染:snap_after_x(x 处理后)
分叉 d1a/d1b;snap_after_d1a(d1a 处理后)分叉 d2a/d2c。探针实测第2 位「首选错次选对」比例
(~11%)高于第1 位(~7%),故 pos1 分支值得抽;但每步多一次 MTP 续抽,pos1=False 时退化为
仅 chainA/chainB(与旧最小树成本一致),便于 A/B 隔离 pos1 增量。
"""
logits1, mh1 = mtp_step(self.mtp, H_last, x_ids, self.lm_head, mtp_cache[0])
lg = logits1[0].reshape(-1)
top2 = [int(i) for i in mx.argsort(lg)[-2:].tolist()][::-1] # [d1a, d1b] 降序
d1a, d1b = top2[0], top2[1]
snap_after_x = _snapshot(mtp_cache) # x 处理后、第1 位分叉前的 MTP 递归态
def _continue(first, h0, n):
"""首 token first 已定,从隐状态 h0 再贪婪续抽 n 个,返回长度 1+n 的链。"""
chain = [first]
h, cur = h0, mx.array([[first]])
for _ in range(n):
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
d = int(mx.argmax(lo[0]))
chain.append(d)
h, cur = mh, mx.array([[d]])
return chain
chainC = None
if pos1:
# 从 d1a 续抽第2 位,捕获第2 位 top-2(d2a/d2c),供 chainA 与 chainC 分叉。
lo2, mh_after_d1a = mtp_step(self.mtp, mh1, mx.array([[d1a]]), self.lm_head, mtp_cache[0])
lg2 = lo2[0].reshape(-1)
top2_p1 = [int(i) for i in mx.argsort(lg2)[-2:].tolist()][::-1] # [d2a, d2c] 降序
d2a, d2c = top2_p1[0], top2_p1[1]
snap_after_d1a = _snapshot(mtp_cache) # d1a 处理后、第2 位分叉前的递归态
chainA = [d1a] + _continue(d2a, mh_after_d1a, K - 2) # [d1a, d2a, d3a...]
_restore(mtp_cache, snap_after_d1a) # 回到 d1a 态,C 链从第2 位次选分叉
chainC = [d1a] + _continue(d2c, mh_after_d1a, K - 2) # [d1a, d2c, d3c...]
else:
chainA = _continue(d1a, mh1, K - 1) # 仅全 top-1 链(旧最小树成本)
# chainB:回到 x 态,从第1 位次选 d1b 分叉续抽。
_restore(mtp_cache, snap_after_x)
chainB = _continue(d1b, mh1, K - 1) # [d1b, d2b, d3b...]
return chainA, chainB, chainC
def draft_paths(self, H_last, x_ids, mtp_cache, K, P):
"""完整树 batch-of-paths:位置1 展开 top-P,返回 P 条链(各长 K)。
位置1 取 MTP top-P 候选 [d1_0..d1_{P-1}];每个候选从 pos1 的共享 MTP 递归态 mh1 分叉,
贪婪续抽 K-1 个 token 成一条链。用 mtp_cache 快照保证各链从同一起点分叉、互不污染。
P=1 退化为普通链;P=2 等价 draft_tree。
"""
logits1, mh1 = mtp_step(self.mtp, H_last, x_ids, self.lm_head, mtp_cache[0])
lg = logits1[0].reshape(-1)
firsts = [int(i) for i in mx.argsort(lg)[-P:].tolist()][::-1] # top-P 降序
snap_pos1 = _snapshot(mtp_cache)
def _continue(first):
chain = [first]
h, cur = mh1, mx.array([[first]])
for _ in range(K - 1):
lo, mh = mtp_step(self.mtp, h, cur, self.lm_head, mtp_cache[0])
d = int(mx.argmax(lo[0]))
chain.append(d)
h, cur = mh, mx.array([[d]])
return chain
paths = []
for j, f in enumerate(firsts):
if j > 0:
_restore(mtp_cache, snap_pos1) # 回到 pos1 态,从同一起点分叉
paths.append(_continue(f))
return paths
def sync(self, prev_H, rH, replay_in, mtp_cache):
"""用已接受 token 的真实主模型 hidden 推进 MTP KV cache。
MTP 在位置 i 消费 (H_i, t_{i+1}) 预测 t_{i+2};因此提交 accepted prefix
`[t_{i+1}, ..., t_{i+n}]` 时,hidden 序列应为 `[H_i, ..., H_{i+n-1}]`。
"""
from mlx_streaming.mtp.qwen3_next_mtp import mtp_advance
h_seq = mx.concatenate([prev_H, rH[:, :-1, :]], axis=1)
H = mtp_advance(self.mtp, h_seq, replay_in, mtp_cache[0])
mx.eval(H)