基于 Hugging Face TRL 的 Prompt 蒸馏


文档摘要

源文件: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

基于 Hugging Face TRL 的 Prompt 蒸馏

本项目演示Prompt 蒸馏——把带长 prompt 的思考模型中的知识,蒸馏到不带 prompt 的非思考模型中,从而大幅加快响应速度。

🎯 主要目标

把推理能力从以下教师蒸馏到学生:

  • 教师:Qwen3-30B-A3B-Thinking-2507,配合一段详尽的 2000+ token prompt
  • 学生:Qwen3-30B-A3B-Instruct-2507,不带任何 prompt

主要收益:

  • 响应速度大幅提升 —— 没有思考开销,也无需处理长 prompt
  • 💰 推理成本更低 —— 每次请求要处理的 token 更少
  • 🎯 能力保持不变 —— 学生模型学会直接作答,无需显式推理
  • 📦 部署更简单 —— 生产环境无需管理长 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!)

方法论

该方法分为两个阶段:

  1. 数据生成(教师模型):一个思考模型借助详尽的 prompt,生成带显式推理的响应。

    • 教师生成:response = thinking_model(long_prompt, query)
  2. 学生训练(蒸馏):微调一个非思考模型,使其在不带 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 模型。差异在于:

  • 教师:思考模型 + 2000+ token prompt → 慢但准
  • 学生:非思考模型 + 无 prompt → 快而直接

不是在做模型规模压缩,而是消除思考开销与 prompt 处理以加快推理。

数据集

本项目使用与 tinker 相同的多语言语言分类任务:

  • 任务:把文本分类为 13 个语言标签
  • 标签ar, de, el, en, es, fr, hi, ru, tr, ur, vi, zh, ot
  • 源数据example-data/multilingual.txt(2101 句)
  • Prompt:详尽的语言分类规则(与 tinker 相同)

安装

前置条件

  1. 安装所需依赖:
pip install -r requirements.txt
  1. 配置 Weights & Biases 以监控训练:
# 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。

系统要求

  • Python 3.10+
  • PyTorch 2.0+
  • CUDA 12.1+(用于 GPU 加速)
  • GPU:H100 80GB(跑 30B 模型),或任意 24GB+ 显存的 GPU 跑更小模型
  • 显存:30B 模型配 LoRA 约需 70-75GB

用法

第 1 步:生成训练数据

用教师模型生成 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)

它会:

  • 从多语言数据集中加载句子
  • 用教师模型配合完整 prompt 生成语言标签
  • 以 JSONL 格式保存训练数据

输出格式:

{ "messages": [ {"role": "user", "content": "Text in some language"}, {"role": "assistant", "content": "en"} ] }

第 2 步:训练学生模型

用 TRL 在蒸馏数据上微调学生模型:

# Single GPU training (recommended - simpler and works reliably) bash train_trl.sh

监控训练:

  • 训练进度默认记录到 Weights & Bases(wandb)
  • https://wandb.ai 查看实时指标
  • 追踪:loss、学习率、吞吐量、GPU 利用率
  • 每一步都会记录,便于细致监控

如需关闭 wandb 日志:

python train_sft_trl.py --report_to none ...other args...

第 3 步:评估模型

训练完成后,评估蒸馏后模型的表现:

# 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_trl
  • 基础模型:Qwen/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:
评估结果会保存:

  • 混淆矩阵:同时以字典和二维数组形式
  • 全部语言:每种语言的完整统计
  • 错误分析:每种语言的错误示例
  • 分语言准确率:从差到好排序

第 4 步:量化前后对比(离线,无需 GPU)

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

它会给出三项指标——全部来自真实数据,不做任何估算:

  1. 输入成本 —— 教师每次调用都要支付完整的分类 prompt;学生只支付原始文本。
  2. 任务质量 —— 学生与教师标签的一致率(蒸馏保真度),从 evaluation_results.json 读取。
  3. 逐案表 —— 若干真实样本并排展示(教师 token / 学生 token / 教师标签 / 学生预测 / 是否一致)。

实测结果(本仓库数据,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/

为什么采用这个方案?

思考模型 → 非思考模型

本实验的主要创新,是从思考模型蒸馏到非思考模型

  1. 思考模型(教师)

    • Qwen3-30B-A3B-Thinking-2507
    • 使用显式推理:<thinking>...</thinking>
    • 需要带详尽指令的长 prompt
    • 较慢但更准
  2. 非思考模型(学生)

    • Qwen3-30B-A3B-Instruct-2507
    • 没有思考标签,直接作答
    • 生产环境无需 prompt
    • 推理速度快 20-30 倍

为什么选 TRL 而不是 verl?

本实现使用 Hugging Face TRL,原因如下:

  1. 更普及:TRL 在社区中被广泛采用
  2. 文档更全:有大量文档与示例
  3. 配置更简单:无需把 JSONL 转 Parquet
  4. 标准工作流:与 HuggingFace 生态无缝衔接
  5. 更易调试:错误信息清晰、工具链更好

在配合 LoRA 做监督微调上,TRL 提供的能力相当,但 API 更为友好。

关键实现细节

数据格式

训练数据使用 TRL/Transformers 所期望的标准对话格式:

{ "messages": [ {"role": "user", "content": "Text to classify"}, {"role": "assistant", "content": "language_code"} ] }

TRL 会自动:

  • 应用模型的对话模板
  • 对格式化后的文本做分词
  • 生成正确的 loss 掩码(只在 assistant 响应上训练)

训练配置

  • 框架:Hugging Face TRL SFTTrainer
  • LoRA:应用于所有线性层以节省显存
  • 梯度检查点:开启以节省显存
  • 混合精度:bfloat16,在现代 GPU 上训练更快

与 Tinker 的对比

本实现紧跟 tinker cookbook 的方法论,并有一处关键增强:

相同:

  • 教师模型:Qwen3-30B-A3B-Thinking(与 tinker 相同)
  • LoRA 配置:rank 32、alpha 16
  • 学习率:2e-4
  • 训练轮数:1
  • 温度:0.15(数据生成)
  • Prompt:完全相同的语言分类 prompt

增强:

  • 学生模型:Qwen3-30B-A3B-Instruct(非思考变体)
    • 移除思考开销以加快推理
    • 同等模型规模,但不带推理 token 直接作答
    • 生产环境比思考模型快 20-30 倍
  • 框架:TRL(比 tinker 内部框架更易上手)
  • 最大长度:2048(学生不需要长上下文)

为什么这样更好:

  • 原 tinker 方案:只蒸馏 prompt
  • 我们的方案:同时蒸馏 prompt 与思考过程
  • 结果:推理大幅加速且质量无损

预期结果

训练完成后,学生模型(Qwen3-30B-A3B-Instruct)应当能:

  • 无需 2000+ token 的详尽 prompt 即可分类语言
  • ✅ 取得与教师模型(思考 + prompt)相当的准确率
  • ✅ 响应快 20-30 倍(无思考过程、无 prompt 处理)
  • ✅ 每次请求占用更少内存(上下文更短)
  • ✅ 推理成本更低(要处理的 token 更少)

输入成本对比(实测,非估算):

运行 python compare.py 即可在本仓库数据上复现真实数字。使用默认的
tiktoken o200k_base 计数器,学生每次调用处理的输入 token 约少 40 倍
(984.9 → 24.7,降幅 97.5%),同时保留 95.19% 与教师标签的一致率。
完整明细见上文 用法 → 第 4 步 下的表格。

这使得蒸馏模型在输入成本敏感的生产部署中颇具吸引力。注意:墙上时钟延迟取决于
推理栈与硬件,必须在 GPU 上实测——本 README 不引用任何编造的延迟数字。

故障排查

显存不足(OOM)

30B 模型需要大显存的 H100 GPU(80GB)。若遇到 OOM:

解决方案:

  1. per_device_train_batch_size 从 4 降到 2 或 1
  2. max_length 从 2048 降到 1024 或 512
  3. 增大 gradient_accumulation_steps 以维持等效 batch size
  4. lora_rank 从 32 降到 16 或 8

备选:使用更小的模型
如果没有 80GB GPU,可换用更小的模型:

  • Qwen2.5-7B-Instruct:约 28GB 内存,大多数 GPU 都能放下
  • Qwen2.5-14B-Instruct:约 50GB 内存,可放在 A100/H100 上
  • 只需在训练脚本里改 --model_name

显存需求:

  • 30B 模型:约 70-75GB(需 H100 80GB)
  • 14B 模型:约 40-50GB(可放在 A100 40GB 或 H100 上)
  • 7B 模型:约 25-30GB(大多数 GPU 都能放下)

数据生成问题

若数据生成失败或很慢:

  1. 增大 tensor_parallel_size 以使用更多 GPU
  2. H100x8 可用并行脚本:bash create_data_h100x8.sh
  3. 测试时缩小数据集规模
  4. nvidia-smi 检查 GPU 显存使用

训练不收敛

若模型学不到东西:

  1. 确认训练数据格式正确
  2. 检查样本是否带有合法的语言标签
  3. 尝试增加训练轮数
  4. 调整学习率(试试 5e-5 或 2e-4)

引用

如果你使用了本代码,请引用原始论文:

@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)。

致谢

  • 原始 tinker cookbook 实现
  • Hugging Face TRL 框架
  • 阿里云的 Qwen 模型家族

发布者: 作者: bojieli 转发
评论区 (0)
U