DeepSpec — Speculative Decoding Training/Eval Codebase (companion to DSpark)

code deepseek-ai-DeepSpec
speculative-decodingdraft-modeleagle3fsdp-trainingflex-attention

DeepSpec — 投机解码 draft model 训练/评测代码库 #

1. TL;DR #

DeepSpec 是训练+评测投机解码 draft model 的全栈 Python 代码库(~6.1k LOC, MIT)。同一套 Qwen3DSparkModel 通过 config 开关切出 DFlash(CE-only)与 DSpark(Markov 头 + L1 分布匹配 + confidence 头),另有独立的 Eagle3(TTT) 路径。三段流水线:数据准备 → FSDP 训练 → rejection-sampling 评测,附 Qwen3/Gemma4 官方 checkpoint。

2. 三问:痛点 / 方法 / 结果 #

Q1 痛点。 投机解码要一个能"猜对 target 下一批 token"的小 draft model,但把 draft 训得既准又快依赖大量工程:怎么从 target 拿监督信号、怎么高效训练一个 block 的多步预测、怎么在评测时既保证输出分布无偏又度量 draft 质量。已有实现(SpecForge/Eagle3、DFlash)各写一套训练/评测栈,算法之间难以公平对比。

Q2 方法。 本仓把三种 draft 算法收敛到统一的 BaseTrainer/BaseEvaluator 骨架上,用 Python config 文件选择 trainer_cls 与超参:

核心技术壁垒:把"draft 一个块内多步预测"从因果链改写成 anchor-block 稀疏注意力训练 —— 用一个手写 FlexAttention mask_mod 让每个 draft 块只看自己的上下文前缀和本块,并配合 block_keep_mask/eval_mask 处理有效 anchor 不足时的 dummy 填充。这块把训练监督密度(一序列 512 个块)与掩码正确性耦合在一起,是最难照抄的部分(详见 §7)。

Q3 结果。 仓库本身不产出 benchmark 数字(那属于 DSpark 论文),但提供可复现能力:一张覆盖 Eagle3/DFlash/DSpark × {Qwen3-4B/8B/14B, Gemma4-12B} 的官方 checkpoint 表,以及 9~11 个投机解码基准(gsm8k/math500/aime/humaneval/mbpp/livecodebench/mt-bench/alpaca/arena-hard-v2)的评测入口。

3. 架构 / 模块图 #

顶层:train.py(训练入口)、eval.py(评测入口,按 draft 配置的 architectures[0] 选 evaluator)、deepspec/ 主包、config/(算法×目标模型的 Python 配置)、scripts/(数据/训练/评测 shell)、eval_datasets/(JSONL 基准)。

flowchart TD subgraph entry [入口] T[train.py
spawn 1 worker/GPU] E[eval.py
按 architectures 选 evaluator] end subgraph cfg [config/] C[dspark/ dflash/ eagle3/
× qwen3_4b/8b/14b, gemma4_12b] end subgraph pkg [deepspec/] DATA[data/
CacheDataset + collators + parser] MOD[modeling/
dspark/ + eagle3/ × qwen3/ gemma4/] TR[trainer/
BaseTrainer + 各算法 trainer] EV[eval/
base_evaluator + dspark/ eagle3/] UT[utils/
config/optim/sampling/metrics/distributed] end subgraph sc [scripts/data/] D1[download_and_split] --> D2[generate_train_data
SGLang 重生成答案] --> D3[prepare_target_cache
~38TB mmap] end C --> T C --> E D3 --> DATA T --> TR TR --> MOD TR --> DATA TR --> UT E --> EV EV --> MOD EV --> DATA

依赖边界config/*.py 是用户主要旋钮(组装 model/train/logging/data 四个 dict + finalize_cfg);trainer_cls 与 evaluator 通过 config 字段 / draft 配置的类名字符串动态选取(EVALUATORS 字典)。BaseTrainer 提供 FSDP 包装、BF16 优化器、梯度累积、可恢复调度、suspend/checkpoint;各算法 trainer 只重写 _build_draft_modelrun_batch

关键文件树(去掉 .git/ 与二进制 PDF):


DeepSpec/
├── train.py / eval.py            # 两个入口
├── deepspec/
│   ├── data/                     # target-cache 数据集 + collators + chat 解析 + CUDA prefetch
│   ├── eval/                     # base_evaluator + dspark/ + eagle3/
│   ├── modeling/                 # dspark/ eagle3/ × qwen3/ gemma4/
│   ├── trainer/                  # BaseTrainer + dspark_/eagle3_ trainer + ckpt_manager
│   └── utils/                    # config/distributed/optim/sampling/metrics/io/suspend
├── config/{dflash,dspark,eagle3}/*.py
├── scripts/{data,train,eval}/
└── eval_datasets/                # JSONL 基准 + 转换脚本

4. 正确性与关键机制 #

无形式化作者证明 — 代码库以工程正确性契约代替定理。以下 6 项是可从代码验证的正确性关键点:

  1. 无偏输出(rejection sampling):验证时对每个提议 token 计算接受概率 $p_{accept}=\min(1, p_{target}/p_{draft})$,accept_mask = (rand < accept_prob),再取 cumprod 得到接受前缀;首个被拒位置用 sample_residual(target, draft) 残差重采样。这保证最终 token 分布等于 target 单独解码的分布——投机解码正确性的核心不变量。
  2. 梯度累积等价性loss = run_batch(batch) / gradient_accumulation_steps 并在非同步 micro-step 上用 model.no_sync(),使 N 个 micro-batch 的梯度和等于一个大 batch;_compute_gradient_accumulation_steps 断言 global_batch_size % (world_size*local_batch_size) == 0,保证整除。
  3. anchor 掩码正确性dspark_mask_mod 里 draft 块只能看到 kv_idx < anchor_pos 的上下文与"同块"draft key,且被 block_keep_mask[b, q_block_id] 门控——dummy anchor 因 keep_mask=False 全部屏蔽,不污染 loss。
  4. train/eval 层对齐守卫assert_no_final_target_layer 拒绝把 target 最后一层放进 target_layer_ids,因为 HF output_hidden_states 在末层存的是 归一化后 隐状态,而 cache 存的是 raw decoder 输出——不守卫会静默错配。
  5. 接受率度量定义:per-position accept rate = $1 - \tfrac{1}{2}\lVert p_{draft} - p_{target}\rVert_1$,clamp 到 $[0,1]$;这是 L1 损失与 confidence 头 BCE 目标共用的量,训练与评测度量一致。
  6. 损失归约的分母一致性compute_dspark_loss 的 CE/L1/BCE 分母都跨 rank all-reduce,配合位置衰减 $\exp(-pos/\gamma)$;Eagle3 侧显式断言禁用 valid_token_mean(注释指出它会降 eval accept-length),只允许 local_mean
  7. DSpark 训练损失(config 驱动的加权和):

    $$L = \alpha_{ce}\cdot \mathrm{CE} + \alpha_{l1}\cdot \mathrm{L1}(p_{draft},p_{target}) + \alpha_{conf}\cdot \mathrm{BCE}(\mathrm{confidence}, \mathrm{acceptrate})$$

    DFlash 即取 $\alpha_{ce}=1,\ \alpha_{l1}=0,\ \alpha_{conf}=0$ 且 markov_rank=0

    5. 评测与可复现能力 #

    仓库提供的是"评测能力"而非论文数字:

    • 评测入口 eval.py:必填 --target_name_or_path / --draft_name_or_path,可调 --max-new-tokens=2048 / --temperature=1.0 / --confidence-threshold=0.0(该阈值为 0 时改为收集 confidence 校准指标,>0 时用于早停)。按 draft 配置 architectures[0]EVALUATORSQwen3/Gemma4 × DSpark/Eagle3 四类 evaluator。
    • 基准集eval.pyTASKS 列出 gsm8k(500)/math500(500)/aime25(30)/humaneval(164)/mbpp(256)/livecodebench(500)/mt-bench(80)/alpaca(500)/arena-hard-v2(500);eval_datasets/ 内含 JSONL 及 aime24、lbpp、swe-bench 等更多集合与转换脚本。
    • 报告指标build_results_table 输出 accept_len / verify_rate / accept_rate@pos;DSpark evaluator 的 ConfidenceHeadRecorder + PerPositionConfidenceMetrics 记录逐位置 ECE/AUROC/Brier 与 reliability diagram。
    • 官方 checkpoint 复现表(README,对应论文 Table 1):
    AlgorithmQwen3-4BQwen3-8BQwen3-14BGemma4-12B
    Eagle3deepseek-ai/eagle3_qwen3_4b_ttt7..._8b_ttt7..._14b_ttt7..._gemma4_12b_ttt7
    DFlashdeepseek-ai/dflash_qwen3_4b_block7..._8b_block7..._14b_block7..._gemma4_12b_block7
    DSparkdeepseek-ai/dspark_qwen3_4b_block7..._8b_block7..._14b_block7..._gemma4_12b_block7

    所有 checkpoint 均在各自 target 用 non-thinking 模式生成的 open-perfectblend 数据上训练,是对应 config 的直接产物。README 明确警告:引用这些结果需对齐本仓训练设置,否则对比无意义。

    • 数据准备(3 段,scripts/data/:① download_and_split.pyopen-perfectblend, test-size 0.05)② generate_train_data.py(用任意 OpenAI-兼容引擎重生成答案,示例 8 路 SGLang,temperature 0.7/top-p 0.8/top-k 20/max-tokens 4096/disable-thinking;SGLang 不在 requirements 需自装)③ prepare_target_cache.py--local-batch-size 16,产出约 38 TB 的 target cache,随数据量/序列长/隐维/target_layer_ids 数扩大)。

    6. 论证链 #

    论点代码支撑
    1三算法应能公平对比统一 BaseTrainer/BaseEvaluator,各算法只重写 _build_draft_model/run_batchconfig/ 下同构的算法×模型矩阵
    2DSpark 是 DFlash 的超集同一 Qwen3DSparkModel;DFlash config 仅置 markov_rank=0+confidence_head_alpha=0+CE-only,说明差异是"加头+加损失项"而非重写模型
    3块内多步预测可离线高密度训练离线 target cache(§5)供监督 + anchor-block FlexAttention(一序列 512 块)+ block_keep_mask 掩码 dummy,把"一次前向"变成"512 个独立块预测"
    4评测既无偏又可诊断rejection sampling 保证输出分布不变(§4-1),confidence 头额外给出逐位置接受概率用于早停或校准记录
    5大规模长跑可承受抢占hfai_suspend + StatelessResumableDistributedSampler + ckpt_manager 记录 next_micro_step,实现 micro-step 精确恢复

    7. 实现 cross-reference(核心文件:行级指引) #

    核心技术壁垒 —— anchor-block FlexAttention 掩码deepspec/modeling/dspark/common.py::create_dspark_attention_mask)。mask_mod 里 draft query 的块号 q_block_id = q_idx // block_size,其上下文只允许 kv_idx < anchor_pos,draft-draft 只允许同块(q_block_id == kv_block_id),整体再乘 block_keep_mask[b, q_block_id]

    
    def dspark_mask_mod(b, h, q_idx, kv_idx):
        q_block_id = q_idx // block_size
        anchor_pos = anchor_positions[b, q_block_id]
        mask_context = (kv_idx < seq_len) & (kv_idx < anchor_pos)
        kv_block_id = (kv_idx - seq_len) // block_size
        mask_draft = (kv_idx >= seq_len) & (q_block_id == kv_block_id)
        return (mask_context | mask_draft) & block_keep_mask[b, q_block_id]
    

    同文件配套:sample_anchor_positions / build_anchor_candidate_mask / build_eval_mask / DSparkForwardOutput(前向输出契约,含 draft_logits[B,num_anchors,block_size,vocab]block_keep_maskeval_mask)。

    其余关键实现位置

    • 训练步与 draft 定义:deepspec/trainer/dspark_trainer.py::Qwen3DSparkTrainer.run_batch(调 compute_dspark_loss);backbone deepspec/modeling/dspark/qwen3/modeling.py::Qwen3DSparkModel(注意力 is_causal=False,K/V 拼 [ctx target_hidden || noise])。
    • 损失:deepspec/modeling/dspark/loss.py::compute_dspark_loss(CE+L1+BCE、位置衰减、跨 rank 分母)。
    • Markov 头:deepspec/modeling/dspark/markov_head.pyVanillaMarkov = Embedding(V,r)Linear(r,V)GatedMarkovHead 加 sigmoid 门;RNNHead 跨块内位置保持递归状态)。
    • 评测热路径:deepspec/eval/base_evaluator.py::generate_decoding_sample(init_context/propose/update 三回调)+ verify_draft_tokens(rejection sampling);DSpark 提议 deepspec/eval/dspark/draft_ops.py::forward_dspark_draft_block / build_dspark_proposal / _confident_prefix_length(confidence 早停)。
    • Eagle3:deepspec/modeling/eagle3/loss.py::FusedLogSoftmaxLoss(Triton 融合 soft-CE,backward 原地写回 logits 存储,故其后不可再读 logits)+ compute_eagle3_loss(展开 ttt_length 步,权重 step_loss_decay ** step_idx)。
    • 数据:deepspec/data/target_cache_dataset.py(mmap 分片读写、manifest/index、validate_train_cache)+ deepspec/data/__*__ CacheDatasetmax_open_shards=4)。

    关键实现细节(易漏)

    1. train.sh/eval.shRANK/WORLD_SIZE 语义是 node_rank/node_count,入口自己 torch.multiprocessing.spawn 每 GPU 一个 worker——不是 torchrun 语义,误当全局 rank 会启动错误的进程拓扑。
    2. BaseTrainer.save_and_eval_checkpoint 调的 _launch_eval 是 stub(只打印提示),自动评测需用户自行接入。