源文件:chapter8/prompt-distillation/README.md 基于 Hugging Face TRL 的 Prompt 蒸馏 本项目演示Prompt 蒸馏——把带长 prompt 的思考模型中的知识,蒸馏到不带 prompt 的非思考模型中,从而大幅加快响应速度。 🎯 主要目标 把推理能力从以下教师蒸馏到学生: 教师:Qwen3-30B-A3B-Thinking-2507,配合一段详尽的 2000+ token prompt 学生:Qwen3-30B-A3B-Instruct-2507,不带任何 prompt 主要收益: ⚡ 响应速度大幅提升 —— 没有思考开销,也无需处理长 prompt 💰 推理成本更低 —— 每次请求要处理的 token 更少 🎯
源文件:chapter8/prompt-distillation/README.md
本项目演示Prompt 蒸馏——把带长 prompt 的思考模型中的知识,蒸馏到不带 prompt 的非思考模型中,从而大幅加快响应速度。
把推理能力从以下教师蒸馏到学生:
主要收益:
Prompt 蒸馏(又称上下文蒸馏)是一种训练方法,它让 LLM 把一段又长又复杂的 prompt 内化到模型参数中。在本实验中,我们还通过从思考模型蒸馏到非思考模型,消除了思考开销。
示例 —— 语言分类:
我们希望把这段详尽的 prompt 内化:
"Classify the language of the provided text into these labels: ar, de, el, en, es, fr, hi, ru, tr, ur, vi, zh, ot. Use these rules: Devanagari script → hi, Greek script → el, Cyrillic script → ru..." (2000+ tokens)
蒸馏前(教师,带思考 + prompt):
System: <2000+ token detailed prompt> User: 一生、バンドしてくれる? Assistant: <thinking>Let me analyze the script... These are Han characters... Based on rule X...</thinking>ja ⏱️ Response time: ~2-3 seconds
蒸馏后(学生,无思考、无 prompt):
User: 一生、バンドしてくれる? Assistant: ja ⏱️ Response time: ~0.1 seconds (20-30x faster!)
该方法分为两个阶段:
数据生成(教师模型):一个思考模型借助详尽的 prompt,生成带显式推理的响应。
response = thinking_model(long_prompt, query)学生训练(蒸馏):微调一个非思考模型,使其在不带 prompt、也不走思考过程的情况下直接预测响应。
non_thinking_model(query) ≈ thinking_model(long_prompt, query)本实现采用 OpenAI Cookbook 的超参数(来自 gpt-oss-20b 示例):
| 参数 | 取值 | 来源 |
|---|---|---|
| 教师模型 | Qwen3-30B-A3B-Thinking-2507 | 带思考能力 + 长 prompt |
| 学生模型 | Qwen3-30B-A3B-Instruct-2507 | 同等规模,无思考、无 prompt |
| LoRA Rank | 32 | tinker |
| LoRA Alpha | 16 | Standard |
| 学习率 | 2e-4 | OpenAI |
| 学习率调度 | cosine_with_min_lr | OpenAI |
| 最小学习率比例 | 0.1 | OpenAI |
| Batch Size | 每块 GPU 4 | OpenAI |
| 梯度累积 | 4 步 | OpenAI |
| 最大长度 | 2048 | OpenAI(学生只需短上下文) |
| 训练轮数 | 1 | OpenAI |
| 温度 | 0.15 | tinker(数据生成) |
| 预热比例 | 0.03 | OpenAI |
| 梯度检查点 | True | OpenAI |
关键设计选择:教师与学生都用同一个 30B 模型。差异在于:
这不是在做模型规模压缩,而是消除思考开销与 prompt 处理以加快推理。
本项目使用与 tinker 相同的多语言语言分类任务:
ar, de, el, en, es, fr, hi, ru, tr, ur, vi, zh, otexample-data/multilingual.txt(2101 句)pip install -r requirements.txt
# Login to wandb (required for training progress tracking) wandb login # Or set your API key as environment variable export WANDB_API_KEY=your_api_key_here
可从 https://wandb.ai/settings 获取你的 API key。
用教师模型生成 Prompt 蒸馏数据:
# Single instance (uses tensor parallelism across GPUs) python create_data.py \ --input_file ./example-data/multilingual.txt \ --output_file ./data/prompt_distillation_lang.jsonl \ --model_name Qwen/Qwen3-30B-A3B-Thinking-2507 \ --temperature 0.15 \ --tensor_parallel_size 4 # For H100x8 users: Run 2 parallel instances to use all 8 GPUs bash create_data_h100x8.sh
选项:
--input_file:输入句子路径(每行一句)--output_file:生成训练数据的保存位置--model_name:教师模型(Qwen3-30B-A3B-Thinking-2507 精度更好)--temperature:采样温度(0.15 与 tinker 一致)--tensor_parallel_size:用于推理的 GPU 数(推荐 4)--max_retries:失败样本的重试次数(默认:3)它会:
输出格式:
{ "messages": [ {"role": "user", "content": "Text in some language"}, {"role": "assistant", "content": "en"} ] }
用 TRL 在蒸馏数据上微调学生模型:
# Single GPU training (recommended - simpler and works reliably) bash train_trl.sh
监控训练:
如需关闭 wandb 日志:
python train_sft_trl.py --report_to none ...other args...
训练完成后,评估蒸馏后模型的表现:
# Evaluate with defaults (uses all defaults) python evaluate.py # Quick evaluation on a subset python evaluate.py --max_samples 100 # Save results to a file python evaluate.py --output_file ./evaluation_results.json # Custom model path python evaluate.py --model_path ./models/my_custom_model
默认值:
./models/prompt_distillation_trlQwen/Qwen3-30B-A3B-Instruct-2507./example-data/multilingual.txt实时输出示例:
Evaluating model... ================================================================================ ✓ [ 1/2100] Pred: ar | GT: ar | Acc: 1/1 (100.0%) | وقال، ماما، لقد عدت للمنزل. ✓ [ 2/2100] Pred: ru | GT: ru | Acc: 2/2 (100.0%) | И той каза: Мамо, у дома съм. ✓ [ 3/2100] Pred: de | GT: de | Acc: 3/3 (100.0%) | und er hat gesagt, Mama ich bin daheim. ✓ [ 4/2100] Pred: el | GT: el | Acc: 4/4 (100.0%) | Και είπε, Μαμά, έφτασα στο σπίτι. ✓ [ 5/2100] Pred: en | GT: en | Acc: 5/5 (100.0%) | And he said, Mama, I'm home. ✗ [ 6/2100] Pred: es | GT: en | Acc: 5/6 ( 83.3%) | Y él dijo: Mamá, estoy en casa. ✓ [ 7/2100] Pred: fr | GT: fr | Acc: 6/7 ( 85.7%) | Et il a dit, maman, je suis à la maison. ✓ [ 8/2100] Pred: hi | GT: hi | Acc: 7/8 ( 87.5%) | और उसने कहा, माँ, मैं घर आया हूं। ✓ [ 9/2100] Pred: ru | GT: ru | Acc: 8/9 ( 88.9%) | И он сказал: Мама, я дома. ... ✗ [2092/2100] Pred: de | GT: ot | Acc: 1994/2092 ( 95.3%) | Hola, mein Freund ✓ [2093/2100] Pred: ru | GT: ru | Acc: 1995/2093 ( 95.3%) | Привет, hello ✗ [2094/2100] Pred: vi | GT: ot | Acc: 1995/2094 ( 95.3%) | Xin chào, merci beaucoup ✗ [2095/2100] Pred: hi | GT: ot | Acc: 1995/2095 ( 95.2%) | नमस्ते, good morning ✓ [2096/2100] Pred: en | GT: en | Acc: 1996/2096 ( 95.2%) | ok ✓ [2097/2100] Pred: en | GT: en | Acc: 1997/2097 ( 95.2%) | yes ✓ [2098/2100] Pred: fr | GT: fr | Acc: 1998/2098 ( 95.2%) | bonjour ✓ [2099/2100] Pred: es | GT: es | Acc: 1999/2099 ( 95.2%) | hola ✗ [2100/2100] Pred: hi | GT: ot | Acc: 1999/2100 ( 95.2%) | namaste ================================================================================ Evaluation completed: 2100 samples processed ================================================================================ CONFUSION MATRIX ================================================================================ ar de el en es fr hi ot ru tr ? ur vi zh | Total ------------------------------------------------------------------ ar | 141 . . . . . . . . . . 5 . . | 146 de | . 135 . . 1 . . . . 1 . . . . | 137 el | . . 144 . . . 3 10 . . . . . . | 157 en | . . . 146 3 . . . . . . . . . | 149 es | . 3 . . 133 . . . . . . . . . | 136 fr | . . . 1 . 139 . . . 3 . . . . | 143 hi | . . . . . . 171 39 . . . . . 1 | 211 ot | . 1 . 4 . . 14 169 1 . . . 1 3 | 193 ru | . . . . . . . . 279 . . . . . | 279 tr | . . . . . . . . . 132 . . . . | 132 ? | . . . . . . . 1 . 1 . . 3 . | 5 ur | . . . . . . 1 . . . . 133 . . | 134 vi | . . . . . . . . . . . . 134 . | 134 zh | . . . . . . . 1 . . . . . 143 | 144 ================================================================================ PER-LANGUAGE ACCURACY ================================================================================ ✗ unknown: 0.0% ( 0/ 5) ⚠️ hi: 81.0% ( 171/ 211) ⚠️ ot: 87.6% ( 169/ 193) ✓ el: 91.7% ( 144/ 157) ✓ ar: 96.6% ( 141/ 146) ✓ fr: 97.2% ( 139/ 143) ✓ es: 97.8% ( 133/ 136) ✓ en: 98.0% ( 146/ 149) ✓ de: 98.5% ( 135/ 137) ✓ ur: 99.3% ( 133/ 134) ✓ zh: 99.3% ( 143/ 144) ✓ ru: 100.0% ( 279/ 279) ✓ tr: 100.0% ( 132/ 132) ✓ vi: 100.0% ( 134/ 134) ================================================================================ MOST PROBLEMATIC LANGUAGES (Top 5) ================================================================================ 1. Language: unknown - Accuracy: 0.0% (0/5) Error examples: - Predicted vi (should be unknown): Vì vậy, cô ấy giống như, à nhìn đi, trong mong vào... - Predicted vi (should be unknown): Thế là, à, tôi, ừ, dù sao, ừ, ừ, đây là ba, ừ, phi... - Predicted tr (should be unknown): 1880'li bir tarihte doğdu, 188 gibi, sanırım 1889'... 2. Language: hi - Accuracy: 81.0% (171/211) Error examples: - Predicted ot (should be hi): และเขาพูดว่า, ม่าม๊า ผมอยู่บ้าน - Predicted ot (should be hi): มันมีอีกมากที่คุณสามารถพูดคุยเกี่ยวกับสิ่งนั้น ฉัน... - Predicted ot (should be hi): และฉันก็แบบว่าตอบตกลงและมันก็เท่านั่น! 3. Language: ot - Accuracy: 87.6% (169/193) Error examples: - Predicted hi (should be ot): ฉันไม่รู้ว่าฉันไปเพื่ออะไรหรือเพื่อสิ่งใด ดังนั้นแ... - Predicted hi (should be ot): วันนี้เขาจะพูดคุยกับเราเกี่ยวกับ Third SS, U2 Quic... - Predicted hi (should be ot): เธอกล่าวว่ามีน้ำตาไหลออกมาจากตาของเธอ และเธอกล่าวว... 4. Language: el - Accuracy: 91.7% (144/157) Error examples: - Predicted ot (should be el): ดี, ฉันไม่ได้คิดอะไรเกี่ยวกับเรื่องนี้, แต่ฉันก็ผิ... - Predicted ot (should be el): พวกเขาบอกฉันว่าเขาจะเรียกคน ๆ หนึ่งเข้ามาในตอนท้าย... - Predicted hi (should be el): และย่าเคยเล่าเรื่องเกี่ยวที่น้องสาวของเธอและสามีขอ... 5. Language: ar - Accuracy: 96.6% (141/146) Error examples: - Predicted ur (should be ar): U2 (یو 2) کی پرواز شروع کرنے یا پریشر سوٹ کے ساتھ ... - Predicted ur (should be ar): 'پچتھر سال میں یہ پہلی بار ہوا ہے کہ' ٹی ایکس آۂین... - Predicted ur (should be ar): میرا مطلب یہ تھا کہ پوری بات. ================================================================================ ============================================================ EVALUATION SUMMARY ============================================================ Model: Qwen/Qwen3-30B-A3B-Instruct-2507 Adapter: ./models/prompt_distillation_trl Performance: Total samples: 2100 Successfully predicted: 2100 Unparseable responses: 0 Parse rate: 100.00% Overall Accuracy: 95.19% Correct: 1999/2100 💡 The model responds directly without the 2000+ token prompt! 📁 Complete results saved to: ./evaluation_results.json Includes: predictions, confusion matrix, per-language stats, error examples
保存为 JSON:
评估结果会保存:
Prompt 蒸馏的全部意义,就浓缩在一次前后对比里:同一个任务,由教师(长 prompt + 思考)完成
对比学生(无 prompt、直接作答)——能省下多少输入成本,又能保留多少质量。compare.py 完全离线地从真实数据集、教师标签和评估结果中计算这一对比(无需下载模型、无需联网):
# Use the default tiktoken counter (works offline, reproducible) python compare.py # For the EXACT Qwen token counts (on a machine with the tokenizer available) python compare.py --tokenizer Qwen/Qwen3-30B-A3B-Instruct-2507 # Show more per-case examples and save the full breakdown python compare.py --num_examples 20 --output_file ./comparison_results.json
它会给出三项指标——全部来自真实数据,不做任何估算:
evaluation_results.json 读取。实测结果(本仓库数据,tiktoken o200k_base 计数器):
| 维度 | 教师(长 prompt + 思考) | 学生(无 prompt) | 变化 |
|---|---|---|---|
| 每次调用平均输入 token | 984.9 | 24.7 | −97.5%(约少 40 倍) |
| 总输入 token(2100 例) | 2,068,204 | 51,913 | −97.5% |
| 任务质量(与教师一致率) | 100%(参考) | 95.19%(1999/2100) | −4.8 pp |
在按输入 token 计费的 API 上,这种输入量的降低大致会按比例拉低成本;教师还会额外消耗
思考(CoT)的输出 token,这里未计入,因此真实差距更大。墙上时钟延迟取决于推理栈,
必须在 GPU 上实测——compare.py 刻意不编造延迟数值。精确的 token 数会因分词器而异;
如需学生模型自身的计数,请传入 --tokenizer。
prompt-distillation/ ├── README.md # This file ├── requirements.txt # Python dependencies ├── create_data.py # Data generation script (Step 1) ├── create_data_h100x8.sh # Parallel data generation for H100x8 ├── train_sft_trl.py # Training script using TRL (Step 2) ├── train_trl.sh # Training script (single GPU) ├── evaluate.py # Evaluation script (Step 3) ├── compare.py # Before/after cost & quality comparison (Step 4, offline) ├── data/ # Generated training data │ └── prompt_distillation_lang.jsonl └── models/ # Trained model checkpoints └── prompt_distillation_trl/
本实验的主要创新,是从思考模型蒸馏到非思考模型:
思考模型(教师):
<thinking>...</thinking>非思考模型(学生):
本实现使用 Hugging Face TRL,原因如下:
在配合 LoRA 做监督微调上,TRL 提供的能力相当,但 API 更为友好。
训练数据使用 TRL/Transformers 所期望的标准对话格式:
{ "messages": [ {"role": "user", "content": "Text to classify"}, {"role": "assistant", "content": "language_code"} ] }
TRL 会自动:
本实现紧跟 tinker cookbook 的方法论,并有一处关键增强:
相同:
增强:
为什么这样更好:
训练完成后,学生模型(Qwen3-30B-A3B-Instruct)应当能:
输入成本对比(实测,非估算):
运行 python compare.py 即可在本仓库数据上复现真实数字。使用默认的tiktoken o200k_base 计数器,学生每次调用处理的输入 token 约少 40 倍
(984.9 → 24.7,降幅 97.5%),同时保留 95.19% 与教师标签的一致率。
完整明细见上文 用法 → 第 4 步 下的表格。
这使得蒸馏模型在输入成本敏感的生产部署中颇具吸引力。注意:墙上时钟延迟取决于
推理栈与硬件,必须在 GPU 上实测——本 README 不引用任何编造的延迟数字。
30B 模型需要大显存的 H100 GPU(80GB)。若遇到 OOM:
解决方案:
per_device_train_batch_size 从 4 降到 2 或 1max_length 从 2048 降到 1024 或 512gradient_accumulation_steps 以维持等效 batch sizelora_rank 从 32 降到 16 或 8备选:使用更小的模型
如果没有 80GB GPU,可换用更小的模型:
--model_name显存需求:
若数据生成失败或很慢:
tensor_parallel_size 以使用更多 GPUbash create_data_h100x8.shnvidia-smi 检查 GPU 显存使用若模型学不到东西:
如果你使用了本代码,请引用原始论文:
@article{askell2021general, title={A general language assistant as a laboratory for alignment}, author={Askell, Amanda and others}, journal={arXiv preprint arXiv:2112.00861}, year={2021} } @article{snell2022learning, title={Learning by distilling context}, author={Snell, Charlie and Klein, Dan and Zhong, Ruiqi}, journal={arXiv preprint arXiv:2209.15189}, year={2022} }
以及 Hugging Face TRL 库:
@software{trl2024, title={TRL: Transformer Reinforcement Learning}, author={TRL contributors}, url={https://github.com/huggingface/trl}, year={2024} }
本项目遵循与 TRL 库相同的许可证(Apache 2.0)。