目录
一、为什么需要模型蒸馏?
1.1 一句话总结
模型蒸馏 = 让一个"小学霸"(小模型)跟着"老教授"(大模型)学习,用小模型的身体继承大模型的智慧。
1.2 现实痛点
假设你用 BERT-base(1.1亿参数)训练了一个新闻标题分类器,效果很好(准确率 95%)。但当你想部署上线时,问题来了:
| 问题 | 具体表现 |
|---|---|
| 推理慢 | BERT 单条推理约 10-30ms,高并发时扛不住 |
| 显存大 | 模型本身 ~440MB,加上 KV Cache 更夸张 |
| 部署难 | 手机端、嵌入式设备根本跑不动 |
蒸馏就是解决方案:训练一个 BiLSTM(可能只有几百万参数),让它学习 BERT 的"思考方式",最终得到一个又快又准的小模型。
1.3 蒸馏 vs 其他压缩方法
模型压缩技术全景:
├── 模型蒸馏(Knowledge Distillation) ← 本文重点
│ └── 用大模型的输出指导小模型训练
├── 模型剪枝(Pruning)
│ └── 去掉不重要的权重/神经元
├── 模型量化(Quantization)
│ └── FP32 → INT8,降低精度
└── 权重共享(Weight Sharing)
└── 多个参数共享同一组权重
🎯 面试考点:蒸馏是唯一一种"跨架构"的压缩方法——教师和学生的模型结构可以完全不同(如 BERT → BiLSTM),而剪枝和量化通常只能压缩同架构模型。
二、核心概念:教师-学生框架
2.1 两个角色
┌─────────────────┐ 知识传递 ┌─────────────────┐
│ 教师模型 │ ─────────────────→ │ 学生模型 │
│ (Teacher) │ │ (Student) │
│ │ │ │
│ BERT-base │ 软标签/特征/注意力 │ BiLSTM │
│ 1.1亿参数 │ │ ~200万参数 │
│ 准确率 95% │ │ 目标准确率 92%+ │
│ 推理慢、大 │ │ 推理快、小 │
└─────────────────┘ └─────────────────┘
2.2 硬标签 vs 软标签(最核心概念!)
这是理解蒸馏的关键钥匙:
硬标签(Hard Label):就是数据集里的标准答案,非 0 即 1。
新闻标题:"华为发布新款手机"
硬标签 = [0, 0, 1, 0, 0] # 科技=1,其他=0
# 只告诉你"这是科技新闻",别的一概不说
软标签(Soft Label):教师模型输出的概率分布,包含丰富的"暗知识"。
新闻标题:"华为发布新款手机"
教师模型(BERT)输出 = [0.02, 0.05, 0.78, 0.12, 0.03]
# 对应类别:[财经, 娱乐, 科技, 数码, 体育]
🔑 暗知识(Dark Knowledge):软标签中非目标类别的概率也包含重要信息!
- 数码类概率 0.12 > 体育类 0.03 → 说明"华为手机"和"数码"有关系,但和"体育"没关系
- 这种类别间的相似性关系,硬标签完全无法提供!
2.3 通俗比喻
考试场景:
硬标签学习(传统训练):
老师只说:"答案是C" → 学生只记住了C
软标签学习(蒸馏训练):
老师说:"答案最可能是C(78%),其次可能是D(12%),
不太可能是B(5%),基本不可能是A(2%)和E(3%)"
→ 学生不仅知道答案,还理解了各选项之间的关系!
三、三大关键公式(面试必考)
公式一:带温度的 Softmax ⭐⭐⭐
这是蒸馏的灵魂公式。普通的 softmax 输出太"尖锐"(概率集中在某一类),无法体现类别间关系。引入温度参数 T 来"软化"概率分布:
$$q_i = \frac{\exp(z_i / T)}{\sum_j \exp(z_j / T)}$$其中:
- $z_i$ 是模型输出的 logit(softmax 之前的原始分数)
- $T$ 是温度参数(Temperature),$T \geq 1$
- $q_i$ 是软化后的概率
温度 T 的效果图解:
假设 logits = [2.0, 1.0, 0.1] 对应 [科技, 数码, 体育]
T=1(正常softmax): [0.66, 0.24, 0.10] ← 太尖锐,暗知识少
T=3(软化后): [0.44, 0.33, 0.23] ← 更平滑,暗知识丰富!
T=20(过度软化): [0.35, 0.34, 0.31] ← 太均匀,信息被稀释
🎯 面试考点:
- T=1 时,就是标准 softmax,等价于普通分类
- T 越大,概率分布越"平坦",非目标类别的概率越大,暗知识越丰富
- T 太大也不行,信息会被过度稀释,一般 T ∈ [2, 10],常用 3~5
公式二:KL 散度(蒸馏损失) ⭐⭐⭐
KL 散度衡量两个概率分布之间的"距离"。在蒸馏中,用来衡量学生软标签和教师软标签的差异:
$$D_{KL}(p \| q) = \sum_i p_i \cdot \log\frac{p_i}{q_i}$$在蒸馏场景中:
- $p_i$ = 教师模型的软标签概率(用温度 T 计算)
- $q_i$ = 学生模型的软标签概率(用温度 T 计算)
蒸馏损失(实际用的形式,乘以 $T^2$):
$$\mathcal{L}_{distill} = T^2 \cdot D_{KL}(p^{teacher} \| q^{student})$$🎯 面试考点:为什么要乘 $T^2$?
这是面试最高频的追问之一!原因如下:
当 T 很大时,soft target 的梯度会按 $1/T^2$ 缩小。如果不乘 $T^2$,蒸馏损失的梯度会远小于硬标签损失的梯度,导致蒸馏"不起作用"。乘 $T^2$ 是为了让两种损失的梯度量级一致,保证蒸馏损失和硬标签损失能平等地贡献梯度。
数学推导:soft target 对 logit 的梯度 $\propto \frac{1}{T^2}$,所以需要乘 $T^2$ 来补偿。
公式三:总损失函数 ⭐⭐⭐
学生模型的最终损失 = 硬标签损失 + 蒸馏损失 的加权组合:
$$\mathcal{L}_{total} = \alpha \cdot \mathcal{L}_{CE}(y, q^{student}) + (1 - \alpha) \cdot \mathcal{L}_{distill}(p^{teacher}, q^{student})$$展开写就是:
$$\mathcal{L}_{total} = \alpha \cdot CE(y, \text{softmax}(z^s)) + (1-\alpha) \cdot T^2 \cdot D_{KL}\left(\text{softmax}\left(\frac{z^t}{T}\right) \Big\| \text{softmax}\left(\frac{z^s}{T}\right)\right)$$其中:
- $y$ = 真实的硬标签(one-hot)
- $z^s$ = 学生模型的 logits
- $z^t$ = 教师模型的 logits
- $\alpha$ = 平衡系数,通常取 0.1~0.5(蒸馏损失权重更大)
总损失示意图:
┌──────────────────────────────────┐
│ 总损失 L_total │
│ = α × L_hard + (1-α) × L_soft │
└──────────┬───────────┬────────────┘
│ │
┌──────────▼───┐ ┌───▼──────────────┐
│ 硬标签损失 │ │ 蒸馏损失(软标签) │
│ L_hard │ │ L_soft │
│ │ │ │
│ 学生输出 │ │ 学生软输出 │
│ vs │ │ vs │
│ 真实标签 │ │ 教师软输出 │
│ │ │ │
│ T=1,普通CE │ │ T>1,KL散度×T² │
│ 学"正确答案" │ │ 学"思考方式" │
└──────────────┘ └───────────────────┘
🎯 面试考点:
- 硬标签部分用 T=1 的标准 softmax
- 蒸馏部分用 T>1 的软化 softmax
- 两部分的 softmax 温度不同,这是很多人容易搞混的点
四、蒸馏的完整流程
Step 1: 训练教师模型(或加载预训练的教师模型)
┌─────────┐ 训练数据 ┌────────────┐
│ 原始数据 │ ──────────→ │ BERT 训练 │──→ 教师模型 ✓
└─────────┘ └────────────┘
Step 2: 用教师模型生成软标签
┌─────────┐ 推理(不更新) ┌────────────┐
│ 训练数据 │ ────────────→ │ BERT 推理 │──→ 软标签
└─────────┘ └────────────┘
Step 3: 训练学生模型(蒸馏训练)
┌─────────┐ ┌────────────┐
│ 训练数据 │ ────────────→ │ BiLSTM │──→ 学生输出
└─────────┘ └────────────┘
│ │
│ 硬标签 │ 学生 logits
▼ ▼
┌─────────┐ ┌────────────┐
│ CE Loss │ │ KL Loss │
│(T=1) │ │(T>1, ×T²) │
└────┬────┘ └─────┬──────┘
│ │
│ 教师软标签 ──────────┘
│ │
└────┬────┘
▼
┌───────────┐
│ L_total │
│ 反向传播 │──→ 更新学生模型参数
└───────────┘
⚠️ 重要:蒸馏过程中,教师模型参数冻结,不更新!只有学生模型在训练。
五、实战代码:BERT → BiLSTM 新闻标题多分类蒸馏
5.0 环境准备
# 安装依赖
# pip install torch transformers numpy pandas scikit-learn
5.1 数据准备(模拟新闻标题数据)
import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from transformers import BertTokenizer, BertForSequenceClassification
import numpy as np
# ============================================================
# 【第一步】准备数据 —— 模拟新闻标题多分类数据集
# ============================================================
# 定义5个新闻类别
LABEL_NAMES = ["科技", "财经", "娱乐", "体育", "教育"]
NUM_CLASSES = len(LABEL_NAMES)
# 模拟一些新闻标题数据(实际项目中从CSV/JSON读取)
# 每条数据 = (新闻标题文本, 类别标签ID)
raw_data = [
("华为发布新一代麒麟芯片性能大幅提升", 0), # 科技
("央行宣布降准0.5个百分点释放长期资金", 1), # 财经
("春节档电影票房突破50亿创新高", 2), # 娱乐
("中国队在世锦赛中获得三枚金牌", 3), # 体育
("教育部发布新课标改革方案", 4), # 教育
("苹果新iPhone搭载A20芯片拍照再升级", 0), # 科技
("A股三大指数集体上涨成交额破万亿", 1), # 财经
("某明星演唱会全国巡演一票难求", 2), # 娱乐
("NBA季后赛湖人队逆转取胜晋级", 3), # 体育
("高考改革新政策综合素质评价权重增加", 4), # 教育
("量子计算突破新型超导材料问世", 0), # 科技
("比特币价格突破新高引发市场热议", 1), # 财经
("国产科幻大片获国际电影节最佳影片提名", 2), # 娱乐
("世界杯预选赛国足关键一战", 3), # 体育
("双一流高校新增人工智能本科专业", 4), # 教育
# ... 实际项目中应有数千~数万条数据
]
# ============================================================
# 【关键概念】自定义Dataset类
# ============================================================
class NewsTitleDataset(Dataset):
"""
新闻标题数据集
- 训练阶段:使用BERT的tokenizer将文本转为input_ids和attention_mask
- 蒸馏阶段:学生模型(BiLSTM)也需要用同一套tokenizer
"""
def __init__(self, data, tokenizer, max_len=32):
"""
参数:
data: list of (text, label) 元组
tokenizer: BERT的分词器(教师和学生共用)
max_len: 文本最大长度(截断或填充到此长度)
"""
self.data = data
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
text, label = self.data[idx]
# 使用BERT tokenizer编码文本
# return_tensors='pt' 返回PyTorch张量
# padding='max_length' 填充到max_len
# truncation=True 超过max_len则截断
encoding = self.tokenizer(
text,
max_length=self.max_len,
padding='max_length',
truncation=True,
return_tensors='pt' # 返回 PyTorch tensor
)
return {
'input_ids': encoding['input_ids'].squeeze(0), # shape: [max_len]
'attention_mask': encoding['attention_mask'].squeeze(0), # shape: [max_len]
'label': torch.tensor(label, dtype=torch.long) # shape: [] (标量)
}
# 加载 BERT tokenizer(中文模型)
# 注意:实际使用时请确保能下载到模型,或使用本地路径
print("正在加载 BERT tokenizer...")
tokenizer = BertTokenizer.from_pretrained('bert-base-chinese')
# 创建数据集和数据加载器
dataset = NewsTitleDataset(raw_data, tokenizer, max_len=32)
dataloader = DataLoader(dataset, batch_size=4, shuffle=True)
print(f"数据集大小: {len(dataset)}")
print(f"类别数: {NUM_CLASSES}")
5.2 训练教师模型(BERT)
# ============================================================
# 【第二步】训练教师模型 —— BERT-base 做新闻标题多分类
# ============================================================
print("\n" + "="*60)
print("第二步:训练教师模型 (BERT)")
print("="*60)
# 加载预训练的BERT模型,num_labels指定分类数
# BertForSequenceClassification = BERT + 一个线性分类头
# 它的输出是 logits(未经softmax的原始分数),这对蒸馏非常重要!
teacher_model = BertForSequenceClassification.from_pretrained(
'bert-base-chinese',
num_labels=NUM_CLASSES # 5个新闻类别
)
# 设置设备(有GPU用GPU,没有就用CPU)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
teacher_model.to(device)
# 定义优化器和损失函数
teacher_optimizer = optim.AdamW(teacher_model.parameters(), lr=2e-5)
# CrossEntropyLoss 内部自动做了 softmax + 负对数似然
# 所以输入是 logits,不是概率!
teacher_criterion = nn.CrossEntropyLoss()
# 训练教师模型
NUM_TEACHER_EPOCHS = 5 # 实际项目中可能需要更多轮次
teacher_model.train() # 设置为训练模式
for epoch in range(NUM_TEACHER_EPOCHS):
total_loss = 0
correct = 0
total = 0
for batch in dataloader:
# 将数据移到GPU/CPU
input_ids = batch['input_ids'].to(device)
attention_mask = batch['attention_mask'].to(device)
labels = batch['label'].to(device)
# 前向传播:BERT输出logits
# outputs.logits 的 shape = [batch_size, num_classes]
outputs = teacher_model(
input_ids=input_ids,
attention_mask=attention_mask
)
logits = outputs.logits # 注意:是logits,不是概率!
# 计算损失(CrossEntropyLoss 内部会做softmax)
loss = teacher_criterion(logits, labels)
# 反向传播 + 更新参数
teacher_optimizer.zero_grad() # 清空梯度
loss.backward() # 计算梯度
teacher_optimizer.step() # 更新参数
# 统计
total_loss += loss.item()
predictions = torch.argmax(logits, dim=1) # 取最大logit作为预测
correct += (predictions == labels).sum().item()
total += labels.size(0)
accuracy = correct / total
avg_loss = total_loss / len(dataloader)
print(f" Epoch {epoch+1}/{NUM_TEACHER_EPOCHS} | Loss: {avg_loss:.4f} | Acc: {accuracy:.4f}")
# 冻结教师模型参数 —— 蒸馏时教师不更新!
# 这是蒸馏的关键:教师是"已毕业的老师傅",不再学习
teacher_model.eval() # 切换到评估模式(关闭Dropout等)
for param in teacher_model.parameters():
param.requires_grad = False # 冻结所有参数
print("\n✅ 教师模型训练完成并已冻结!")
print(f" 教师模型参数量: {sum(p.numel() for p in teacher_model.parameters()):,}")
5.3 定义学生模型(BiLSTM)
# ============================================================
# 【第三步】定义学生模型 —— BiLSTM(轻量级序列分类模型)
# ============================================================
class BiLSTMClassifier(nn.Module):
"""
BiLSTM 文本分类器(学生模型)
架构:Embedding → BiLSTM → 取最后时刻隐状态 → 全连接分类头
与 BERT 的区别:
- BERT: Transformer架构,自注意力机制,~1.1亿参数
- BiLSTM: RNN架构,双向LSTM,~200万参数(小50倍!)
但通过蒸馏,BiLSTM可以学到BERT的"思考方式",
效果远好于只用硬标签训练的BiLSTM!
"""
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes,
num_layers=2, dropout=0.3):
"""
参数:
vocab_size: 词表大小(要和BERT tokenizer的词表一致)
embed_dim: 词向量维度(如256)
hidden_dim: LSTM隐藏层维度(如256)
num_classes: 分类类别数
num_layers: LSTM层数
dropout: Dropout比率
"""
super(BiLSTMClassifier, self).__init__()
# 词嵌入层:将每个token ID映射为一个稠密向量
# 注意:这里复用了BERT tokenizer的词表,所以vocab_size要和BERT一致
self.embedding = nn.Embedding(
num_embeddings=vocab_size, # 词表大小(BERT中文约21128)
embedding_dim=embed_dim, # 每个词用embed_dim维向量表示
padding_idx=0 # padding位置(ID=0)的嵌入不参与训练
)
# 双向LSTM层
# bidirectional=True → 同时有前向和后向LSTM
# 输出维度 = hidden_dim * 2(因为双向拼接)
self.lstm = nn.LSTM(
input_size=embed_dim, # 输入维度 = 词向量维度
hidden_size=hidden_dim, # LSTM每方向的隐藏维度
num_layers=num_layers, # LSTM层数
batch_first=True, # 输入shape: [batch, seq_len, embed_dim]
bidirectional=True, # 双向!前向+后向
dropout=dropout if num_layers > 1 else 0 # 多层时才用dropout
)
# Dropout层:防止过拟合
self.dropout = nn.Dropout(dropout)
# 分类头:将LSTM输出映射到类别空间
# 输入维度 = hidden_dim * 2(双向拼接)
self.classifier = nn.Linear(hidden_dim * 2, num_classes)
def forward(self, input_ids, attention_mask=None):
"""
前向传播
参数:
input_ids: shape [batch_size, seq_len] - token ID序列
attention_mask: shape [batch_size, seq_len] - 注意力掩码(可选)
返回:
logits: shape [batch_size, num_classes] - 未经softmax的原始分数
"""
# Step 1: 词嵌入
# input_ids: [batch, seq_len] → embedded: [batch, seq_len, embed_dim]
embedded = self.embedding(input_ids)
# Step 2: BiLSTM 编码
# embedded: [batch, seq_len, embed_dim]
# lstm_out: [batch, seq_len, hidden_dim*2](双向拼接)
# hidden: 最终隐藏状态(这里不用)
# cell: 最终细胞状态(这里不用)
lstm_out, (hidden, cell) = self.lstm(embedded)
# Step 3: 取最后一个时刻的输出(也可以用hidden拼接)
# 方案A:取序列最后一个时刻 → lstm_out[:, -1, :]
# 方案B:取最后一层的双向hidden拼接 → 更常用
# hidden shape: [num_layers*2, batch, hidden_dim]
# 取最后一层的前向和后向hidden拼接
# hidden[-2] = 最后一层前向, hidden[-1] = 最后一层后向
last_hidden = torch.cat([hidden[-2], hidden[-1]], dim=1)
# last_hidden shape: [batch, hidden_dim*2]
# Step 4: Dropout + 分类
dropped = self.dropout(last_hidden) # [batch, hidden_dim*2]
logits = self.classifier(dropped) # [batch, num_classes]
return logits # 返回logits!不是概率!和教师模型保持一致
# 创建学生模型
# vocab_size 必须和 BERT tokenizer 的词表大小一致
student_model = BiLSTMClassifier(
vocab_size=tokenizer.vocab_size, # BERT中文词表大小 ≈ 21128
embed_dim=256, # 词向量维度
hidden_dim=256, # LSTM隐藏维度
num_classes=NUM_CLASSES, # 5个类别
num_layers=2, # 2层LSTM
dropout=0.3 # 30% dropout
)
student_model.to(device)
# 打印参数量对比
teacher_params = sum(p.numel() for p in teacher_model.parameters())
student_params = sum(p.numel() for p in student_model.parameters())
print(f"\n📊 模型参数量对比:")
print(f" 教师模型(BERT): {teacher_params:>12,} 参数")
print(f" 学生模型(BiLSTM): {student_params:>12,} 参数")
print(f" 压缩比: {teacher_params/student_params:.1f}x")
5.4 核心:蒸馏损失函数(最关键代码!)
# ============================================================
# 【第四步】定义蒸馏损失函数 —— 这是整个蒸馏的核心!
# ============================================================
class DistillationLoss(nn.Module):
"""
知识蒸馏损失函数
总损失 = α × 硬标签损失 + (1-α) × 蒸馏损失(软标签)
硬标签损失:学生的标准输出 vs 真实标签 → 学"正确答案"
蒸馏损失: 学生的软输出 vs 教师的软输出 → 学"思考方式"
"""
def __init__(self, temperature=3.0, alpha=0.3):
"""
参数:
temperature (T): 温度参数,控制软标签的平滑程度
- T=1: 标准softmax,不平滑(等价于普通分类)
- T=3~5: 常用范围,适度平滑
- T=10: 高度平滑,暗知识多但信息可能被稀释
alpha: 硬标签损失的权重
- alpha=0.3: 表示30%权重给硬标签,70%给蒸馏
- 通常蒸馏损失权重更大,因为软标签信息更丰富
- 常见取值:0.1 ~ 0.5
"""
super(DistillationLoss, self).__init__()
self.temperature = temperature
self.alpha = alpha
# KL散度损失:衡量两个概率分布的差异
# reduction='batchmean':对batch取平均(推荐用法)
# 注意:PyTorch的KLDivLoss 期望输入是 log-probability!
self.kl_div = nn.KLDivLoss(reduction='batchmean')
# 交叉熵损失:用于硬标签
self.ce_loss = nn.CrossEntropyLoss()
def forward(self, student_logits, teacher_logits, labels):
"""
计算蒸馏总损失
参数:
student_logits: 学生模型的输出 logits [batch, num_classes]
teacher_logits: 教师模型的输出 logits [batch, num_classes]
labels: 真实标签 [batch]
返回:
total_loss: 标量,总损失值
"""
T = self.temperature
# ==========================================
# Part 1: 硬标签损失 (Hard Label Loss)
# ==========================================
# 学生logits(T=1,标准softmax) vs 真实标签
# CrossEntropyLoss 内部自动做 softmax + NLL
# 所以直接传 logits 和 labels 即可
hard_loss = self.ce_loss(student_logits, labels)
# 注意:这里学生logits没有除以T,因为硬标签用T=1
# ==========================================
# Part 2: 蒸馏损失 (Soft Label / Distillation Loss)
# ==========================================
# 教师软标签:教师的logits除以T,然后做softmax
# 用 detach() 确保不计算教师的梯度(教师参数已冻结,但这是好习惯)
# softmax(dim=1) 对类别维度做softmax
teacher_soft = F.softmax(teacher_logits.detach() / T, dim=1)
# teacher_soft shape: [batch, num_classes]
# 这就是教师的"软标签",包含了暗知识
# 学生软标签:学生的logits除以T,然后做log_softmax
# ⚠️ 重要:PyTorch的KLDivLoss要求输入是log概率!
# 所以这里用 log_softmax 而不是 softmax
student_log_soft = F.log_softmax(student_logits / T, dim=1)
# student_log_soft shape: [batch, num_classes]
# KL散度: KL(teacher_soft || student_soft)
# KLDivLoss(input, target) 中:
# input = log概率(学生的log_softmax)
# target = 概率(教师的softmax)
kl_loss = self.kl_div(student_log_soft, teacher_soft)
# 乘以 T²:补偿温度缩放导致的梯度缩小
# 这是Hinton论文中的关键技巧!
# 不乘T²的话,蒸馏损失的梯度会远小于硬标签损失
soft_loss = kl_loss * (T ** 2)
# ==========================================
# Part 3: 加权组合
# ==========================================
# total = α × hard_loss + (1-α) × soft_loss
total_loss = self.alpha * hard_loss + (1 - self.alpha) * soft_loss
return total_loss
# 创建蒸馏损失函数
# T=4: 温度适中,能提取丰富的暗知识
# alpha=0.3: 30%硬标签 + 70%蒸馏(让蒸馏发挥主要作用)
distill_criterion = DistillationLoss(temperature=4.0, alpha=0.3)
print("\n✅ 蒸馏损失函数已创建!")
print(f" 温度 T = {distill_criterion.temperature}")
print(f" 硬标签权重 α = {distill_criterion.alpha}")
print(f" 蒸馏损失权重 (1-α) = {1 - distill_criterion.alpha}")
5.5 蒸馏训练循环(核心训练逻辑)
# ============================================================
# 【第五步】蒸馏训练 —— 教师指导学生,最关键的训练循环!
# ============================================================
print("\n" + "="*60)
print("第五步:开始蒸馏训练 (BERT → BiLSTM)")
print("="*60)
# 学生模型的优化器
# 注意:学习率可以比训练BERT时大,因为BiLSTM从头训练
student_optimizer = optim.Adam(student_model.parameters(), lr=1e-3)
NUM_DISTILL_EPOCHS = 20
# 蒸馏训练循环
student_model.train() # 学生设为训练模式
teacher_model.eval() # 教师保持评估模式(确保不更新)
print(f"\n{'Epoch':>5} | {'Total Loss':>10} | {'Hard Loss':>10} | {'Soft Loss':>10} | {'Acc':>6}")
print("-" * 65)
for epoch in range(NUM_DISTILL_EPOCHS):
total_loss_sum = 0
correct = 0
total = 0
for batch in dataloader:
input_ids = batch['input_ids'].to(device)
attention_mask = batch['attention_mask'].to(device)
labels = batch['label'].to(device)
# ------ Step A: 教师模型前向传播(不计算梯度) ------
# 用 torch.no_grad() 确保教师模型不参与梯度计算
# 这很重要:教师是"已毕业的老师",参数不再更新
with torch.no_grad():
teacher_outputs = teacher_model(
input_ids=input_ids,
attention_mask=attention_mask
)
teacher_logits = teacher_outputs.logits # [batch, num_classes]
# ------ Step B: 学生模型前向传播(需要计算梯度) ------
student_logits = student_model(
input_ids=input_ids,
attention_mask=attention_mask # BiLSTM中可选使用
) # [batch, num_classes]
# ------ Step C: 计算蒸馏损失 ------
# 蒸馏损失 = α × CE(student, labels) + (1-α) × T² × KL(teacher_soft, student_soft)
loss = distill_criterion(student_logits, teacher_logits, labels)
# ------ Step D: 反向传播,更新学生参数 ------
student_optimizer.zero_grad() # 清空上一步的梯度
loss.backward() # 计算梯度(只更新学生参数!)
student_optimizer.step() # 更新学生模型参数
# ------ 统计 ------
total_loss_sum += loss.item()
predictions = torch.argmax(student_logits, dim=1)
correct += (predictions == labels).sum().item()
total += labels.size(0)
# 每个epoch打印训练信息
avg_loss = total_loss_sum / len(dataloader)
accuracy = correct / total
print(f" {epoch+1:>3}/{NUM_DISTILL_EPOCHS} | {avg_loss:>10.4f} | "
f"{'---':>10} | {'---':>10} | {accuracy:>6.4f}")
print("\n✅ 蒸馏训练完成!")
5.6 对比实验:普通训练 vs 蒸馏训练
# ============================================================
# 【第六步】对比实验 —— 证明蒸馏的价值!
# ============================================================
print("\n" + "="*60)
print("第六步:对比实验 —— 蒸馏 vs 普通训练")
print("="*60)
# --- 对照组:只用硬标签训练的BiLSTM(不使用蒸馏) ---
print("\n📌 训练对照组:BiLSTM + 纯硬标签(无蒸馏)...")
# 创建一个全新的BiLSTM(从零开始)
baseline_model = BiLSTMClassifier(
vocab_size=tokenizer.vocab_size,
embed_dim=256,
hidden_dim=256,
num_classes=NUM_CLASSES,
num_layers=2,
dropout=0.3
).to(device)
baseline_optimizer = optim.Adam(baseline_model.parameters(), lr=1e-3)
baseline_criterion = nn.CrossEntropyLoss() # 只用标准交叉熵
baseline_model.train()
for epoch in range(NUM_DISTILL_EPOCHS): # 同样的轮次
for batch in dataloader:
input_ids = batch['input_ids'].to(device)
attention_mask = batch['attention_mask'].to(device)
labels = batch['label'].to(device)
logits = baseline_model(input_ids, attention_mask)
loss = baseline_criterion(logits, labels)
baseline_optimizer.zero_grad()
loss.backward()
baseline_optimizer.step()
# --- 评估对比 ---
print("\n" + "="*60)
print("📊 最终评估结果:")
print("="*60)
def evaluate_model(model, dataloader, device, name):
"""评估模型准确率"""
model.eval()
correct = 0
total = 0
with torch.no_grad():
for batch in dataloader:
input_ids = batch['input_ids'].to(device)
attention_mask = batch['attention_mask'].to(device)
labels = batch['label'].to(device)
logits = model(input_ids, attention_mask)
predictions = torch.argmax(logits, dim=1)
correct += (predictions == labels).sum().item()
total += labels.size(0)
accuracy = correct / total
return accuracy
teacher_acc = evaluate_model(teacher_model, dataloader, device, "BERT (教师)")
student_acc = evaluate_model(student_model, dataloader, device, "BiLSTM (蒸馏)")
baseline_acc = evaluate_model(baseline_model, dataloader, device, "BiLSTM (普通)")
print(f"\n{'模型':<25} | {'准确率':>8} | {'参数量':>12} | {'推理速度':>8}")
print("-" * 65)
print(f"{'BERT (教师模型)':<20} | {teacher_acc:>8.4f} | {teacher_params:>12,} | {'慢':>8}")
print(f"{'BiLSTM (蒸馏训练)':<20} | {student_acc:>8.4f} | {student_params:>12,} | {'快':>8}")
print(f"{'BiLSTM (普通训练)':<20} | {baseline_acc:>8.4f} | {student_params:>12,} | {'快':>8}")
print(f"\n💡 关键观察:")
print(f" - 蒸馏后的BiLSTM 准确率通常 > 普通训练的BiLSTM")
print(f" - 但参数量只有BERT的 1/{teacher_params//student_params}")
print(f" - 推理速度大幅提升,适合线上部署!")
5.7 保存和加载蒸馏后的学生模型
# ============================================================
# 【第七步】保存蒸馏后的学生模型(部署用)
# ============================================================
# 保存学生模型
save_path = "bilstm_student_distilled.pth"
torch.save({
'model_state_dict': student_model.state_dict(),
'model_config': {
'vocab_size': tokenizer.vocab_size,
'embed_dim': 256,
'hidden_dim': 256,
'num_classes': NUM_CLASSES,
'num_layers': 2,
'dropout': 0.3,
},
'label_names': LABEL_NAMES,
}, save_path)
print(f"\n✅ 蒸馏后的学生模型已保存到: {save_path}")
# 推理示例
def predict_news(title, model, tokenizer, device, max_len=32):
"""用蒸馏后的学生模型预测新闻类别"""
model.eval()
# 编码输入
encoding = tokenizer(
title,
max_length=max_len,
padding='max_length',
truncation=True,
return_tensors='pt'
)
input_ids = encoding['input_ids'].to(device)
# 推理(不需要梯度)
with torch.no_grad():
logits = model(input_ids)
probs = F.softmax(logits, dim=1) # 转为概率
pred_id = torch.argmax(probs, dim=1).item()
confidence = probs[0][pred_id].item()
return LABEL_NAMES[pred_id], confidence
# 测试
test_titles = [
"腾讯发布新一代AI大模型",
"央行下调存款利率",
"春节档票房冠军出炉",
]
print("\n🔮 学生模型预测结果:")
for title in test_titles:
category, conf = predict_news(title, student_model, tokenizer, device)
print(f" 「{title}」 → {category} (置信度: {conf:.2%})")
六、面试高频考点总结
🎯 考点 1:什么是知识蒸馏?解决什么问题?
标准回答模板:
知识蒸馏是 Hinton 在 2015 年提出的模型压缩技术。核心思想是用一个已训练好的大型教师模型(如 BERT)的软标签来指导小型学生模型(如 BiLSTM)的训练。
它解决的问题是:大模型效果好但部署成本高(推理慢、显存大),蒸馏可以让小模型在保持接近大模型性能的同时,大幅降低推理开销。
关键创新在于"软标签"——教师模型输出的概率分布包含了类别间的相似性关系(暗知识),这比硬标签(one-hot)提供的信息更丰富。
🎯 考点 2:温度参数 T 的作用?怎么选?
- T 的作用:控制 softmax 输出概率分布的"平滑度"。T 越大,分布越平滑,非目标类别的概率越大,暗知识越丰富。
- T=1:标准 softmax,等价于普通分类,没有软化效果。
- T 过大:分布过于均匀,接近均匀分布,有效信息被稀释。
- 经验选择:一般 T ∈ [2, 10],常用 3~5,通过验证集调参。
🎯 考点 3:为什么蒸馏损失要乘 T²?(超高频!)
当使用温度 T 进行 softmax 软化时,soft target 对 logit 的梯度会按 $1/T^2$ 的尺度缩小。如果不乘 $T^2$ 进行补偿:
- 蒸馏损失的梯度会远小于硬标签损失的梯度
- 导致蒸馏损失在总损失中"话语权"过低,蒸馏效果大打折扣
乘 $T^2$ 后,两种损失的梯度量级一致,蒸馏才能真正发挥作用。
🎯 考点 4:蒸馏的三种类型
| 类型 | 蒸馏什么 | 代表方法 |
|---|---|---|
| Logit 蒸馏 | 教师最终输出的 logits/概率 | Hinton KD(本文方法) |
| 特征蒸馏 | 教师中间层的特征表示 | FitNets, PKT |
| 关系蒸馏 | 样本之间的关系/注意力模式 | RKD, Attention Transfer |
本文实现的是最经典的 Logit 蒸馏。
🎯 考点 5:Online vs Offline 蒸馏
Offline 蒸馏(本文方法):
1. 先训练好教师模型
2. 冻结教师,训练学生
→ 教师能力固定,简单稳定
Online 蒸馏:
教师和学生同时训练,互相学习
→ 不需要预训练好的教师,但训练更复杂
→ 代表:Deep Mutual Learning (DML)
Self 蒸馏:
模型自身的深层指导浅层
→ 不需要额外的教师模型
→ 代表:Born-Again Networks
🎯 考点 6:蒸馏 vs 微调 vs 从头训练
性能对比(一般规律):
BERT (教师) ████████████████████████ 95%
BiLSTM (蒸馏) ██████████████████████ 90% ← 蒸馏
BiLSTM (微调/硬标签) ████████████████████ 85% ← 普通训练
BiLSTM (随机初始化) █████████████████ 80% ← 从零开始
蒸馏 > 普通训练 > 从零开始
七、常见误区与调参技巧
❌ 误区 1:教师模型越大越好
正解:教师模型只需要"足够好",过大的教师反而可能导致软标签过于自信(概率集中),暗知识不够丰富。有时候用一个中等大小但效果好的教师,蒸馏效果反而更好。
❌ 误区 2:温度 T 越高越好
正解:T 太高会导致概率分布接近均匀分布,所有类别概率都差不多,反而失去了区分度。需要通过实验选择最佳 T。
❌ 误区 3:只要蒸馏损失,不需要硬标签损失
正解:硬标签损失($\alpha > 0$)通常是必要的。纯蒸馏($\alpha = 0$)会让学生过度拟合教师的"偏见",而硬标签可以起到纠偏作用。
✅ 调参技巧清单
1. 温度 T: 从 3 开始尝试,范围 [2, 5, 8, 10]
2. 权重 α: 从 0.3 开始,范围 [0.1, 0.3, 0.5, 0.7]
3. 学生容量: 不能太小,否则"装不下"教师的知识
4. 学习率: 学生可以用比教师大的学习率(学生是从头训练)
5. 训练轮次: 蒸馏通常需要比普通训练更多的epoch
6. 数据增强: 蒸馏时可以用更多无标签数据(教师生成伪标签)
总结:一张图记住知识蒸馏
┌─────────────────────────────────────────────────────────────┐
│ 知识蒸馏全景图 │
│ │
│ ┌──────────┐ 冻结 ┌──────────────────────────────┐ │
│ │ BERT │ ───────→ │ 教师 logits ÷ T → softmax │ │
│ │ (教师) │ │ = 软标签(含暗知识) │ │
│ │ 1.1亿参数 │ └──────────────┬───────────────┘ │
│ └──────────┘ │ │
│ ▼ │
│ ┌─────────────┐ │
│ │ KL散度 × T² │ ← 蒸馏损失 │
│ └──────┬──────┘ │
│ │ │
│ ┌──────────┐ 训练 ┌──────────┐ │ │
│ │ BiLSTM │ ──────→ │ 学生 │───┘ │
│ │ (学生) │ │ logits │ │
│ │ 200万参数 │ │ │──→ softmax → CE Loss │
│ └──────────┘ └──────────┘ (T=1) ← 硬标签损失│
│ │
│ 总损失 = α × 硬标签损失 + (1-α) × 蒸馏损失 │
│ │
│ 结果:小模型学到"正确答案" + "思考方式" = 又快又准! │
└─────────────────────────────────────────────────────────────┘
最后一句话:知识蒸馏的本质,不是让小模型"抄答案",而是让小模型学会大模型的思维方式——不仅知道"是什么",还知道"为什么"和"其他选项有多接近"。这就是暗知识的力量。