ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

17.知识蒸馏

17.知识蒸馏 1.知识蒸馏Distillation用大模型教师 Teacher去训练一个更小的模型学生 Student把大模型的思考方式、输出风格迁移给小模型得到一个更快、更省显存的轻量模型类比学霸大教师不直接把脑子复制给学徒而是大量做题把自己的思考过程、倾向性、答案写成笔记学徒拿着这套笔记重新学习学徒本身规模很小但学会学霸的做题习惯。2. 核心概念1)教师模型 Teacher大参数量、效果好、推理慢、吃显存蒸馏阶段一般冻结只做推理输出样本不参与参数更新。2)学生模型 Student参数量小新模型要训练更新参数目标尽量复刻老师的行为。3)软标签 Soft Label / 暗知识 Dark Knowledge普通训练只用硬标签只有标准答案。蒸馏用老师输出的完整概率分布不只看最终答案还看老师认为各个备选答案的可能性里面藏着模型的推理倾向叫暗知识。4)蒸馏损失学生一边对齐真实标准答案一边对齐老师输出的概率分布两者加权一起训练学生模型。3. 标签标签就是给每个候选 token 的目标分数3.1 硬标签目标值只有 0 或者 1。例子正确 token 是猫猫:1 狗:0 鸟:0目标输出尽量变成 [1,0,0]。 这里的标签就是一组目标分数非 0 即 1。3.2 软标签目标是一堆小数概率不是只有 0 和 1猫:0.88 狗:0.11 鸟:0.01目标输出尽量贴近 [0.88,0.11,0.01]。 这组小数就是软标签。4. LLM 行业里两种常见蒸馏4.1 Logit 蒸馏原生知识蒸馏学生直接学老师输出 token 的完整概率分布对训练算力要求高效果上限高。LLM 每一步要选下一个 token词汇表有几万候选词模型给每一个候选词打一个原始分数这个分数就叫 logit。4.2 输出蒸馏行业最常用也叫 SFT 蒸馏拿一堆 prompt调用教师 API 拿到回答把promptteacher回答做成 SFT 数据集拿这份数据集微调学生模型。现在很多开源小模型就是这么做的调用 DeepSeek、GPT‑4 大量生成样本拿样本训自己的 7B/14B 小底座。 这是学输出文本不学原始 logit 概率工程简单但上限低于原生 logit 蒸馏。SFT Supervised Fine‑Tuning 监督微调用一批「用户问题 理想回答」的标注数据继续微调预训练大模型让模型学会按对话格式输出。5.蒸馏 vs 量化对比点蒸馏 Distillation量化 Quantization做什么训练出新的更小模型知识迁移不改模型结构把参数的比特位降低FP16→INT4/INT8是否重新训练必须训练学生模型PTQ 后量化不需要训练QAT 量化感知训练才要训改变参数量可以从 70B→7B参数量直接变小参数量不变只是每个数字存得更粗糙目的把大模型能力教给小模型同一个模型压缩体积、加速推理关系经常组合先蒸馏得到小模型再对小模型做 INT4 量化进一步压缩部署举个端侧 AI 部署链路例子 70B 大模型 →【蒸馏】→7B 学生模型 →【INT4 量化】→本地端侧可跑的小权重文件。6.输出蒸馏介绍输出蒸馏也叫硬样本蒸馏、API 蒸馏DeepSeek‑R1‑Distill、很多开源小推理模型都是这套流程流程准备一批 prompt业务数据 / 公开数据集调用教师 API生成完整回答包含 CoT 思考过程清洗、过滤、去重做成 prompt→response SFT 数据集拿这份数据集对学生底座做 SFT 微调。输出蒸馏代码# 硬样本蒸馏链路相关代码 import json import os from openai import OpenAI # ---------- 配置教师APIDeepSeek / 豆包都兼容openai格式 ---------- client OpenAI( api_keyos.environ[DEEPSEEK_API_KEY], base_urlhttps://api.deepseek.com/v1 ) def load_my_business_prompt(): 加载业务prompt列表可以从txt/json读取这里模拟 return [ 写一段python快速排序, 解释什么是logit, 简单讲什么是知识蒸馏, 写rust实现二分查找 ] prompts load_my_business_prompt() distill_data [] # ---------- 步骤1调用教师API生成蒸馏硬样本 ---------- for q in prompts: resp client.chat.completions.create( messages[{role:user, content: q}], modeldeepseek-chat, temperature0.7 ) ans resp.choices[0].message.content distill_data.append({ messages: [ {role:user, content: q}, {role:assistant, content: ans} ] }) # ---------- 步骤2简单清洗过滤保存数据集json ---------- def filter_sample(sample): 简单过滤逻辑回答为空、过短直接丢弃业务可再加更多校验 assistant_text sample[messages][1][content] if not assistant_text or len(assistant_text.strip()) 10: return False return True distill_data [s for s in distill_data if filter_sample(s)] with open(hard_distill_dataset.json, w, encodingutf‑8) as f: json.dump(distill_data, f, ensure_asciiFalse, indent2) print(f生成蒸馏数据集完成共 {len(distill_data)} 条保存 hard_distill_dataset.json) # ---------- 步骤3调用 LLaMA‑Factory 执行SFT训练硬样本蒸馏 ---------- 注意llama‑factory是命令行工具不在python内直接import跑 两种方式 A) 脚本调用subprocess拉起llamafactory-cli命令下面示例 B) 终端直接执行llamafactory-cli train sft_distill.yaml import subprocess cmd [ llamafactory-cli, train, sft_distill.yaml ] print(开始执行LLaMA‑Factory SFT蒸馏训练...) subprocess.run(cmd, checkTrue) print(训练完成蒸馏后的学生模型输出在配置的 output_dir 目录)输出学生模型位置配置在sft_distill.yaml### sft_distill.yaml model_name_or_path: Qwen/Qwen2‑1.5B‑Instruct dataset: hard_distill_dataset dataset_dir: ./ template: qwen finetuning_type: lora lora_target: all stage: sft do_train: true overwrite_cache: true cutoff_len: 1024 max_samples: null per_device_train_batch_size: 2 gradient_accumulation_steps: 2 learning_rate: 5e‑5 num_train_epochs: 3 logging_steps: 10 save_steps: -1 save_total_limit: 2 output_dir: ./student_hard_distill_lora bf16: true fp16: false optim: paged_adamw_8bit val_size: 0.0 report_to: none训练结束后导出完整权重代码# 训练完合并LoRA导出完整学生模型 cmd_merge [ llamafactory-cli, export, --model_name_or_path, Qwen/Qwen2‑1.5B‑Instruct, --adapter_name_or_path, ./student_hard_distill_lora, --export_dir, ./student_distilled_full, --export_size, 2 ] subprocess.run(cmd_merge, checkTrue)
返回列表