【Bug已解决】[Bug]: MTP speculative decoding crash with illegal memory access on long sequences (Qwen3.6-27B-FP8, v0.19.1) 解决方案
一、现象长什么样
在Qwen3.6-27B-FP8+ vLLM 0.19.1 上开MTP(multi-token prediction,多 token 预测)投机解码,短序列一切正常,但一旦序列接近 KV 缓存上限(长文档、长对话),forward 中途进程崩溃,报非法内存访问:
torch.OutOfMemoryError: CUDA out of memory (有时) RuntimeError: CUDA error: an illegal memory access was encountered或者更明确指向投机解码层:
illegal memory access at kernel mtp_draft_forward: slot index 16384 >= num_kv_slots 16384几个特征:
- 短序列(比如 2k token 以内)完全正常;长序列(接近
--max-model-len)必崩。 - 只在开 MTP 时崩;关掉 MTP(
speculative_config=None)长序列也能跑。 - 崩的位置是 MTP 的 draft 前向,不是主模型前向。
- 报错有时是
illegal memory access,有时是out of memory,本质是同一件事:MTP 往「超出已分配 KV 槽位」的地方写了数据。
本质:MTP 一次草拟k个 token,会把当前位置往后推k格去写 KV 缓存;当序列已经很长、剩余 KV 槽位不足k个时,草拟位置越界,kernel 写到了未分配显存 → 非法内存访问。
二、背景
MTP 的做法是:主模型算出下一个 token 后,MTP 头基于「已生成的序列 + 刚算出的 token」再一次性草拟出接下来 k 个 token(比如 k=3),然后主模型并行验证这 k 个。为了草拟第k个 token,MTP 前向需要把位置pos, pos+1, ..., pos+k-1的 KV 都写进缓存、并读出来算下一层。
问题在于「写 KV 缓存」这一步和「KV 缓存容量」的耦合:
- KV 缓存是按
max_model_len预分配好固定槽位的(比如 16384 个 slot)。 - 普通解码每次只推进 1 个位置,永远不会越界(因为调度器保证序列长度 ≤ max_model_len)。
- 但 MTP 一次要推进
k个位置。调度器在计算「这个序列还能不能接着生成」时,往往只按「主模型 +1」来算剩余槽位,没把 MTP 要额外占的k-1个槽位算进去。于是当序列长度 =max_model_len - 2时,调度器认为「还能生成」,MTP 却要写pos到pos+2共 3 个槽位,最后一个槽位pos+2 = max_model_len已经越界 → kernel 写未分配显存 → 非法内存访问。
这和普通「序列超长」不同:普通情况调度器会拦下;但 MTP 把「一次占用的槽位数」从 1 变成了 k,调度器的边界判断没跟着改,漏洞就出现了。
三、根因
根因是MTP 草拟长度k没有被纳入 KV 槽位的边界核算,导致长序列末尾草拟位置越界,三层:
第一层(主因):调度器的「剩余槽位」判断没加 MTP 的k余量。调度器决定「这个序列还能不能生成下一个 token」时,检查的是seq_len + 1 <= max_model_len。但 MTP 实际上需要seq_len + k <= max_model_len。差了k-1个槽位,序列在max_model_len - k < seq_len <= max_model_len - 1这段区间里,调度器放行、MTP 越界。
第二层:MTP draft 前向没有对 slot 做边界断言。draft kernel 拿到pos和k后直接kv_cache[pos + i] = ...,没有任何if pos + i >= num_slots: 截断/报错的防护。它假设「调用方保证槽位够」,但调用方(调度器)的保证是错的,于是越界写直接发生。
第三层:错误表现不稳定(IMA vs OOM)。越界写的后果取决于「越界到哪」:若越界到同一块已分配显存的邻近区域,可能只是静默污染(偶尔还能跑完但结果错);若越界到未映射显存,就是illegal memory access;若越界触发了一次额外的显存分配,就是out of memory。同一个根因,三种表象,增加排查难度。
一句话:MTP 的草拟长度没被调度器算进 KV 边界,长序列末尾草拟越界写未分配显存,表现为非法内存访问(或偶发 OOM)。
四、最小可运行复现
下面用纯 Python 模拟「MTP 草拟 k 个 token,但调度器只按 +1 判断边界,长序列末尾越界」的控制流,不需要 GPU:
class KVCache: def __init__(self, num_slots): self.slots = [None] * num_slots self.num = num_slots def write(self, pos, k, value): # MTP draft 前向:写 pos .. pos+k-1 for i in range(k): idx = pos + i if idx >= self.num: # 原版没有这个检查,直接越界 raise IndexError(f"slot {idx} >= num_slots {self.num}") self.slots[idx] = value def can_generate(seq_len, max_len, k, speculative): # 调度器的边界判断 needed = seq_len + (k if speculative else 1) return needed <= max_len def main(): max_len, k = 16, 3 cache = KVCache(max_len) # 序列化到 seq_len = 14(max_len - 2) seq_len = max_len - 2 speculative = True # 调度器认为:14 + 1 = 15 <= 16,放行 print("调度器放行:", can_generate(seq_len, max_len, k, speculative)) # 但 MTP 要写 14,15,16 -> 16 越界 try: cache.write(seq_len, k, "draft_token") print("写成功(实际会越界)") except IndexError as e: print("复现成功:", e) if __name__ == "__main__": main()跑出来会打印调度器放行: True然后复现成功: slot 16 >= num_slots 16——调度器以为能生成、MTP 却越界,和线上「长序列末尾崩溃」完全一致。
五、解决方案(第一层:最小直接修复)
最省事的救火:关掉 MTP(退回普通解码),长序列立刻能跑。代价是吞吐下降(失去投机加速):
llm = LLM( model="Qwen3.6-27B-FP8", # speculative_config=None # 不启用 MTP )或者把--max-model-len调大一点,给 MTP 的k余量留出空间(代价是 KV 缓存显存变大):
vllm serve Qwen3.6-27B-FP8 \ --speculative-config '{"method":"mtp","num_speculative_tokens":3}' \ --max-model-len 16384 \ --gpu-memory-utilization 0.8 # 留出 KV 余量更精准的临时规避:限制 MTP 只在「剩余槽位充足」时启用,剩余不足k就退回单 token 解码。这是第一层的「带保护」版本:
def safe_num_draft(seq_len, max_len, k): # 剩余槽位不足以支撑 k 个草拟时,自动缩减到 1(普通解码) remaining = max_len - seq_len return min(k, max(1, remaining))六、解决方案(第二层:结构性改进)
第一层是「避开/手动留余量」,第二层是「让调度器和 MTP 用同一套边界规则」——核心是把 MTP 的k纳入「可生成判定」和「KV 槽位核算」的单一事实来源:
from dataclasses import dataclass @dataclass class SeqBounds: max_len: int num_speculative: int = 1 def can_generate(self, seq_len: int) -> bool: # 单一边界规则:主模型 + 全部草拟 token 都必须落在 max_len 内 needed = seq_len + self.num_speculative return needed <= self.max_len def draft_slots_ok(self, pos: int, k: int) -> bool: # MTP draft 前向的边界检查:pos .. pos+k-1 必须全部合法 return pos + k <= self.max_len def check_position(self, pos: int, k: int) -> None: assert self.draft_slots_ok(pos, k), ( f"MTP draft 越界: pos={pos} k={k} 需要槽位 {pos+k} " f"但 max_len={self.max_len}" ) def check_expert(self, pos: int, k: int) -> None: # 专家路由侧的同样检查(MoE 下 token 也要落到合法 slot) self.check_position(pos, k)调度器在决定是否继续生成时,统一调用can_generate,把num_speculative算进去;MTP draft kernel 入口先check_position(pos, k)再写 KV:
def mtp_draft_forward(kv_cache, pos, k, bounds: SeqBounds): bounds.check_position(pos, k) # 越界立刻报错,绝不写未分配显存 for i in range(k): kv_cache.write(pos + i, compute_token(pos + i))这样「边界规则」只有一份,调度器和 MTP 不可能再各算各的。
七、解决方案(第三层:断言 / CI 守护)
把「MTP 不越界」「调度器按 k 判断」「长序列末尾安全降级」固化成测试:
import pytest def test_draft_within_bounds_ok(): b = SeqBounds(max_len=16, num_speculative=3) b.check_position(10, 3) # 10..12 <= 16,应通过 def test_draft_at_boundary_raises(): b = SeqBounds(max_len=16, num_speculative=3) with pytest.raises(AssertionError): b.check_position(14, 3) # 14..16 越界 def test_scheduler_accounts_for_k(): b = SeqBounds(max_len=16, num_speculative=3) # seq_len=14 时,14+3=17>16,调度器应拒绝继续生成 assert b.can_generate(14) is False assert b.can_generate(13) is True # 13+3=16 <= 16 def test_long_seq_tail_safe_degrade(): # 长序列末尾,MTP 自动退化成单 token,不越界 b = SeqBounds(max_len=16, num_speculative=3) seq_len = 15 k = min(b.num_speculative, 16 - seq_len) # k=1 b.check_position(seq_len, k) # 15..15 合法 def test_no_ima_on_max_len(): # 端到端:在 max_len 处停止草拟,不应触发越界 b = SeqBounds(max_len=16, num_speculative=3) for seq_len in range(0, 16): k = 3 if b.can_generate(seq_len) else 0 if k: b.check_position(seq_len, k) assert True再加一个端到端回归:长序列 + MTP 跑到max_model_len不崩:
def test_mtp_long_sequence_no_ima(): engine = make_engine(model="Qwen3.6-27B-FP8", speculative={"method": "mtp", "num_speculative_tokens": 3}, max_model_len=16384) out = engine.generate("超长文档..." * 500, max_tokens=16384) assert out is not None # 不应 illegal memory access八、排查清单
- 看报错是否
illegal memory access/out of memory且栈指向 MTP draft 前向 → 坐实本问题。 - 短序列能跑、长序列崩,且只在开 MTP 时崩 → 基本是 MTP 越界。
- 临时救火:关 MTP,或调大
--max-model-len,或在长序列末尾手动降num_speculative。 - 检查调度器「剩余槽位」判断是否包含 MTP 的
k余量(最常见疏漏)。 - 长期修复:边界规则单一化(调度器与 MTP 共用
SeqBounds),draft 前向前做check_position。 - 升级 vLLM 到合了 MTP 边界修复的版本,并跑上面的长序列回归。
- 若用
CUDA_LAUNCH_BLOCKING=1+TORCH_USE_CUDA_DSA=1复现,能让越界错误定位到精确 kernel 行。
九、小结
MTP 长序列非法内存访问,不是 FP8 或 Qwen 的锅,而是MTP 一次草拟 k 个 token,但调度器的 KV 边界判断只按 +1 算,长序列末尾草拟位置越界写未分配显存。最小修复是关 MTP / 调大 max_model_len / 末尾降 k;结构性修复是把边界规则收敛成单一SeqBounds、draft 前向前做check_position;最后用 pytest 把「不越界」「调度器按 k 判断」「长序列安全降级」锁死。配合CUDA_LAUNCH_BLOCKING=1能快速定位越界 kernel。抓住「投机解码一次性占用的槽位数 ≠ 1」这条,所有 spec decode 的边界坑都能照此排查。