## English # Prompt Distillation with Hugging Face TRL This project demonstrates **prompt distillation** - a technique to distill knowledge from a **thinking model with long prompts** into a **non-thinking model without prompts**, making responses dramatically faster. ## 🎯 Main Goal Distill the reasoning capability from: - **Teacher**: Qwen3-30B-A3B-**Thinking**-2507 with a detailed 2000+ token prompt - **Student**: Qwen3-30B-A3B-**Instruct**-2507 without any prompt **Key Benefits:** - ⚡ **Much faster response time** - No thinking overhead, no long prompt processing - 💰 **Lower inference cost** - Fewer tokens to process per request - 🎯 **Same capability** - Student model learns to respond directly without explicit reasoning - 📦 **Easier deployment** - No need to manage long prompts in production ## What is Prompt Distillation? Prompt Distillation (also known as **context distillation**) is a training method that makes an LLM internalize a long and complex prompt into its parameters. In this experiment, we also remove the thinking overhead by distilling from a thinking model to a non-thinking model. **Example - Language Classification:** We want to internalize this detailed 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)* **Before distillation (Teacher with thinking + prompt):** ```text System: <2000+ token detailed prompt> User: 一生、バンドしてくれる? Assistant: Let me analyze the script... These are Han characters... Based on rule X...ja ⏱️ Response time: ~2-3 seconds ``` **After distillation (Student, no thinking, no prompt):** ```text User: 一生、バンドしてくれる? Assistant: ja ⏱️ Response time: ~0.1 seconds (20-30x faster!) ``` ## Methodology The method involves two stages: 1. **Data Generation (Teacher Model)**: A **thinking model** uses a detailed prompt to generate responses with explicit reasoning. - Teacher generates: `response = thinking_model(long_prompt, query)` 2. **Student Training (Distillation)**: A **non-thinking model** is fine-tuned to predict responses directly without the prompt or thinking process. - Student learns: `non_thinking_model(query) ≈ thinking_model(long_prompt, query)` - Result: Fast, direct responses with internalized reasoning capability ## Hyperparameters This implementation uses **OpenAI Cookbook hyperparameters** (from gpt-oss-20b example): | Parameter | Value | Source | |-----------|-------|--------| | **Teacher Model** | Qwen3-30B-A3B-**Thinking**-2507 | With thinking capability + long prompt | | **Student Model** | Qwen3-30B-A3B-**Instruct**-2507 | Same size, no thinking, no prompt | | **LoRA Rank** | 32 | tinker | | **LoRA Alpha** | 16 | Standard | | **Learning Rate** | 2e-4 | OpenAI | | **LR Schedule** | cosine_with_min_lr | OpenAI | | **Min LR Rate** | 0.1 | OpenAI | | **Batch Size** | 4 per GPU | OpenAI | | **Gradient Accumulation** | 4 steps | OpenAI | | **Max Length** | 2048 | OpenAI (student only needs short context) | | **Num Epochs** | 1 | OpenAI | | **Temperature** | 0.15 | tinker (data generation) | | **Warmup Ratio** | 0.03 | OpenAI | | **Gradient Checkpointing** | True | OpenAI | **Key Design Choice**: We use the same 30B model for both teacher and student. The difference is: - **Teacher**: Thinking model + 2000+ token prompt → Slow but accurate - **Student**: Non-thinking model + no prompt → Fast and direct This is **not** about model size compression, but about **removing thinking overhead and prompt processing** for faster inference. ## Dataset The project uses the same multilingual language classification task as tinker: - **Task**: Classify text into 13 language labels - **Labels**: `ar, de, el, en, es, fr, hi, ru, tr, ur, vi, zh, ot` - **Source Data**: `example-data/multilingual.txt` (2,101 sentences) - **Prompt**: Detailed language classification rules (same as tinker) ## Installation ### Prerequisites 1. Install the required dependencies: ```bash # From the repository root: use a separate Linux/CUDA project-local environment. # The root ch8 extra intentionally excludes this training stack; ordinary # Chapter 8 experiments should not pull TRL/PEFT/vLLM or CUDA-oriented deps. # requirements.txt includes vLLM, which is Linux/GPU-only in this repository. cd chapter8/prompt-distillation python -m venv .venv-prompt-distillation source .venv-prompt-distillation/bin/activate python -m pip install --upgrade pip python -m pip install -r requirements.txt ``` 2. Setup Weights & Biases for training monitoring: ```bash # 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 ``` You can get your API key from [https://wandb.ai/settings](https://wandb.ai/settings) ### System Requirements - Python 3.10+ - PyTorch 2.0+ - CUDA 12.1+ (for GPU acceleration; vLLM path is Linux/GPU) - **GPU**: H100 80GB (for 30B model) or any GPU with 24GB+ for smaller models - **Memory**: ~70-75GB VRAM for 30B model with LoRA ## Usage ### Step 1: Generate Training Data Generate prompt distillation data using the teacher model: ```bash # 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 ``` **Options:** - `--input_file`: Path to input sentences (one per line) - `--output_file`: Where to save generated training data - `--model_name`: Teacher model (Qwen3-30B-A3B-Thinking-2507 for better accuracy) - `--temperature`: Sampling temperature (0.15 matches tinker) - `--tensor_parallel_size`: Number of GPUs for inference (4 recommended) - `--max_retries`: Number of retry attempts for failed samples (default: 3) This will: - Load sentences from the multilingual dataset - Use the teacher model to generate language labels with the full prompt - Save training data in JSONL format **Output format:** ```json { "messages": [ {"role": "user", "content": "Text in some language"}, {"role": "assistant", "content": "en"} ] } ``` ### Step 2: Train the Student Model Fine-tune the student model on the distilled data using TRL: ```bash # Single GPU training (recommended - simpler and works reliably) bash train_trl.sh ``` **Monitoring Training:** - Training progress is logged to **Weights & Biases** (wandb) by default - View real-time metrics at: [https://wandb.ai](https://wandb.ai) - Tracks: loss, learning rate, throughput, GPU utilization - Every step is logged for detailed monitoring **To disable wandb logging:** ```bash python train_sft_trl.py --report_to none ...other args... ``` ### Step 3: Evaluate Your Model After training, evaluate the distilled model's performance: ```bash # 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 ``` **Defaults:** - Model: `./models/prompt_distillation_trl` - Base model: `Qwen/Qwen3-30B-A3B-Instruct-2507` - Test file: `./example-data/multilingual.txt` **Real-time Output Example:** ```text 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 ``` **Saved to JSON:** The evaluation results are saved with: - **Confusion matrix**: Both as dict and 2D array format - **All languages**: Complete statistics for every language - **Error analysis**: Example errors for each language - **Per-language accuracy**: Sorted from worst to best ### Step 4: Quantify the Before/After (offline, no GPU needed) The whole point of prompt distillation is captured in one before/after comparison: the same task, done by the **teacher (long prompt + thinking)** vs the **student (no prompt, direct answer)** — how much input cost is saved, and how much quality is retained. `compare.py` computes this **entirely offline** from the real dataset, the teacher labels, and the evaluation results (no model download, no network): ```bash # 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 ``` It reports three things — all from real data, nothing estimated: 1. **Input cost** — teacher pays the full classification prompt on every call; student pays only the raw text. 2. **Task quality** — the student's agreement rate with the teacher's labels (distillation fidelity), read from `evaluation_results.json`. 3. **Per-case table** — several real samples side by side (teacher tokens / student tokens / teacher label / student prediction / match). **Measured result** (this repo's data, `tiktoken o200k_base` counter): | Dimension | Teacher (long prompt + thinking) | Student (no prompt) | Change | |-----------|----------------------------------|---------------------|--------| | Avg input tokens / call | 984.9 | 24.7 | **−97.5% (≈40× fewer)** | | Total input tokens (2,100 cases) | 2,068,204 | 51,913 | −97.5% | | Task quality (agreement w/ teacher) | 100% (reference) | **95.19%** (1999/2100) | −4.8 pp | On a per-input-token-billed API this input reduction lowers cost roughly proportionally; the teacher additionally spends thinking (CoT) **output** tokens that are not counted here, so the real gap is larger. Wall-clock latency depends on the serving stack and must be measured on GPU — `compare.py` deliberately does **not** fabricate a latency number. Exact token counts vary by tokenizer; pass `--tokenizer` for the student model's own count. ## Project Structure ```text 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/ ``` A book-audited CUDA run is retained under `validation/exp8-8-kimi3-smollm2-20260730/`: SmolLM2-135M-Instruct student trained on Kimi K3 teacher labels (160 train / 80 test rows). The completed campaign retains all 160 training and 80 held-out teacher receipts. Held-out results: teacher 100%, baseline 0%, trained 95%; ~197× latency speedup; ~75% input-token reduction. All eight evidence gates pass; see `manifest.json` for the content-hashed evidence package. ## Why This Approach? ### Thinking Model → Non-Thinking Model The main innovation in this experiment is distilling from a **thinking model** to a **non-thinking model**: 1. **Thinking Model (Teacher)**: - Qwen3-30B-A3B-**Thinking**-2507 - Uses explicit reasoning: `...` - Requires long prompts with detailed instructions - Slower but more accurate 2. **Non-Thinking Model (Student)**: - Qwen3-30B-A3B-**Instruct**-2507 - No thinking tags, direct responses - No prompts needed in production - **20-30x faster inference** ### Why TRL Instead of verl? We use Hugging Face TRL for this implementation because: 1. **More Common**: TRL is widely adopted in the community 2. **Better Documentation**: Extensive docs and examples 3. **Simpler Setup**: No need to convert JSONL to Parquet 4. **Standard Workflow**: Works seamlessly with HuggingFace ecosystem 5. **Easier to Debug**: Clear error messages and better tooling TRL provides the same capabilities for supervised fine-tuning with LoRA, but with a much more user-friendly API. ## Key Implementation Details ### Data Format The training data uses the standard chat format that TRL/Transformers expects: ```json { "messages": [ {"role": "user", "content": "Text to classify"}, {"role": "assistant", "content": "language_code"} ] } ``` TRL automatically: - Applies the model's chat template - Tokenizes the formatted text - Creates proper loss masks (only trains on assistant responses) ### Training Configuration - **Framework**: Hugging Face TRL SFTTrainer - **LoRA**: Applied to all linear layers for memory efficiency - **Gradient Checkpointing**: Enabled to save memory - **Mixed Precision**: bfloat16 for faster training on modern GPUs ## Comparison to Tinker This implementation closely follows the tinker cookbook methodology with a key enhancement: **Same:** - Teacher model: Qwen3-30B-A3B-Thinking (same as tinker) - LoRA configuration: rank 32, alpha 16 - Learning rate: 2e-4 - Training epochs: 1 - Temperature: 0.15 (data generation) - Prompt: Identical language classification prompt **Enhanced:** - **Student model**: Qwen3-30B-A3B-**Instruct** (non-thinking variant) - Removes thinking overhead for faster inference - Same model size, but direct responses without reasoning tokens - 20-30x faster than thinking model in production - **Framework**: TRL (more accessible than tinker's internal framework) - **Max length**: 2048 (student doesn't need long context) **Why This Is Better:** - Original tinker approach: Distill prompt only - Our approach: **Distill both prompt AND thinking process** - Result: Dramatically faster inference with no quality loss ## Expected Results After training, the student model (Qwen3-30B-A3B-Instruct) should: - ✅ Classify languages **without** the 2000+ token detailed prompt - ✅ Achieve similar accuracy to the teacher model (thinking + prompt) - ✅ Respond **20-30x faster** (no thinking process, no prompt processing) - ✅ Use **much less memory** per request (shorter context) - ✅ Lower inference cost (fewer tokens to process) **Input-Cost Comparison (measured, not estimated):** Run `python compare.py` to reproduce the real numbers on this repo's data. With the default `tiktoken o200k_base` counter, the student processes **≈40× fewer input tokens per call** (984.9 → 24.7, a 97.5% reduction) while retaining **95.19%** agreement with the teacher's labels. See the table under *Usage → Step 4* for the full breakdown. This makes the distilled model attractive for **production deployment** where input cost matters. Note: wall-clock latency depends on the serving stack and hardware and must be measured on GPU — this README does not quote a fabricated latency figure. ## Troubleshooting ### Out of Memory (OOM) The 30B model requires a large H100 GPU (80GB). If you encounter OOM errors: **Solutions:** 1. Reduce `per_device_train_batch_size` from 4 to 2 or 1 2. Reduce `max_length` from 2048 to 1024 or 512 3. Increase `gradient_accumulation_steps` to maintain effective batch size 4. Reduce `lora_rank` from 32 to 16 or 8 **Alternative: Use a Smaller Model** If you don't have an 80GB GPU, use a smaller model: - **Qwen2.5-7B-Instruct**: ~28GB memory, fits on most GPUs - **Qwen2.5-14B-Instruct**: ~50GB memory, fits on A100/H100 - Just change `--model_name` in the training script **Memory Requirements:** - 30B model: ~70-75GB (requires H100 80GB) - 14B model: ~40-50GB (fits on A100 40GB or H100) - 7B model: ~25-30GB (fits on most GPUs) ### Data Generation Issues If data generation fails or is slow: 1. Increase `tensor_parallel_size` to use more GPUs 2. Use the parallel script for H100x8: `bash create_data_h100x8.sh` 3. Reduce the dataset size for testing 4. Check GPU memory usage with `nvidia-smi` ### Training Not Converging If the model doesn't learn: 1. Verify training data format is correct 2. Check that examples have valid language labels 3. Try increasing the number of training epochs 4. Adjust the learning rate (try 5e-5 or 2e-4) ## Citation If you use this code, please cite the original papers: ```bibtex @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} } ``` And the Hugging Face TRL library: ```bibtex @software{trl2024, title={TRL: Transformer Reinforcement Learning}, author={TRL contributors}, url={https://github.com/huggingface/trl}, year={2024} } ``` ## License This project follows the same license as the TRL library (Apache 2.0). ## Acknowledgments - Original tinker cookbook implementation - Hugging Face TRL framework - Qwen model family by Alibaba Cloud --- ## 中文 # 使用 Hugging Face TRL 进行快速蒸馏 该项目演示了**即时蒸馏** - 一种将知识从**具有长提示**的思维模型提炼为**无提示的非思维模型**的技术,从而使响应速度显著加快。 ## 🎯 主要目标 从以下内容中提取推理能力: - **老师**:Qwen3-30B-A3B-**Thinking**-2507,详细2000+代币提示 - **学生**:Qwen3-30B-A3B-**Instruct**-2507 无任何提示 **主要优点:** - ⚡ **响应时间更快** - 无需思考开销,无需长时间的提示处理 - 💰 **降低推理成本** - 每个请求处理的令牌更少 - 🎯 **相同的能力** - 学生模型学会直接回应而无需明确的推理 - 📦 **更轻松的部署** - 无需在生产中管理长提示 ## 什么是快速蒸馏? 提示蒸馏(也称为**上下文蒸馏**)是一种训练方法,使大语言模型将长而复杂的提示内化为其参数。在这个实验中,我们还通过从思维模型提炼为非思维模型来消除思维开销。 **示例 - 语言分类:** 我们想要内化这个详细的提示: > “将所提供文本的语言分类为以下标签:ar、de、el、en、es、fr、hi、ru、tr、ur、vi、zh、ot。使用这些规则:梵文脚本 → hi、希腊脚本 → el、西里尔脚本 → ru...” *(2000+ 个标记)* **蒸馏前(老师思考+提示):** ```text System: <2000+ token detailed prompt> User: 一生、バンドしてくれる? Assistant: Let me analyze the script... These are Han characters... Based on rule X...ja ⏱️ Response time: ~2-3 seconds ``` **蒸馏后(学生,无思考,无提示):** ```text User: 一生、バンドしてくれる? Assistant: ja ⏱️ Response time: ~0.1 seconds (20-30x faster!) ``` ## 方法论 该方法涉及两个阶段: 1. **数据生成(教师模型)**:**思维模型**使用详细的提示来生成具有明确推理的响应。 - 教师生成:`response = thinking_model(long_prompt, query)` 2. **学生训练(蒸馏)**:**非思考模型**经过微调,可直接预测响应,无需提示或思考过程。 - 学生学习:`non_thinking_model(query) ≈ thinking_model(long_prompt, query)` - 结果:快速、直接的反应以及内在的推理能力 ## 超参数 此实现使用 **OpenAI Cookbook 超参数**(来自 gpt-oss-20b 示例): |参数|价值|来源 | |-----------|-------|--------| | **教师模型** | Qwen3-30B-A3B-**Thinking**-2507 |具备思考能力+长提示| | **学生模型** | Qwen3-30B-A3B-**Instruct**-2507 |大小相同,无需思考,无需提示 | | **LoRA 排名** | 32 | 32 Tinker | | **洛拉阿尔法** | 16 | 16标准| | **学习率** | 2e-4 | 2e-4 OpenAI | | **LR 时间表** |余弦_with_min_lr | OpenAI | | **最低 LR 率** | 0.1 | 0.1 OpenAI | | **批量大小** |每个 GPU 4 个 | OpenAI | | **梯度累积** | 4 步骤 | OpenAI | | **最大长度** | 2048 | 2048 OpenAI(学生只需要简短的上下文)| | **历元数** | 1 | OpenAI | | **温度** | 0.15 | 0.15 Tinker(数据生成)| | **预热比率** | 0.03 | 0.03 OpenAI | | **梯度检查点** |真实| OpenAI | **关键设计选择**:我们为教师和学生使用相同的 30B 模型。区别在于: - **老师**:思维模型+2000+代币提示→慢而准 - **学生**:无思维模型+无提示→快速直接 这**不是**关于模型大小压缩,而是关于**消除思维开销和提示处理**以加快推理速度。 ## 数据集 该项目使用与tinker相同的多语言语言分类任务: - **任务**:将文本分类为 13 种语言标签 - **标签**:`ar, de, el, en, es, fr, hi, ru, tr, ur, vi, zh, ot` - **源数据**:`example-data/multilingual.txt`(2,101 句) - **提示**:详细的语言分类规则(与tinker相同) ## 安装 ### 先决条件 1. 安装所需的依赖项: ```bash # 从仓库根目录开始:请使用单独的 Linux/CUDA 项目本地环境。 # 根目录 ch8 extra 有意不包含本训练栈;普通第 8 章实验不应拉取 # TRL/PEFT/vLLM 或 CUDA 取向依赖。 # requirements.txt 包含 vLLM;本仓库将 vLLM 视为 Linux/GPU-only 依赖。 cd chapter8/prompt-distillation python -m venv .venv-prompt-distillation source .venv-prompt-distillation/bin/activate python -m pip install --upgrade pip python -m pip install -r requirements.txt ``` 2. 设置训练监控的权重和偏差: ```bash # 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](https://wandb.ai/settings) 获取您的 API 密钥 ### 系统要求 - Python 3.10+ - PyTorch 2.0+ - CUDA 12.1+(用于 GPU 加速;vLLM 路径面向 Linux/GPU) - **GPU**:H100 80GB(适用于 30B 型号)或任何具有 24GB+ 的 GPU(适用于较小型号) - **内存**:带有 LoRA 的 30B 型号约为 70-75GB VRAM ## 用法 ### 第 1 步:生成训练数据 使用教师模型生成即时蒸馏数据: ```bash # 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匹配修补匠) - `--tensor_parallel_size`:用于推理的 GPU 数量(推荐 4 个) - `--max_retries`:失败样本的重试次数(默认值:3) 这将: - 从多语言数据集中加载句子 - 使用教师模型生成带有完整提示的语言标签 - 以 JSONL 格式保存训练数据 **输出格式:** ```json { "messages": [ {"role": "user", "content": "Text in some language"}, {"role": "assistant", "content": "en"} ] } ``` ### 第 2 步:训练学生模型 使用 TRL 根据蒸馏数据微调学生模型: ```bash # Single GPU training (recommended - simpler and works reliably) bash train_trl.sh ``` **监控培训:** - 训练进度默认记录到 **权重和偏差** (wandb) - 查看实时指标:[https://wandb.ai](https://wandb.ai) - 跟踪:损失、学习率、吞吐量、GPU 利用率 - 记录每个步骤以进行详细监控 **禁用 wandb 日志记录:** ```bash python train_sft_trl.py --report_to none ...other args... ``` ### 第 3 步:评估您的模型 训练后,评估蒸馏模型的性能: ```bash # 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` **实时输出示例:** ```text 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) 快速蒸馏的全部要点可以在前后对比中得到体现: 同样的任务,由**老师完成(长提示+思考)** vs **学生 (无提示,直接回答)**——节省了多少投入成本,质量有多少 被保留。 `compare.py` 根据真实数据集**完全离线**计算, 老师标签,以及评估结果(无模型下载,无网络): ```bash # 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. **投入成本** — 教师在每次通话时支付全部分类提示费用; 学生只需支付原始文本费用。 2. **任务质量**——学生对老师标签的同意率 (蒸馏保真度),从 `evaluation_results.json` 读取。 3. **按案例表** — 几个并排的真实样本(教师标记/ 学生标记/教师标签/学生预测/匹配)。 **测量结果**(本仓库的数据,`tiktoken o200k_base`计数器): |尺寸|老师(长提示+思考)|学生(无提示)|改变| |-----------|----------------------------------|---------------------|--------| |平均输入令牌/调用| 984.9 | 24.7 | **−97.5%(≈40×更少)** | |输入令牌总数(2,100 例)| 2,068,204 | 51,913 | −97.5% | |任务质量(与老师达成一致)| 100%(参考)| **95.19%** (1999/2100) | −4.8 个百分点 | 在按输入代币计费的 API 上,这种输入减少大致降低了成本 按比例;老师还额外花费思考(CoT)**输出**代币 这里没有计算在内,所以真正的差距更大。挂钟延迟取决于 在服务堆栈上并且必须在 GPU 上测量 - `compare.py` 故意这样做 **不**捏造延迟数字。确切的令牌计数因令牌生成器而异;经过 `--tokenizer`为学生模特自己算的。 ## 项目结构 ```text 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 - 使用显式推理:`...` - 需要长提示和详细说明 - 更慢但更准确 2. **非思考模型(学生)**: - Qwen3-30B-A3B-**Instruct**-2507 - 没有思考标签,直接回应 - 生产中无需提示 - **推理速度加快 20-30 倍** ### 为什么 TRL 而不是 verl? 我们在此实现中使用 Hugging Face TRL,因为: 1. **更常见**:TRL在社区中被广泛采用 2. **更好的文档**:广泛的文档和示例 3. **更简单的设置**:无需将 JSONL 转换为 Parquet 4. **标准工作流程**:与 HuggingFace 生态系统无缝协作 5. **更容易调试**:清除错误消息和更好的工具 TRL 提供与 LoRA 相同的监督微调功能,但具有更加用户友好的 API。 ## 关键实施细节 ### 数据格式 训练数据使用 TRL/Transformers 期望的标准聊天格式: ```json { "messages": [ {"role": "user", "content": "Text to classify"}, {"role": "assistant", "content": "language_code"} ] } ``` 自动TRL: - 应用模特的聊天模板 - 对格式化文本进行标记 - 创建适当的损失掩模(仅训练助理响应) ### 训练配置 - **框架**:Hugging Face TRL SFTTrainer - **LoRA**:应用于所有线性层以提高内存效率 - **梯度检查点**:启用以节省内存 - **混合精度**:bfloat16 可在现代 GPU 上实现更快的训练 ## 与 Tinker 的比较 此实现紧密遵循 Tinker Cookbook 方法,并进行了关键增强: **相同的:** - 教师模型:Qwen3-30B-A3B-Thinking(与tinker相同) - LoRA配置:等级32,阿尔法16 - 学习率:2e-4 - 训练时期:1 - 温度:0.15(数据生成) - 提示:相同语言分类提示 **增强:** - **学生模型**:Qwen3-30B-A3B-**Instruct**(非思维变体) - 消除思维开销以加快推理速度 - 相同的模型大小,但无需推理标记即可直接响应 - 比生产中的思维模型快 20-30 倍 - **框架**:TRL(比tinker的内部框架更容易访问) - **最大长度**:2048(学生不需要长上下文) **为什么这样更好:** - 原始修补方法:仅提取提示 - 我们的方法:**提炼即时和思考过程** - 结果:推理速度显著加快,且没有质量损失 ## 预期结果 训练后,学生模型 (Qwen3-30B-A3B-Instruct) 应: - ✅ 对语言进行分类 **没有** 2000+ token 详细提示 - ✅ 达到与教师模型相似的准确性(思考+提示) - ✅ 响应速度**20-30 倍**(无需思考过程,无需提示处理) - ✅ 每个请求使用**更少的内存**(更短的上下文) - ✅ 降低推理成本(需要处理的代币更少) **投入成本比较(测量,而非估计):** 运行 `python compare.py` 来重现此存储库数据上的实数。随着 默认`tiktoken o200k_base`计数器,学生处理**≈40×更少的输入 每次调用代币**(984.9 → 24.7,减少 97.5%),同时保留 **95.19%** 与老师的标注一致。请参阅*用法 → 步骤 4* 下的表格了解 完整明细。 这使得蒸馏模型对于**生产部署**具有吸引力,其中输入 成本很重要。注意:挂钟延迟取决于服务堆栈和硬件, 必须在 GPU 上测量——本自述文件并未引用捏造的延迟数字。 ## 故障排除 ### 内存不足 (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`以保持有效的批量大小 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. 尝试增加训练epoch数 4. 调整学习率(尝试5e-5或2e-4) ## 引文 如果您使用此代码,请引用原始论文: ```bibtex @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 库: ```bibtex @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模型家族