news 2026/9/10 9:09:11

SGLang 中的 Reasoning-Aware Compression:用模型自身的思维链校准一次性剪枝推理模型

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
SGLang 中的 Reasoning-Aware Compression:用模型自身的思维链校准一次性剪枝推理模型

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.93623.3 min
C40.744135.0 min
仅任务提示词0.812115.6 min
RAC(提示词 + on-policy CoT)0.90035.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 500

prompt_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 明确给出):

  1. Phase I 的 decode 占比应明显高于 50%——这个缺口正是 prompt-only 校准丢弃的激活质量。运行结束时脚本会打印=== RAC calibration set ===报告块,包含 rows、prompt tokens、decode tokens、total tokens、decode share 与墙钟时间(见 rac_collect_traces.py)。
  2. Phase II 应报告与目标稀疏度几乎一致的 realized sparsity——rac_prune.py在剪枝后会对所有非lm_headLinear层统计零权重占比并打印(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.pyrac_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.Engineskip_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.6top_p=0.95max_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_tokensnum_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--methodsparsegpt默认 /wanda)实例化SparseGPTModifierWandaPruningModifier,目标是全部Linear层,忽略lm_headignore=["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-500testsplit 上评估 500 题(--num-problems默认 500,即完整 MATH-500 规模),--max-new-tokens默认 8192(README 注明论文评估时使用 32768),采样参数与 Phase I 一致(T=0.6, top_p=0.95)。每个模型用独立的sgl.Engine生成并shutdownevaluate函数,同文件 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 9:09:01

Swin-Transformer水果图像分类迁移学习实践指南

简介&#xff1a;面向图像分类与迁移学习实践者&#xff0c;这是一套基于Swin-Transformer的水果十二分类图像识别项目&#xff0c;可直接运行并支持替换为自己的数据集。数据集涵盖香蕉、苹果、西瓜等12类水果&#xff0c;包含2340张训练图片与581张预测图片&#xff1b;模型采…

作者头像 李华
网站建设 2026/9/10 9:06:18

结果驱动的夹具动态选择:测试用例依赖与资源装配实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/10 9:06:02

2026 专业语言服务商测评:国内高适配翻译机构选型参考

一、引言根据中国翻译协会发布的《2026 中国翻译行业发展报告》公开统计数据&#xff0c;2025 年国内翻译行业总产值达到701.2 亿元&#xff0c;全行业从业人员规模686.7 万人&#xff0c;专职翻译人员113.5 万人&#xff0c;主营翻译业务的市场主体共计14393 家。进入 2026 年…

作者头像 李华
网站建设 2026/9/10 9:05:53

LeetCode-Go 题解:127. Word Ladder 单词接龙 BFS 最短转换序列

LeetCode-Go 题解&#xff1a;127. Word Ladder 单词接龙 BFS 最短转换序列 【免费下载链接】LeetCode-Go ✅ Solutions to LeetCode by Go, 100% test coverage, runtime beats 100% | LeetCode 题解 项目地址: https://gitcode.com/GitHub_Trending/le/LeetCode-Go 导…

作者头像 李华
网站建设 2026/9/10 9:04:55

CANN/ge算子参数更新API

aclopUpdateParams 【免费下载链接】ge GE&#xff08;Graph Engine&#xff09;是面向昇腾的图编译器和执行器&#xff0c;提供了计算图优化、多流并行、内存复用和模型下沉等技术手段&#xff0c;加速模型执行效率&#xff0c;减少模型内存占用。 GE 提供对 PyTorch、TensorF…

作者头像 李华
网站建设 2026/9/10 9:03:30

基于context-mode的LLM上下文管理:分层、归档与召回实战

我之前负责过一个文档问答机器人&#xff0c;上线三个月后被投诉最多的问题就是“聊着聊着它就忘了”。用户早上进来问合同审核清单&#xff0c;下午回来接着问&#xff0c;模型已经完全不记得合同附件里写了什么&#xff0c;甚至会把另一个项目的条款内容混进来。一开始我以为…

作者头像 李华