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。
Q1 痛点。 投机解码要一个能"猜对 target 下一批 token"的小 draft model,但把 draft 训得既准又快依赖大量工程:怎么从 target 拿监督信号、怎么高效训练一个 block 的多步预测、怎么在评测时既保证输出分布无偏又度量 draft 质量。已有实现(SpecForge/Eagle3、DFlash)各写一套训练/评测栈,算法之间难以公平对比。
Q2 方法。 本仓把三种 draft 算法收敛到统一的 BaseTrainer/BaseEvaluator 骨架上,用 Python config 文件选择 trainer_cls 与超参:
target_layer_ids)预计算成 mmap 分片存储(Qwen3-4B 默认约 38 TB),训练时只读 cache。block_size=7 的 draft 块),叠加低秩 Markov 头(按前一 token 给每位置 logit 偏置)、L1 分布匹配损失、以及预测每位置接受率的 confidence 头。DFlash = 关掉两个头、纯 CE 的同一模型。核心技术壁垒:把"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)的评测入口。
顶层:train.py(训练入口)、eval.py(评测入口,按 draft 配置的 architectures[0] 选 evaluator)、deepspec/ 主包、config/(算法×目标模型的 Python 配置)、scripts/(数据/训练/评测 shell)、eval_datasets/(JSONL 基准)。
依赖边界:config/*.py 是用户主要旋钮(组装 model/train/logging/data 四个 dict + finalize_cfg);trainer_cls 与 evaluator 通过 config 字段 / draft 配置的类名字符串动态选取(EVALUATORS 字典)。BaseTrainer 提供 FSDP 包装、BF16 优化器、梯度累积、可恢复调度、suspend/checkpoint;各算法 trainer 只重写 _build_draft_model 与 run_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 基准 + 转换脚本
无形式化作者证明 — 代码库以工程正确性契约代替定理。以下 6 项是可从代码验证的正确性关键点:
accept_mask = (rand < accept_prob),再取 cumprod 得到接受前缀;首个被拒位置用 sample_residual(target, draft) 残差重采样。这保证最终 token 分布等于 target 单独解码的分布——投机解码正确性的核心不变量。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,保证整除。dspark_mask_mod 里 draft 块只能看到 kv_idx < anchor_pos 的上下文与"同块"draft key,且被 block_keep_mask[b, q_block_id] 门控——dummy anchor 因 keep_mask=False 全部屏蔽,不污染 loss。assert_no_final_target_layer 拒绝把 target 最后一层放进 target_layer_ids,因为 HF output_hidden_states 在末层存的是 归一化后 隐状态,而 cache 存的是 raw decoder 输出——不守卫会静默错配。compute_dspark_loss 的 CE/L1/BCE 分母都跨 rank all-reduce,配合位置衰减 $\exp(-pos/\gamma)$;Eagle3 侧显式断言禁用 valid_token_mean(注释指出它会降 eval accept-length),只允许 local_mean。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。
仓库提供的是"评测能力"而非论文数字:
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] 从 EVALUATORS 选 Qwen3/Gemma4 × DSpark/Eagle3 四类 evaluator。eval.py 的 TASKS 列出 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。| Algorithm | Qwen3-4B | Qwen3-8B | Qwen3-14B | Gemma4-12B |
|---|---|---|---|---|
| Eagle3 | deepseek-ai/eagle3_qwen3_4b_ttt7 | ..._8b_ttt7 | ..._14b_ttt7 | ..._gemma4_12b_ttt7 |
| DFlash | deepseek-ai/dflash_qwen3_4b_block7 | ..._8b_block7 | ..._14b_block7 | ..._gemma4_12b_block7 |
| DSpark | deepseek-ai/dspark_qwen3_4b_block7 | ..._8b_block7 | ..._14b_block7 | ..._gemma4_12b_block7 |
所有 checkpoint 均在各自 target 用 non-thinking 模式生成的 open-perfectblend 数据上训练,是对应 config 的直接产物。README 明确警告:引用这些结果需对齐本仓训练设置,否则对比无意义。
scripts/data/):① download_and_split.py(open-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 数扩大)。| 步 | 论点 | 代码支撑 |
|---|---|---|
| 1 | 三算法应能公平对比 | 统一 BaseTrainer/BaseEvaluator,各算法只重写 _build_draft_model/run_batch;config/ 下同构的算法×模型矩阵 |
| 2 | DSpark 是 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 精确恢复 |
核心技术壁垒 —— 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_mask、eval_mask)。
其余关键实现位置:
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 分母)。deepspec/modeling/dspark/markov_head.py(VanillaMarkov = 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 早停)。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/__*__ CacheDataset(max_open_shards=4)。关键实现细节(易漏):
train.sh/eval.sh 里 RANK/WORLD_SIZE 语义是 node_rank/node_count,入口自己 torch.multiprocessing.spawn 每 GPU 一个 worker——不是 torchrun 语义,误当全局 rank 会启动错误的进程拓扑。BaseTrainer.save_and_eval_checkpoint 调的 _launch_eval 是 stub(只打印提示),自动评测需用户自行接入。