SGLang 中的 Reasoning-Aware Compression:用模型自身的思维链校准一次性剪枝推理模型
【免费下载链接】sglangSGLang is a high-performance serving framework for large language models and multimodal models.项目地址: https://gitcode.com/GitHub_Trending/sg/sglang
推理模型(Reasoning Model)会为每个问题生成成千上万个思维链(Chain-of-Thought, CoT)token,这颠覆了传统剪枝校准所依赖的 token 分布假设——直接套用 SparseGPT、Wanda 等经典一次性剪枝流程,会让剪枝后的推理模型"胡言乱语"(ramble):思维链更长、准确率更低、端到端反而更慢。Reasoning-Aware Compression(RAC,出自论文Reasoning Models Can be Accurately Pruned Via Chain-of-Thought Reconstruction,ICLR 2026,arXiv:2509.12464)给出了一个只改一行校准数据的修复方案:用稠密模型自身采样的思维链轨迹来构建剪枝校准集。SGLang 仓库以三个可直接运行的脚本完整实现了这一配方,本文带你从原理、源码到实战命令完整掌握这套流程。
问题:提示词校准与推理模型的分布错配
一次性逐层剪枝(one-shot layer-wise pruning)选择要删除的权重,本质上是求解一个逐层重建误差最小化问题:
min_{W'} || W X - W' X ||_F^2 s.t. ||W'||_0 <= S其中X是校准激活矩阵,W是原始权重,W'是稀疏化后的权重,S是稀疏度约束(非零权重数量上限)。经典流程(SparseGPT、Wanda 等)都用提示词(prompt)token来构建X——无论是 C4 语料还是任务提示词。当|prompt| >> |output|时,这是一个合理的代理。
但推理模型把比例反转过来了:每个查询会产出成千上万个思维链 token,剪枝后的模型将来运行的前向传播,几乎全部作用在模型自己生成的 token上。仅仅用提示词校准,等于让求解器为一个模型几乎不会遇到的分布做优化。
这种校准错配的后果比单纯的精度下降更糟:校准不良的剪枝推理模型会ramble——它输出更多的思考 token,却回答得更不准确,于是剪枝反而让模型变慢。论文在 DeepSeek-R1-Distill-Qwen-7B、MATH-500、SparseGPT 50% 稀疏度、1M 校准 token 的设定下给出了直接对比:
| 校准集 | acc@1 | 评估墙钟时间 |
|---|---|---|
| Dense(未剪枝) | 0.936 | 23.3 min |
| C4 | 0.744 | 135.0 min |
| 仅任务提示词 | 0.812 | 115.6 min |
| RAC(提示词 + on-policy CoT) | 0.900 | 35.3 min |
可以看到,C4 校准的 50% 稀疏模型评估耗时是稠密模型的近 6 倍(135.0 min vs 23.3 min);而 RAC 校准在几乎追平稠密准确率(0.900 vs 0.936)的同时,把耗时压到了 35.3 min。
RAC 的修复:校准集的"drop-in 替换"
RAC 的修复只动算法的一行:采样稠密模型自身的 rollout(on-policy 采样),让校准激活矩阵同时包含提示词与解码(decode)两部分的激活:
X_RAC = [ X_prompt , X_decode ]剪枝求解器完全不动——RAC 对 SparseGPT、Wanda 等而言就是一次校准集的即插即用替换(drop-in calibration-set swap)。这正是其工程价值的核心:你已有的任何 SparseGPT/Wanda 剪枝流水线,只需要把校准数据换成 RAC 轨迹即可获得收益。
为什么这个配方会落在 SGLang 里
收集 rollout 是论文 Algorithm 1 的Phase I,也是整个流程最昂贵的一半:论文的预算是每个校准集1M 个 on-policy CoT token。这是批量自回归生成(batched autoregressive generation),恰恰是 SGLang 的强项。剪枝求解器本身与推理引擎无关,因此Phase II委托给llm-compressor(SGLang 并不依赖它),最后由 SGLang 服务剪枝结果。
SGLang 仓库把这一配方以三个可运行脚本封装在 examples/usage/reasoning_aware_compression 目录下:
rac_collect_traces.py Phase I sgl.Engine 采样 on-policy CoT -> traces.jsonl rac_prune.py Phase II llm-compressor SparseGPT/Wanda -> 剪枝 checkpoint rac_serve_and_eval.py Phase III sgl.Engine 评分 MATH-500 -> acc + CoT 长度 + 运行时间三个阶段各司其职,下文逐一展开源码级讲解。
环境准备
Phase I 与 Phase III 只需要 SGLang 本身。Phase II 额外需要llm-compressor,它不是SGLang 的依赖:
pip install "llmcompressor>=0.12.0"仓库内实测针对llmcompressor0.12.0 版本(README 明确标注 "Tested againstllmcompressor0.12.0")。若未安装,rac_prune.py会抛出带有安装提示的ImportError(源码中的INSTALL_HINT字符串,见 rac_prune.py)。
完整运行:复现论文的 50% 稀疏度设定
下面三条命令复现论文 DeepSeek-R1-Distill-Qwen-1.5B 在 50% 稀疏度下的整条流水线(论文的全部一次性剪枝实验均在单张 H100 上完成):
cd examples/usage/reasoning_aware_compression # Phase I —— 1M on-policy CoT token(论文的预算),T_max = 8192, T = 0.6, top_p = 0.95。 python rac_collect_traces.py \ --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ --dataset open-r1/OpenR1-Math-220k \ --prompt-column problem \ --target-tokens 1000000 \ --output-dir ./rac_traces_math # Phase II —— SparseGPT 50% 非结构化稀疏度,用上面的轨迹校准。 python rac_prune.py \ --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ --calibration ./rac_traces_math/traces.jsonl \ --sparsity 0.5 \ --output-dir ./rac_pruned_50 # Phase III —— 同时报告准确率、CoT 长度与墙钟时间。 python rac_serve_and_eval.py --model-path ./rac_pruned_50 --num-problems 500提示:README 与文档正文中的完整运行示例默认使用
rac_serve_and_eval.py直接评分。若想以标准服务方式部署剪枝 checkpoint,也可以在 Phase II 之后运行python -m sglang.launch_server --model-path ./rac_pruned_50(该命令在 rac_prune.py 的收尾输出中直接给出)。
对照实验:构建 prompt-only 基线
要看清 RAC 真正买到了什么,可以用同一批提示词构建论文的 prompt-only 基线,再直接对比两个 checkpoint:
python rac_collect_traces.py \ --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ --dataset open-r1/OpenR1-Math-220k --prompt-column problem \ --calibration-mode prompt_only \ --target-tokens 1000000 \ --output-dir ./prompt_only_traces_math python rac_prune.py \ --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ --calibration ./prompt_only_traces_math/traces.jsonl \ --sparsity 0.5 --output-dir ./prompt_only_pruned_50 python rac_serve_and_eval.py \ --model-path ./prompt_only_pruned_50 ./rac_pruned_50 \ --num-problems 500prompt_only模式完全跳过生成,只付出 tokenization 的代价(README 原话:it costs nothing but the tokenization pass)。源码层面,该模式在 rac_collect_traces.py 中直接把 decode 输出替换为空列表[[] for _ in batch],且不会创建推理引擎(llm = sgl.Engine(...) if args.calibration_mode == "rac" else None,见同文件 L377)。
冒烟测试:先验证管道再烧 1M token
在投入 1M token 的大规模运行之前,先用几分钟在单 GPU 上验证整条管道是否通畅:
python rac_collect_traces.py --model-path Qwen/Qwen3-0.6B \ --dataset open-r1/OpenR1-Math-220k --prompt-column problem \ --target-tokens 20000 --max-new-tokens 1024 --output-dir /tmp/rac_traces python rac_prune.py --model-path Qwen/Qwen3-0.6B \ --calibration /tmp/rac_traces/traces.jsonl --sparsity 0.5 --output-dir /tmp/rac_pruned python rac_serve_and_eval.py --model-path /tmp/rac_pruned --num-problems 50 --max-new-tokens 2048冒烟测试的两个验收标准(README 明确给出):
- Phase I 的 decode 占比应明显高于 50%——这个缺口正是 prompt-only 校准丢弃的激活质量。运行结束时脚本会打印
=== RAC calibration set ===报告块,包含 rows、prompt tokens、decode tokens、total tokens、decode share 与墙钟时间(见 rac_collect_traces.py)。 - Phase II 应报告与目标稀疏度几乎一致的 realized sparsity——
rac_prune.py在剪枝后会对所有非lm_head的Linear层统计零权重占比并打印(measure_sparsity函数,见 rac_prune.py)。
模型与数据集选择
论文评估了 DeepSeek-R1-Distill-Qwen(1.5B/7B/14B/32B)与 Qwen3(1.7B/8B/14B),在 20–50% 稀疏度下剪枝。以上任意一个都可以在这里使用;更大的模型通过--tp-size参数做张量并行分片(rac_collect_traces.py与rac_serve_and_eval.py均支持该参数)。
校准提示词遵循论文约定:
- 数学任务:
open-r1/OpenR1-Math-220k数据集,--prompt-column problem; - 代码任务:CodeForces 提示词集,
--prompt-column prompt; --dataset也接受本地.jsonl文件路径——源码中load_rows会先用os.path.exists(dataset)判断:本地路径存在则按 JSON 数据集加载,否则按 Hugging Face 数据集 id 加载(见 rac_collect_traces.py)。
三个阶段的源码级深入
Phase I:rac_collect_traces.py—— 用 sgl.Engine 采样校准轨迹
这一脚本完成论文 Algorithm 1 的 Phase I(采样 rollout),是成本最高的一半。其核心设计与实现要点:
批量自回归生成。引擎通过sgl.Engine以skip_tokenizer_init=True启动(见 rac_collect_traces.py),提示词由外层 tokenizer 预先模板化并 tokenize 成input_ids,随后直接以 token id 喂给llm.generate(input_ids=...)(rollout函数,同文件 L242-L247)。skip_tokenizer_init是 SGLang 服务端参数之一(字段定义于 python/sglang/srt/arg_groups/fields/serving.py),此处由 Phase I 自行负责模板化,服务端不再重复初始化 tokenizer。
on-policy 采样参数与论文保持一致:temperature=0.6、top_p=0.95、max_new_tokens=8192(T_max,单条 rollout 的 CoT 长度上限)、num_generations=2(每个提示词采样 2 条 rollout,论文同款),默认随机种子 42。采样参数以字典形式传入llm.generate(L360-L364)。
逐块(chunk)懒加载。数学语料有 22 万行,远大于任何 token 预算,因此iter_question_chunks按--chunk-size(默认 256)逐块读取并生成,1M token 的运行只触碰实际用到的提示词(见 rac_collect_traces.py)。达到--target-tokens预算即停止(collect_traces主循环中的if total >= target_tokens: break,同文件 L316-L317)。注意 chunk 粒度决定了超出 token 预算的上界——README 称之为 "Bounds how far past the token budget a run can overshoot"。
输出格式:每条校准序列一行 JSON,包含input_ids(模板化提示词 + 模型自身续写的 token id 拼接)、num_prompt_tokens、num_decode_tokens,以及可选的text字段(可用--no-text关闭以缩小文件,因为剪枝器实际只读 token id)。输出为traces.jsonl,同时在同一目录写入rac_manifest.json记录完整 provenance(模型、数据集、模式、token 统计、采样参数、耗时等,由TraceManifestmsgspec 结构体序列化,见 L80-L98 与 L397-L418)。结束后脚本会打印下一步剪枝命令的提示。
校准模式:--calibration-mode二选一,rac(默认,提示词 + on-policy CoT,对应论文 Eq. 7)或prompt_only(仅提示词,作为论文的消融基线)。prompt_only模式不创建引擎、不采样。
Phase II:rac_prune.py—— 委托 llm-compressor 求解
该脚本刻意保持"薄":RAC 的全部贡献在于求解器重建哪些激活(提示词 + 模型自身 decode token),因此这里只是把 RAC 校准集交给llm-compressor的 SparseGPT 或 Wanda 实现,并把结果保存为 SGLang 可服务的 checkpoint。
校准序列加载:load_calibration_sequences逐行读取traces.jsonl,取input_ids并按--max-seq-length(默认 8192)截断;--num-samples可限制使用的序列数(默认全部使用),抽样前用种子固定的random.Random(seed).shuffle打乱(见 rac_prune.py)。
校准 DataLoader 的 batch size 恒为 1:源码注释明确指出,对不同长度的序列做 batch 需要 padding,而 pad token 的激活会混入逐层 Hessian,仿佛它们是真实激活一样——这正是 RAC 想要避免的校准污染(见 rac_prune.py)。collate 函数把单个序列包装为(input_ids, attention_mask)张量,attention_mask全 1。
求解器配方:build_recipe按--method(sparsegpt默认 /wanda)实例化SparseGPTModifier或WandaPruningModifier,目标是全部Linear层,忽略lm_head(ignore=["re:.*lm_head"],见 rac_prune.py)。
单 GPU 适配:--pipeline sequential(默认)让llm-compressor每次只驻留一个 decoder 层的 Hessian,这是它能在单卡上跑通的关键(README 注释:"sequential keeps only one decoder layer's Hessians resident, which is what fits on one GPU")。其他参数还包括--dtype(默认bfloat16)、--device-map(默认auto)、--seed。
Phase III:rac_serve_and_eval.py—— 同时报告准确率、CoT 长度与墙钟时间
该脚本的核心意图是报告一对数字:只报准确率会掩盖校准漂移导致的 ramble 故障。因此每一行都同时给出 acc@1、平均完成长度(mean completion tokens)与墙钟时间(见report函数,rac_serve_and_eval.py),输出类似:
=== MATH-500 === model acc@1 mean CoT tokens wall clock ...多 checkpoint 头对头对比:--model-path接受多个路径,依次评估并排显示(示例中的./pruned_prompt_only ./pruned_rac对比)。如果某个剪枝模型分数更差且产生更多 CoT token,脚本会打印提示:这正是 RAC 针对的失败模式——校准漂移让它 ramble,既更不准确又更慢。
轻量判分:extract_boxed提取最后一个花括号配平的\boxed{...}内容,normalize_answer去除 LaTeX 噪声(\left、\right、\!、\,、$,并把\dfrac/\tfrac归一为\frac,见 rac_serve_and_eval.py)。这种 boxed-answer 匹配足以在开发期对 checkpoint 排序;要得到论文级别的数字,请使用 RAC 与 open-r1 仓库所用的lightevalharness。
评估配置:默认在HuggingFaceH4/MATH-500的testsplit 上评估 500 题(--num-problems默认 500,即完整 MATH-500 规模),--max-new-tokens默认 8192(README 注明论文评估时使用 32768),采样参数与 Phase I 一致(T=0.6, top_p=0.95)。每个模型用独立的sgl.Engine生成并shutdown(evaluate函数,同文件 L157-L195)。
关键设计注意点
以下是 README "Notes" 小节逐条列举的、直接影响方法有效性的设计约束:
- Chat 模板必须保持一致。轨迹通过模型自身的 chat 模板配合 open-r1 system prompt 生成——这正是参考实现发布轨迹所用的配置。校准分布本身就是方法,因此修改
--system-prompt会改变结果。默认的 open-r1 GRPO system prompt 完整定义在 rac_collect_traces.py:要求模型先内部思考再以<think>...</think><answer>...</answer>格式回答。传入空字符串可省略 system message。 - 存 token id 而非文本。Phase I 输出 token id,Phase II 直接消费它们,因此剪枝器重建的序列与模型实际产生的序列严格一致——没有 detokenize/retokenize 漂移。
- 校准期间 batch size 恒为 1。padding token 会像真实激活一样进入逐层 Hessian,这正是 RAC 要避免的污染。
2:4掩码。传--mask-structure 2:4可获得半结构化掩码;论文的头条结果是非结构化(0:0,默认值)。- 幅度剪枝(magnitude pruning)不在其中。参考实现里有,但这里没有暴露:
llm-compressor的 magnitude modifier 是渐进式的训练期 modifier,而非一次性求解器,而 RAC 是一次性方法。 - 判分口径。
rac_serve_and_eval.py做轻量 boxed-answer 匹配,足以对 checkpoint 排序;论文级数字请用 RAC 与 open-r1 仓库使用的lightevalharness。
相关文档
- 本文对应的高级特性文档:docs/docs/advanced_features/reasoning_aware_compression.mdx(含问题描述、修复原理与三阶段使用说明)
- 量化的另一条压缩轴(服务期应用):docs/docs/advanced_features/quantization.mdx
引用
如需引用论文:
@inproceedings{lucas2026reasoning, title = {Reasoning Models Can be Accurately Pruned Via Chain-of-Thought Reconstruction}, author = {Lucas, Ryan and Behdin, Kayhan and Wang, Zhipeng and Tang, Shao and Song, Qingquan and Mazumder, Rahul}, booktitle = {International Conference on Learning Representations (ICLR)}, year = {2026} }【免费下载链接】sglangSGLang is a high-performance serving framework for large language models and multimodal models.项目地址: https://gitcode.com/GitHub_Trending/sg/sglang
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考