BERT 模型蒸馏实战教程:让大模型教出一个小而准的学生模型
配套项目:今日头条新闻 10 分类(BERT 教师 + BiLSTM 学生)。
一、为什么需要模型蒸馏
你已经训练出一个很准的 BERT 模型(409MB,准确率 95%+),但上线时遇到问题:
| 问题 | 影响 |
|---|---|
| 模型太大(409MB) | 部署成本高、内存吃紧 |
| 推理太慢 | BERT 前向传播慢,用户等待久 |
| 依赖 GPU | 硬件成本高 |
量化 能把模型压到 1/4,但如果想换一个结构完全不同、体积小几十倍的模型呢?------这就是**蒸馏(Knowledge Distillation)**要解决的问题。
二、什么是知识蒸馏
2.1 直观理解:导师带学生
想象一位博士生导师 (知识渊博,但很贵很慢)和一位研究生(年轻便宜,但经验少):
- 导师不是只丢给学生一个"标准答案"
- 而是把自己多年的判断逻辑、思考方式传授给学生
- 学生虽然脑子小,但学到"真传"后,水平能接近导师
知识蒸馏就是让大模型(教师)把"内功"传给小模型(学生)。
2.2 关键:软标签(Soft Label)才是精华
这是蒸馏最核心的洞察。看看同一个新闻,两种标签的区别:
硬标签(Hard Label) ------ 普通训练用的
"皇马输球替补席闹丑闻" → sports
只有一个答案,信息量少。
软标签(Soft Label) ------ 教师模型输出的概率分布
"皇马输球替补席闹丑闻" → sports: 0.85, game: 0.08, entertainment: 0.05, 其他: 0.02
每个类别的概率都有! 这里面藏着"知识":
- 它告诉学生:"这是体育,但和游戏、娱乐有一点像"
- 而"财经"的概率几乎为 0,说明"和财经完全不像"
类比 :老师批改作业时,不只是打"对/错",而是说"这道题你思路对了一半,但把 A 当成 B 了"。这种细节反馈比单纯的对错更有教学价值。
软标签 = 教师的"细节反馈",教会学生类别之间的相似性关系。
2.3 温度参数 T:让知识"露出来"
问题来了:教师模型的概率分布往往太极端:
sports: 0.9999, game: 0.00005, entertainment: 0.00003 ...
接近 0 的部分看不出信息。
温度 T 的作用:把概率分布"软化",让隐藏的小概率也浮现出来。
T = 1(原始): sports: 0.9999, game: 0.00005
T = 2(软化): sports: 0.85, game: 0.08 ← 知识浮现了!
公式 :softmax(logits / T) ------ T 越大,分布越平滑,暴露的"暗知识"越多。
类比:温度高,冰山融化,藏在水下的大部分(暗知识)就露出来了。
三、蒸馏的整体架构
训练数据(新闻标题)
│
┌──────┴──────┐
▼ ▼
┌───────────┐ ┌───────────┐
│ 教师模型 │ │ 学生模型 │
│ BERT │ │ BiLSTM │
│ 409MB │ │ 几MB │
│ 权重冻结 │ │ 需要训练 │
└─────┬─────┘ └─────┬─────┘
│ │
软标签 学生输出
│ │
└──────┬───────┘
▼
损失 = α·软标签损失(KL散度) + (1-α)·硬标签损失(交叉熵)
│
▼
更新学生模型参数
要点:
- 教师模型权重冻结 (
teacher.eval()+torch.no_grad()),只负责"输出软标签" - 只训练学生模型
- 总损失 = 软标签损失 + 硬标签损失
四、损失函数详解(蒸馏的核心)
4.1 总公式
Loss = α × soft_loss + (1 - α) × hard_loss
| 项 | 含义 | 作用 |
|---|---|---|
soft_loss |
软标签损失(KL 散度) | 学教师怎么想(概率分布) |
hard_loss |
硬标签损失(交叉熵) | 学正确答案是什么 |
α |
权重(如 0.7) | 软标签占多大比重 |
4.2 软标签损失:KL 散度
KL 散度衡量"两个概率分布有多不一样"。我们用它让学生分布贴近教师分布:
python
teacher_log_probs = F.log_softmax(teacher_logits / T, dim=1) # 教师软标签
student_log_probs = F.log_softmax(student_logits / T, dim=1) # 学生软标签
soft_loss = F.kl_div(student_log_probs, teacher_log_probs,
log_target=True, reduction='batchmean') * (T * T)
为什么要 * (T * T)?
因为除以 T 会让梯度变小,为了保持损失量级,需要乘以 T² 补偿。
4.3 硬标签损失:交叉熵
python
hard_loss = criterion(student_logits, teacher_preds) # 交叉熵
注意:这里用的是教师模型的预测结果 (teacher_preds)作为标签,而不是原始真实标签------这叫"教师伪标签",也是蒸馏的一种实践。
五、环境准备
bash
pip install torch transformers scikit-learn tqdm
六、项目结构
bert_distll/
├── bert_classifer_model.py # 教师模型:BERT 分类器
├── bilstm_classifier.py # 学生模型:BiLSTM 分类器
├── config.py # 配置(含 BiLSTM 超参)
├── utils.py # 数据加载
├── soft_label_distillation.py # ★ 蒸馏主程序
├── predict_fun.py # 蒸馏模型预测
└── models_save/
└── student_distll.pt # 蒸馏出的学生模型
七、教师模型(BERT)
python
from transformers import BertModel
from config import Config
import torch
conf = Config()
class BertClassifier(torch.nn.Module):
"""教师模型:BERT + 全连接分类层"""
def __init__(self):
super(BertClassifier, self).__init__()
self.bert = BertModel.from_pretrained(conf.bert_path)
self.fc = torch.nn.Linear(conf.hidden_size, conf.class_num)
def forward(self, input_ids, attention_mask):
_, pooled = self.bert(input_ids=input_ids,
attention_mask=attention_mask,
return_dict=False)
return self.fc(pooled)
教师模型直接加载已训练好的权重 (
bert20250521_.pt),不再训练。
八、学生模型(BiLSTM)
学生模型结构简单得多------Embedding + 双向 LSTM + 全连接:
python
import torch
import torch.nn as nn
from config import Config
conf = Config()
class BiLSTMClassifier(nn.Module):
"""学生模型:词嵌入 + 双向LSTM + 最大池化 + 全连接"""
def __init__(self, config):
super(BiLSTMClassifier, self).__init__()
# ① 词嵌入层:token ID → 向量
self.embedding = nn.Embedding(
num_embeddings=config.tokenizer.vocab_size, # 复用 BERT 词表大小
embedding_dim=config.embed_size # 256
)
# ② 双向 LSTM:提取序列特征(双向 → 输出维度 x2)
self.lstm = nn.LSTM(
input_size=config.embed_size, # 256
hidden_size=config.hidden_size_lstm, # 512
num_layers=config.num_layers, # 4
bidirectional=True,
batch_first=True
)
# ③ 全连接层:映射到 10 个类别
self.fc = nn.Linear(config.hidden_size_lstm * 2, config.class_num)
# ④ Dropout 防过拟合
self.dropout = nn.Dropout(p=config.dropout)
def forward(self, input_ids, attention_mask):
# 过滤 [CLS](101) 和 [SEP](102)
cls_sep_mask = (input_ids != 101) & (input_ids != 102)
valid_mask = attention_mask & cls_sep_mask
# 嵌入 + 屏蔽无效 token
embed = self.embedding(input_ids) * valid_mask.unsqueeze(-1)
# 双向 LSTM
lstm_out, _ = self.lstm(embed)
# 屏蔽无效位置
masked_output = lstm_out * valid_mask.unsqueeze(-1)
# 最大池化取最强特征
hidden, _ = masked_output.max(dim=1)
# Dropout + 全连接
return self.fc(self.dropout(hidden))
对应配置(config.py 中的 BiLSTM 超参)
python
self.embed_size = 256 # 词嵌入维度
self.hidden_size_lstm = 512 # LSTM 隐藏层维度
self.num_layers = 4 # LSTM 层数
self.dropout = 0.3 # Dropout 概率
self.class_num = 10 # 类别数
参数量对比 :BERT ~1.1 亿参数 → BiLSTM 约几百万参数,小了几十倍。
九、核心:蒸馏训练代码(soft_label_distillation.py)
python
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.optim import AdamW
from sklearn.metrics import classification_report, f1_score, accuracy_score, precision_score
from tqdm import tqdm
from config import Config
from utils import build_dataloader, get_time_diff
from bert_classifer_model import BertClassifier
from bilstm_classifier import BiLSTMClassifier
import time
import warnings
warnings.filterwarnings("ignore")
conf = Config()
def model2train():
# ===== 步骤1:准备数据 =====
train_dataloader, test_dataloader, dev_dataloader = build_dataloader()
# ===== 步骤2:定义教师模型(加载已训练权重,冻结)=====
teacher_model = BertClassifier()
state_dict = torch.load(conf.model_save_path, map_location=conf.device)
teacher_model.load_state_dict(state_dict)
# ===== 步骤3:定义学生模型 =====
student_model = BiLSTMClassifier(conf)
# ===== 步骤4:配置参数 =====
device = conf.device
num_epochs = conf.num_epochs
learning_rate = conf.learning_rate
T = 2.0 # ★ 温度参数:软化概率分布,暴露暗知识
alpha = 0.7 # ★ 软标签损失权重
step = 0
best_dev_f1 = 0.0
save_path = conf.save_model_path3 + "/student_distll.pt"
# ===== 步骤5:模型搬到设备 =====
teacher_model.to(device)
student_model.to(device)
# ===== 步骤6:优化器和损失函数 =====
optimizer = AdamW(student_model.parameters(), lr=learning_rate) # 只优化学生!
criterion = nn.CrossEntropyLoss() # 硬标签损失
# ===== 步骤7:训练循环 =====
for epoch in range(num_epochs):
student_model.train()
teacher_model.eval() # ★ 教师只推理,不更新
for batch_index, (input_ids, attention_mask, labels) in enumerate(
tqdm(train_dataloader, desc=f"软标签蒸馏训练 Epoch {epoch+1}/{num_epochs}")):
# 7.1 数据搬到设备
input_ids = input_ids.to(device)
attention_mask = attention_mask.to(device)
labels = labels.to(device)
# 7.2 清空梯度
optimizer.zero_grad()
# 7.3 教师前向(不计算梯度)
with torch.no_grad():
teacher_logits = teacher_model(input_ids, attention_mask)
teacher_preds = torch.argmax(teacher_logits, dim=1) # 教师硬标签
# 7.4 学生前向
student_logits = student_model(input_ids, attention_mask)
# 7.5 ★ 软标签损失(KL 散度)
teacher_log_probs = F.log_softmax(teacher_logits / T, dim=1)
student_log_probs = F.log_softmax(student_logits / T, dim=1)
soft_loss = F.kl_div(student_log_probs, teacher_log_probs,
log_target=True, reduction='batchmean') * (T * T)
# 7.6 硬标签损失(交叉熵)
hard_loss = criterion(student_logits, teacher_preds)
# 7.7 ★ 总损失 = 加权和
loss = alpha * soft_loss + (1 - alpha) * hard_loss
# 7.8 反向传播 + 更新
loss.backward()
optimizer.step()
# 7.9 每 2 批验证一次
if batch_index % 2 == 0:
report, f1score, accuracy, precision = model2dev(
student_model, dev_dataloader, device)
print(f"Step {step}, Epoch {epoch+1} - Dev F1: {f1score:.4f}, "
f"Acc: {accuracy:.4f}")
student_model.train() # 切回训练模式
if f1score > best_dev_f1: # 保存最优学生模型
best_dev_f1 = f1score
torch.save(student_model.state_dict(), save_path)
# 7.10 每个 epoch 结束再验证一次
report, f1score, accuracy, precision = model2dev(student_model, dev_dataloader, device)
print(f"Epoch {epoch+1} - Dev F1: {f1score:.4f}, Acc: {accuracy:.4f}")
student_model.train()
def model2dev(model, data_loader, device):
"""评估函数"""
model.eval()
preds, true_labels = [], []
with torch.no_grad():
for batch in tqdm(data_loader, desc="Evaluating ......"):
input_ids, attention_mask, labels = batch
input_ids = input_ids.to(device)
attention_mask = attention_mask.to(device)
labels = labels.to(device)
logits = model(input_ids, attention_mask)
batch_preds = torch.argmax(logits, dim=1)
preds.extend(batch_preds.cpu().numpy())
true_labels.extend(labels.cpu().numpy())
report = classification_report(true_labels, preds)
f1score = f1_score(true_labels, preds, average='micro')
accuracy = accuracy_score(true_labels, preds)
precision = precision_score(true_labels, preds, average='micro')
return report, f1score, accuracy, precision
if __name__ == "__main__":
start_time = time.time()
model2train()
print(f"Training completed in {get_time_diff(start_time)}")
十、代码关键点深度解析
① 教师模型必须"冻结"
python
teacher_model.eval() # 评估模式
with torch.no_grad(): # 不计算梯度
teacher_logits = teacher_model(...)
教师只提供"参考答案",不参与训练、不更新权重。
② 优化器只优化学生
python
optimizer = AdamW(student_model.parameters(), lr=learning_rate)
^^^^^^^^^^^^^ 只有学生的参数
如果写成 teacher_model.parameters() 就完全错了。
③ 软标签:除以 T 再 softmax
python
teacher_log_probs = F.log_softmax(teacher_logits / T, dim=1)
student_log_probs = F.log_softmax(student_logits / T, dim=1)
两边都要除以 T,保证在同一个"软化尺度"上比较。
④ KL 散度 + T² 补偿
python
soft_loss = F.kl_div(student_log_probs, teacher_log_probs,
log_target=True, reduction='batchmean') * (T * T)
log_target=True:表示第二个参数传入的是 log 概率reduction='batchmean':按批求平均* (T*T):补偿除以 T 造成的梯度缩小
⑤ 加权组合
python
loss = alpha * soft_loss + (1 - alpha) * hard_loss
α=0.7 表示"70% 学教师的思考方式,30% 学正确答案"。
十一、蒸馏后的预测
python
import torch
from config import Config
from bilstm_classifier import BiLSTMClassifier
conf = Config()
device = conf.device
model = BiLSTMClassifier(conf)
model.load_state_dict(torch.load(conf.save_model_path3 + "/student_distll.pt",
map_location=device))
model.to(device)
model.eval()
def predict(data):
text = data["text"]
encoded = conf.tokenizer.encode_plus(text, return_tensors="pt")
input_ids = encoded["input_ids"].to(device)
attention_mask = encoded["attention_mask"].to(device)
with torch.no_grad():
logits = model(input_ids, attention_mask)
pred_idx = torch.argmax(logits, dim=1).item()
pred_class = conf.class_list[pred_idx]
return {"text": text, "pred_class": pred_class}
if __name__ == "__main__":
print(predict({"text": "中华女子学院:本科层次仅1专业招男生"}))
十二、超参数调优建议
| 参数 | 作用 | 建议值 | 调整方向 |
|---|---|---|---|
| T(温度) | 软化概率分布 | 2.0 ~ 4.0 | 太小暗知识不显现;太大噪声多 |
| α(软标签权重) | 软/硬标签比重 | 0.5 ~ 0.9 | 教师强 → 调大;教师弱 → 调小 |
| 学习率 | 学生训练速度 | 5e-5 ~ 1e-3 | 学生比教师可以用大一点 |
| 学生模型容量 | 决定上限 | --- | 太小会"学不会"教师 |
经验 :教师越强、任务越难,α 应该越大(多学教师的智慧)。
十三、效果对比
| 模型 | 参数量 | 模型大小 | 推理速度 | 准确率 |
|---|---|---|---|---|
| 教师 BERT | ~1.1 亿 | 409 MB | 慢 | 95%+ |
| 学生 BiLSTM(无蒸馏) | 几百万 | 几十 MB | 快 | 较低 |
| 学生 BiLSTM(有蒸馏) | 几百万 | 几十 MB | 快 | 接近教师 |
蒸馏的价值:用小模型拿到接近大模型的精度。
十四、常见问题(FAQ)
Q1:教师模型需要重新训练吗?
不需要。直接用已经训练好的 BERT,加载权重即可。
Q2:为什么 α 是 0.7,软标签比硬标签权重还高?
因为软标签信息量更大(包含类别相似性)。经验上软标签权重大一些效果更好,但也要保留一定硬标签防止过拟合。
Q3:温度 T 取多少合适?
一般 2~4。T=1 就是原始分布(蒸馏意义不大),T 太大分布太平(噪声多)。
Q4:学生模型可以用别的结构吗?
可以。任何比教师小的模型都行(BiLSTM、TextCNN、小 Transformer 等),只要输入输出对齐。
Q5:蒸馏比量化好在哪?
- 量化:同结构压缩(BERT → 小 BERT)
- 蒸馏:跨结构压缩(BERT → BiLSTM),能小几十倍
Q6:训练很久没提升?
- 检查教师模型是否真的训练好了
- 提高 T 或 α
- 增加学生模型容量(embed_size / hidden_size)
- 学习率调大一点
Q7:RuntimeError: Expected all tensors on same device
→ 检查教师模型、学生模型、数据是否都 .to(device)。
十五、蒸馏 vs 量化 vs 剪枝
| 维度 | 蒸馏 | 量化 | 剪枝 |
|---|---|---|---|
| 原理 | 大模型教小模型 | 降精度(32→8位) | 删冗余参数 |
| 结构变化 | 换结构 | 不变 | 变瘦 |
| 压缩比 | 几十倍 | ~4 倍 | 可变 |
| 是否需训练 | 需要 | 不需要 | 通常需微调 |
| 精度 | 接近教师 | 几乎不变 | 略降 |
| 复杂度 | 高 | 极低 | 中 |
十六、总结
训练好的 BERT(教师,409MB,95%+)
│ 输出软标签(含类别相似性知识)
▼
BiLSTM 学生(几MB)
│ Loss = α·KL散度(软) + (1-α)·交叉熵(硬)
▼
蒸馏后的学生模型(小而准,适合部署)
核心一句话:
知识蒸馏 = 让大模型(教师)通过"软标签"把类别相似性等暗知识传授给小模型(学生) ,学生用
α×软损失 + (1-α)×硬损失学习,最终用小几十倍的体积拿到接近大模型的精度。
三个关键参数 :温度 T (软化分布)、权重 α (软硬比重)、教师模型的质量。
附:蒸馏速查代码
python
# 1. 教师(冻结)
teacher = TeacherModel().to(device)
teacher.load_state_dict(torch.load("teacher.pt"))
teacher.eval()
# 2. 学生(训练)
student = StudentModel().to(device)
optimizer = AdamW(student.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
T, alpha = 2.0, 0.7
for x, y in dataloader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
with torch.no_grad():
t_logits = teacher(x)
s_logits = student(x)
# 软标签损失(KL散度)
soft = F.kl_div(F.log_softmax(s_logits/T, dim=1),
F.log_softmax(t_logits/T, dim=1),
log_target=True, reduction='batchmean') * T * T
# 硬标签损失
hard = criterion(s_logits, y)
loss = alpha * soft + (1 - alpha) * hard
loss.backward()
optimizer.step()
祝你蒸馏出好学生!🚀