BERT模型-蒸馏实战

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()

祝你蒸馏出好学生!🚀

相关推荐
ai大模型-探索者4 天前
BERT模型压缩-量化实战
bert
月疯13 天前
bert的架构解析
人工智能·深度学习·bert
Jialu.16 天前
模型压缩实战:BERT 量化从 390MB 到 146MB 的实践
人工智能·深度学习·bert
Kobebryant-Manba16 天前
学习Bert微调
人工智能·学习·bert
Jialu.20 天前
中文 BERT 多任务分类项目:从模型结构到训练细节
人工智能·pytorch·分类·微软·nlp·bert
Jialu.20 天前
从 102M 到 34M:BERT 到 BiLSTM 的完整知识蒸馏实现
深度学习·nlp·bert
傲笑风21 天前
【openvino】tinybert基于openvino服务化部署(四)
人工智能·python·自然语言处理·nlp·bert·openvino
Lee_jerome1 个月前
python神经网络编程入门(四十四)——微调预训练 BERT 做下游任务
微调·bert·迁移学习·文本分类·预训练·小样本·冻结特征
Tbisnic1 个月前
BGE-M3 算法详解:从模型架构到三种检索方式的数学原理
算法·自然语言处理·大模型·bert·transformer·注意力机制