知识蒸馏如何压缩大模型体积
知识蒸馏(Knowledge Distillation)是将大模型(Teacher)的推理能力迁移到小模型(Student)的技术路线。核心思路是让Student不仅学习真实标签,还模仿Teacher在logits层输出的软标签分布,捕获类别间的相似关系。这种软标签携带的暗知识(Dark Knowledge)远比one-hot硬标签信息量大,是小模型超越同参数量从零训练模型的关键。
Hinton在2015年提出蒸馏框架时定义了损失函数:
L = alpha * KL(softmax(z_t/T), softmax(z_s/T)) + (1-alpha) * CE(y, z_s)
其中z_t和z_s分别是Teacher和Student的logits输出,T是温度参数,控制软标签的平滑程度。T越高,分布越均匀,暗知识越容易被Student捕获。alpha通常取0.5到0.9之间,平衡蒸馏损失和真实标签损失。
温度参数对蒸馏效果的影响机制
温度T的取值直接影响Student学到的信息量。当T=1时,蒸馏退化为标准训练;当T趋近无穷时,softmax输出接近均匀分布,所有类别的概率差异被抹平。实际工程中T的调优范围在2-20之间,需要根据Teacher模型的logits分布特征进行调整。
实验数据表明,在ImageNet分类任务上,T=4时ResNet-18作为Student从ResNet-152蒸馏获得的精度提升最显著,Top-1准确率从69.8%提升至72.3%。T过低时暗知识传递不充分,T过高则噪声占比上升,Student容易被误导。
import torch
import torch.nn.functional as F
def distillation_loss(student_logits, teacher_logits, labels, T=4, alpha=0.7):
soft_loss = F.kl_div(
F.log_softmax(student_logits / T, dim=1),
F.softmax(teacher_logits / T, dim=1),
reduction="batchmean"
) * (T * T)
hard_loss = F.cross_entropy(student_logits, labels)
return alpha * soft_loss + (1 - alpha) * hard_loss
多阶段蒸馏策略与中间层对齐
单纯的logits层蒸馏对深层Transformer模型效果有限,因为大模型的中间隐藏层包含大量结构性知识。多阶段蒸馏在logits蒸馏之外引入中间层特征对齐,强制Student的特定层输出接近Teacher对应层的分布。
具体做法是选择Teacher的若干中间层作为hint layers,Student对应层作为guided layers,用均方误差或余弦相似度作为对齐损失:
def feature_alignment_loss(student_features, teacher_features, proj_weight, proj_bias):
student_proj = student_features @ proj_weight + proj_bias
teacher_norm = F.normalize(teacher_features, dim=-1)
student_norm = F.normalize(student_proj, dim=-1)
return F.mse_loss(student_norm, teacher_norm.detach())
投影层(Adaptor)的作用是将Student的隐藏维度映射到与Teacher一致,避免维度不匹配。这个投影层只在训练阶段存在,推理时移除,不增加Student的推理开销。
TinyLLM蒸馏实战:7B到1.5B的能力迁移
以Llama系列为蓝本,将7B参数量的Teacher模型蒸馏到1.5B的Student模型。训练流程分为三个阶段:
第一阶段用大规模无标注语料做软标签蒸馏,Teacher对每条文本生成logits,Student在离线蒸馏模式下学习。这种离线蒸馏方式不需要实时调用Teacher,训练速度取决于磁盘IO和Student的前向计算。
第二阶段用高质量指令数据做在线蒸馏,Teacher和Student同步前向计算,动态生成软标签。在线蒸馏的梯度更新更及时,对齐效果更好,但显存开销翻倍。
第三阶段用DPO或RLHF做对齐微调,让Student的输出风格与Teacher保持一致。这个阶段不涉及蒸馏,而是用Teacher生成的偏好数据直接训练Student。
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch, json
teacher = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7b-chat-hf",
torch_dtype=torch.float16,
device_map="auto"
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-7b-chat-hf")
def generate_soft_labels(text, T=4):
inputs = tokenizer(text, return_tensors="pt").to(teacher.device)
with torch.no_grad():
outputs = teacher(**inputs)
logits = outputs.logits / T
soft_probs = torch.softmax(logits, dim=-1)
return soft_probs.cpu().numpy()
蒸馏模型部署与推理加速
蒸馏后的1.5B模型在精度下降可控的前提下,推理速度提升4-5倍,显存占用从14GB降至3GB左右。配合INT4量化可以进一步压缩到1GB以内,在消费级GPU甚至CPU上流畅运行。
蒸馏模型的瓶颈通常不在计算而在内存带宽。对于batch推理场景,KV Cache的显存管理是核心问题。PagedAttention将KV Cache分页管理,按需分配和释放,避免预分配带来的显存浪费,在vLLM框架中已成为标配。
from vllm import LLM, SamplingParams
llm = LLM(
model="./distilled-1.5b",
tensor_parallel_size=1,
gpu_memory_utilization=0.9,
enforce_eager=True
)
params = SamplingParams(temperature=0.7, top_p=0.9, max_tokens=512)
outputs = llm.generate(["用Python写一个快速排序"], [params])
实际测试中,蒸馏模型在代码生成和文本总结任务上的BLEU/ROUGE指标与7B模型差距在5%以内,但推理吞吐量提升3.8倍。在边缘部署场景下,这种精度与速度的权衡是合理且必要的。
原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/da-mo-xing-zhi-shi-zheng-liu-ji-shu-yuan-li-yu-tinyllm-qing/