LLaMA-13B 解码一 Token 约 1.2ms。一秒钟只能出 800 个字——不够快。为什么不能一次出多个 Token?因为解码是自回归的——下一个 Token 依赖上一个的输出。

推测解码(Speculative Decoding)打破了自回归瓶颈。它用一个超小的 Draft Model 提前"猜"一串 Token,然后用大模型并行验证。猜对了就一次吐出多个 Token。


为什么大模型输出速度慢

# 标准 Decode——逐 Token 自回归
def standard_decode(model, prompt, max_tokens=128):
    output = prompt[:]
    for step in range(max_tokens):
        # 每次只能算 1 个 Token
        logits = model.infer(output)  # 1 次推理 → 1 个 Token
        next_token = sample(logits[-1])
        output.append(next_token)
        
        if next_token == EOS:
            break
    return output
# 128 个 Token = 128 次推理 → 128 × 1.2ms = 154ms

自回归限制:第 5 个 Token 的推理必须等第 1-4 个算完。NPU 的 24 TFLOPS 算力在解码阶段只用了 35%。

Draft Model 如何提前生成 Token

# 推测解码——Draft Model 猜一串,大模型批量验证

def speculative_decode(large_model, draft_model, prompt, gamma=5):
    """
    large_model: LLaMA-13B(目标模型)
    draft_model: 一个小模型(如 LLaMA-68M,只有目标模型的 0.5% 参数)
    gamma: 每次猜测几 Token
    """
    output = prompt[:]
    
    while len(output) < max_tokens:
        # Step 1: Draft Model 快速猜 gamma 个 Token
        draft_tokens = []
        draft_state = draft_model.init_state(output)
        
        for i in range(gamma):
            logits = draft_model.infer(draft_state)  # 0.01ms(68M 模型)
            next_tok = sample(logits[-1])
            draft_tokens.append(next_tok)
            draft_state = draft_model.append_state(draft_state, next_tok)
        
        # Step 2: 大模型并行验证所有 Draft Token
        candidate_seq = output + draft_tokens
        
        # 一次推理算完所有 candidate 的 logits!
        large_logits = large_model.infer(candidate_seq)
        # 不是 gamma 次推理——只推理 1 次
        
        # Step 3: 逐个验证——保留通过验证的前 k 个 Token
        accepted = 0
        for i in range(gamma):
            # 大模型的 predicted 分布 vs Draft Model 的分布
            large_probs = softmax(large_logits[len(output) + i - 1])
            draft_probs = softmax(draft_logits[i])
            
            if draft_tokens[i] == sample(large_probs):
                # Draft 猜对了 → 接受
                accepted += 1
            else:
                # Draft 猜错了 → 从大模型的正确分布重新采样
                output.append(sample(large_probs))
                break
        
        # 如果全部猜对——一次验证拿到了 gamma+1 个 Token
        if accepted == gamma:
            output.extend(draft_tokens)
            output.append(sample(large_logits[-1]))
    
    return output
# 如果 gamma=5、接受率 80%——每步平均产出 4 Token
# 128 Token 从 128 次推理降到 32 次 → 延迟从 154ms 降到 38ms

Draft Model 的关键:它算得快但不是必须准——猜错了只是浪费一次验证,不会降低输出质量。大模型的验证保证最终输出跟不用推测解码完全一致。

昇腾NPU如何并行验证

// CANN Runtime 上的并行验证——Draft Token 在同一个 Batch 中推理
// draft_tokens = [a, b, c, d, e]

// 不推测时:5 次解码推理
for (int i = 0; i < 5; i++) {
    aclmdlExecute(model, single_token_input, output);
}

// 推测解码:拼接 draft_tokens 做一次推理
// 输入序列 = prompt + [a, b, c, d, e]
// GE 会执行 Prefill 路径(输入多 Token)——比 Decode 路径快 N 倍
int16_t candidate[] = {45, 892, 312, 67, 1289, 34, a, b, c, d, e};
aclmdlExecute(model, candidate_input, output);
// 一次推理产出所有 candidate 位置的 logits
// 时间 = 一次 Prefill ≈ 多次 Decode

// 验证在 CPU 上做——跟推理异步
for (int i = 0; i < gamma; i++) {
    if (reject(output[i], draft_logits[i])) {
        output[output_len] = resample(output[i]);
        break;
    }
    output[output_len + i] = draft_tokens[i];
}

关键点:大模型用一次 Prefill 推理验证了 gamma 个 Token。Prefill 的 Cube 利用率(80%)远高于 Decode(35%)——GPU 算力的浪费减少了。

性能提升

# 推测解码的实测性能(LLaMA-13B, Ascend 910)
results = {
    "gamma_1":  {"tokens/s": 800,  "accept_rate": 1.0},  # 等价于标准解码
    "gamma_3":  {"tokens/s": 1800, "accept_rate": 0.75}, # 猜 3 个接受 2.25 个
    "gamma_5":  {"tokens/s": 2400, "accept_rate": 0.65}, # 猜 5 个接受 3.25 个
    "gamma_8":  {"tokens/s": 2800, "accept_rate": 0.55}, # 接受率下降后收益饱和
}
# 最佳 gamma≈5:速度提升约 3 倍

Draft Model 的参数量=目标模型的 0.5%-1%。太大就没有速度优势。太小接受率太低。LLaMA-13B 的常用 Draft Model 是 LLaMA-68M——0.5% 的参数量,推理快了 100 倍,接受率约 65%。

// Draft Model 的 CANN 部署——跟大模型共享 Runtime 实例
// 两个模型用一个 Device、不同 Context

aclrtContext draft_ctx, main_ctx;
aclrtCreateContext(&draft_ctx, 0);
aclrtCreateContext(&main_ctx, 0);

uint32_t draft_model_id, main_model_id;
aclmdlLoadFromFile("draft_68m.om", &draft_model_id);
aclmdlLoadFromFile("llama_13b.om", &main_model_id);

// 交替推理:Draft → 大模型 → Draft → 大模型
for (int step = 0; step < gamma; step++) {
    aclrtSetContext(draft_ctx);
    aclmdlExecute(draft_model_id, draft_input, draft_output);  // 0.01ms
}

aclrtSetContext(main_ctx);
aclmdlExecute(main_model_id, verify_input, verify_output);  // 1 次替代 5 次

参考仓库

Multi-Stream 并行调度

FlashDecode 优化

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐