阶段三 · BERT 微调与意图识别

模型蒸馏与轻量化

一句话总结

BERT-base 有 1.1 亿参数,推理一次需要 10-50ms. 对群聊管理系统来说太重了.
模型蒸馏和量化能把模型压缩到原来的 1/3 ~ 1/10,推理速度提升 2-5 倍,效果只下降 1-3 个点.

前置回顾

前三篇完成了 BERT 微调和意图识别,拿到了可用的分类模型. 但 BERT-base 的 400MB 模型文件和每次 10-50ms 的推理延迟,在高并发的群聊场景下是个问题. 本篇解决部署前的最后一环:把模型压小,推快,为阶段四的 Go 集成做准备.

为什么需要轻量化

群聊管理系统的性能需求:

场景: 一个 bot 管理 500 个群,每个群日均 2000 条消息
日处理量: 500 × 2000 = 100 万条/天
峰值 QPS: 假设 80% 消息集中在 8 小时 → ~28 条/秒

BERT-base 推理(单条,CPU):
延迟: 30-50ms
吞吐: 20-33 条/秒
→ 刚好够用,但没有余量

BERT-base 推理(单条,GPU):
延迟: 5-10ms
吞吐: 100-200 条/秒
→ 够用,但需要 GPU 服务器(成本高)
方案 模型大小 CPU 延迟 效果损失
BERT-base(原始) 400MB 30-50ms 基准
BERT-small(6 层) 200MB 15-25ms F1 -1~2
DistilBERT(蒸馏) 250MB 15-20ms F1 -1~3
BERT-tiny(蒸馏) 50MB 3-5ms F1 -3~5
BERT-base + INT8 量化 100MB 10-15ms F1 -0.5~1
BERT-tiny + INT8 量化 17MB 1-3ms F1 -4~6

BERT-base = 全功能 SUV. 动力强,空间大,但油耗高停车难.
蒸馏模型 = 紧凑型轿车. 日常通勤完全够用,油耗减半.
量化 = 把汽油换成混动. 同一辆车,能耗进一步降低.
群聊管理这条路不需要 SUV--选对车型比堆排量重要.

知识蒸馏:让小模型学大模型

核心思想

知识蒸馏(Knowledge Distillation)的本质:用一个训练好的大模型(Teacher)去教一个小模型(Student).

传统训练:
训练数据(hard labels) → 小模型
标签: [1, 0, 0](确定的类别)

问题: 小模型容量有限,从 hard label 中学不到足够的信息

蒸馏训练:
训练数据 + Teacher 的 soft labels → 小模型
标签: [0.85, 0.10, 0.05](Teacher 的概率分布)

优势: soft label 包含了类间关系的信息

为什么 soft label 比 hard label 好? 看一个例子:

输入: "你有毛病吧"

Hard label: [1, 0, 0] → 对骂=1, 擦边=0, 正常=0
告诉学生: 这是对骂. 句号.

Soft label: [0.85, 0.12, 0.03] → Teacher 的输出
告诉学生: 这主要是对骂(0.85),但和擦边有点像(0.12),
和正常完全不像(0.03).
→ "有毛病"和"废物"比,和"无聊"比,"有毛病"更接近擦边

这种类间关系信息(dark knowledge)是 hard label 无法传递的.

Hard label = 标准答案只写 "A". 学生只知道选 A.
Soft label = 老师批改时写 "A 对,B 也有一定道理,C 完全不对". 学生不仅知道正确答案,还理解了为什么 B 有点像,C 为什么错. 学到的知识更丰富.

温度参数(Temperature)

蒸馏中有一个关键超参数:温度 T. 它控制 Teacher 输出的"软度":

原始 softmax(T=1):
logits = [3.0, 1.5, 0.2]
→ softmax → [0.85, 0.12, 0.03]
分布尖锐,主要类别占主导

高温 softmax(T=3):
logits = [3.0, 1.5, 0.2]
→ softmax(logits / 3) → [0.55, 0.28, 0.17]
分布平滑,次要类别的信息被放大

T 越大,分布越"软":
T=1: [0.85, 0.12, 0.03] ← 太尖锐,几乎和 hard label 一样
T=3: [0.55, 0.28, 0.17] ← 适中,暴露了类间关系
T=10: [0.38, 0.33, 0.29] ← 太平滑,信息被稀释
T=20: [0.35, 0.33, 0.32] ← 接近均匀分布,无用

经验值:T 通常设为 3-5. 太小 = soft label 退化成 hard label,太大 = 所有类别概率趋同.

蒸馏损失函数

训练 Student 时使用两个损失的加权和:

总损失 = α × 蒸馏损失 + (1-α) × 标准损失

蒸馏损失 = KL散度(Student_soft, Teacher_soft)
→ 让 Student 的概率分布接近 Teacher 的概率分布

标准损失 = 交叉熵(Student_hard, true_label)
→ 让 Student 的预测接近真实标签

α 通常取 0.5 ~ 0.7(蒸馏损失占主导)

蒸馏实战:从 BERT-base 到 BERT-tiny

方案一:用现成的蒸馏模型

最省事的方式是直接用别人蒸馏好的小模型,在自己的数据上微调:

from transformers import AutoModelForSequenceClassification, AutoTokenizer

# 蒸馏好的中文小模型
MODEL_NAME = "hfl/chinese-roberta-wwm-ext-small" # 6层,参数量减半

tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForSequenceClassification.from_pretrained(
MODEL_NAME, num_labels=3
)

# 后续训练流程和第 12 篇完全一样
# 只需要换一个 model name,其他代码不变

常用的中文小模型:

模型 层数 隐藏维度 参数量 大小
hfl/chinese-roberta-wwm-ext 12 768 102M ~400MB
hfl/chinese-roberta-wwm-ext-small 6 768 60M ~230MB
hfl/chinese-electra-180g-small 12 256 12M ~48MB
uer/chinese_roberta_L-4_H-512 4 512 24M ~95MB

方案二:自己做蒸馏

当现成模型效果不满足需求时,可以用自己微调好的 Teacher 做蒸馏:

import torch
import torch.nn.functional as F
from transformers import (
AutoModelForSequenceClassification,
AutoTokenizer,
Trainer,
TrainingArguments,
)

TEMPERATURE = 4.0
ALPHA = 0.5

teacher_model = AutoModelForSequenceClassification.from_pretrained(
"./chat-moderation-bert" # 微调好的 BERT-base
)
teacher_model.eval()

student_model = AutoModelForSequenceClassification.from_pretrained(
"hfl/chinese-electra-180g-small", num_labels=3
)

class DistillationTrainer(Trainer):
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
labels = inputs.pop("labels")
student_outputs = model(**inputs)
student_logits = student_outputs.logits

with torch.no_grad():
teacher_outputs = teacher_model(**inputs)
teacher_logits = teacher_outputs.logits

hard_loss = F.cross_entropy(student_logits, labels)

soft_student = F.log_softmax(student_logits / TEMPERATURE, dim=-1)
soft_teacher = F.softmax(teacher_logits / TEMPERATURE, dim=-1)
distill_loss = F.kl_div(soft_student, soft_teacher, reduction="batchmean")
distill_loss = distill_loss * (TEMPERATURE ** 2)

loss = ALPHA * distill_loss + (1 - ALPHA) * hard_loss

return (loss, student_outputs) if return_outputs else loss

trainer = DistillationTrainer(
model=student_model,
args=TrainingArguments(
output_dir="./distilled-model",
learning_rate=5e-5, # Student 用稍大的学习率
num_train_epochs=5, # 蒸馏训练多几轮
per_device_train_batch_size=32,
eval_strategy="epoch",
save_strategy="epoch",
load_best_model_at_end=True,
metric_for_best_model="f1",
),
train_dataset=tokenized["train"],
eval_dataset=tokenized["test"],
compute_metrics=compute_metrics,
)

trainer.train()

蒸馏的数据量要求

蒸馏训练可以用无标注数据--只需要 Teacher 的 soft label,不需要真实标签. 实践中常用策略:把大量未标注的群聊消息过一遍 Teacher 模型,生成 soft label,作为蒸馏的训练数据. 数据量越大,Student 学得越好.

量化:用更少的位数表示参数

原理

模型参数默认是 FP32(32 位浮点数). 量化就是把它们转成更少位数的表示:

FP32(32位):  精度最高,体积最大
1.234567890123... → 需要 4 字节

FP16(16位): 精度略降,体积减半
1.23457... → 需要 2 字节

INT8(8位): 精度有损,体积 1/4
值被映射到 [-128, 127] 的整数范围
1.234... → 映射为整数 79 → 需要 1 字节
量化类型 模型大小 推理速度 精度损失
FP32(原始) 400MB 基准
FP16 200MB 快 1.5-2x 几乎无
INT8 动态量化 100MB 快 2-3x 极小(F1 -0.5~1)
INT8 静态量化 100MB 快 2-4x 小(需要校准数据)

动态量化(最简单)

PyTorch 内置动态量化,一行代码完成:

import torch

model = AutoModelForSequenceClassification.from_pretrained("./chat-moderation-bert")
model.eval()

quantized_model = torch.quantization.quantize_dynamic(
model,
{torch.nn.Linear}, # 量化所有 Linear 层
dtype=torch.qint8
)

# 对比模型大小
import os

torch.save(model.state_dict(), "original.pt")
torch.save(quantized_model.state_dict(), "quantized.pt")

original_size = os.path.getsize("original.pt") / 1024 / 1024
quantized_size = os.path.getsize("quantized.pt") / 1024 / 1024
print(f"原始: {original_size:.1f}MB → 量化: {quantized_size:.1f}MB")
# → 原始: 418.2MB → 量化: 108.3MB

动态 vs 静态量化

动态量化在推理时动态计算每个 tensor 的量化参数(scale 和 zero_point),不需要校准数据,一行代码搞定.

静态量化需要提前用一批校准数据(calibration data)统计每一层的数值范围,然后固化量化参数. 速度更快,但需要额外步骤. 群聊场景用动态量化就够了.

方案选型:什么场景用什么

决策树:

你的任务对延迟有多敏感?

┌─────────────────────────┐
│ < 5ms(实时过滤) │ → BERT-tiny + INT8 量化
│ │ 或 ONNX Runtime 优化
│ 5-20ms(准实时) │ → DistilBERT / BERT-small
│ │ + 可选 INT8
│ 20-50ms(可接受) │ → BERT-base + INT8
│ │ 或 BERT-base + batch
│ > 50ms(不敏感) │ → BERT-base 原始
└─────────────────────────┘

群聊管理的实际选择建议:

组件 推荐方案 理由
对骂检测(全量消息) BERT-small + INT8 高吞吐,每条消息都要过
意图识别(仅管理员消息) BERT-base + INT8 QPS 低,可以用更大模型
实时关键词过滤 不用模型,纯规则 延迟 < 1ms

蒸馏效果对比实验

做一组完整的对比实验,量化不同方案的效果:

import time

def benchmark(model, tokenizer, test_texts, num_runs=3):
"""测量推理延迟和吞吐量"""
model.eval()
encoded = tokenizer(
test_texts, padding=True, truncation=True,
max_length=64, return_tensors="pt"
)

latencies = []
for _ in range(num_runs):
start = time.perf_counter()
with torch.no_grad():
model(**encoded)
latencies.append(time.perf_counter() - start)

avg_ms = (sum(latencies) / num_runs) * 1000
per_sample = avg_ms / len(test_texts)
return {
"batch_ms": f"{avg_ms:.1f}",
"per_sample_ms": f"{per_sample:.2f}",
"throughput": f"{1000 / per_sample:.0f} samples/s",
}

预期结果(CPU,batch_size=32,max_length=64):

模型 参数量 大小 延迟(每条) F1 相对 F1
BERT-base 102M 400MB 8ms 0.87 基准
BERT-base + INT8 102M 100MB 4ms 0.865 -0.5
BERT-small(6 层) 60M 230MB 4ms 0.855 -1.5
BERT-tiny(蒸馏) 12M 48MB 1ms 0.83 -4.0
BERT-tiny + INT8 12M 17MB 0.5ms 0.825 -4.5

蒸馏不是万能的

如果 BERT-base 在某个任务上只有 0.75 的 F1(数据或标注质量差),蒸馏后可能掉到 0.70 以下--变得不可用. 蒸馏的前提是 Teacher 本身足够好. 优先保证 Teacher 的效果,再考虑蒸馏压缩.

衔接阶段四:从 Python 到 Go

模型压缩完成后,下一步是部署. Go 后端不能直接运行 PyTorch 模型. 阶段四会解决这个问题:

阶段三(本篇完成):
Python 训练/蒸馏 → 得到轻量化模型(.pt 或 .safetensors)

阶段四(即将开始):
.pt → 导出为 ONNX → Go 调用 ONNX Runtime → 推理服务

或者:
.pt → 导出为 ONNX → 部署到 Triton Server → Go 通过 gRPC 调用

ONNX(Open Neural Network Exchange)是模型的"跨平台二进制格式"--类似 protobuf 对数据做的事情. 导出为 ONNX 后,模型不再依赖 PyTorch,可以被任何语言的 ONNX Runtime 加载运行.

PyTorch 模型 = Go 源码. 只能在 Go 编译器环境中运行.
ONNX 模型 = 编译好的二进制文件. 任何操作系统都能运行,不需要编译器.
ONNX Runtime = 操作系统的运行时. 负责加载二进制并高效执行.

快速回顾

  • 知识蒸馏:大模型(Teacher)教小模型(Student),用 soft label 传递类间关系信息
  • 温度参数:T=3~5,控制 soft label 的平滑度. 太小退化成 hard label,太大信息被稀释
  • INT8 量化:模型体积缩小到 1/4,推理速度提升 2-3 倍,精度损失极小
  • 动态量化:一行代码 quantize_dynamic,不需要校准数据,群聊场景够用
  • 方案选型:全量消息过滤用 BERT-tiny + INT8(快),管理员意图识别用 BERT-base + INT8(准)
  • 下一步:ONNX 导出 → Go 集成(阶段四)

动手练习

  1. 用现成小模型微调:用 hfl/chinese-electra-180g-small 替换 BERT-base,在相同数据上训练,对比 F1 和推理延迟
  2. 动态量化:对微调好的 BERT-base 做 INT8 动态量化,对比量化前后的模型大小,推理延迟,和 F1 值
  3. 蒸馏训练:用 DistillationTrainer 从自己的 BERT-base Teacher 蒸馏一个 BERT-tiny Student,调整 T 和 α 参数,找到 F1 损失 < 3 个点的最小模型
  4. 延迟基准测试:用 benchmark 函数测量 BERT-base, BERT-small, BERT-tiny 在 CPU 上的推理延迟(batch_size=1 和 batch_size=32),画对比图
  5. 端到端选型报告:综合 F1,延迟,模型大小三个指标,为群聊管理系统选出最优方案,写一份一页纸的选型报告