ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

知识蒸馏实战:用Pytorch将BERT压缩为TextCNN的文本分类方案

知识蒸馏实战:用Pytorch将BERT压缩为TextCNN的文本分类方案 简介一套基于Pytorch的中文文本分类知识蒸馏实践项目面向具备一定深度学习基础、希望掌握大模型轻量化落地技巧的技术人员。项目核心是将Hugging Face的bert-base-chinese蒸馏至BiLSTM同时给出梯度累加、混合精度训练、对抗训练三种扩展实验适合作为模型压缩与训练优化的参考模板。压缩包共43个文件以22个Python脚本为主干覆盖蒸馏主流程及多策略配置9个pkl文件为预训练词表或中间产物txt与json分别存放数据说明与参数配置整体大小63.85MB。目录按data、config、models、checkpoints、processor分层便于按模块定位数据、模型定义与训练入口。除基础kd_main.py外还提供main_with_attack.py、main_with_apex.py等变体入口可直接对比不同策略下的训练效果。学习者可从中梳理BERT到BiLSTM蒸馏的完整数据流与训练流程也能借鉴对抗训练、混合精度等技巧在中文分类场景中的集成方式。目前已有312人学习适合用于课程设计、算法对比或个人项目起步。1. 知识蒸馏在中文文本分类里到底值不值得做从BERT到小模型的距离中文文本分类做到要上线的阶段常会撞上一个尴尬微调后的BERT在验证集上表现很好但推理慢、显存占CPU服务器上扛不住并发。知识蒸馏Distillation这时候比任何调参都管用——用一个大模型当教师把它的判断倾向教给一个轻量学生模型。基于Pytorch实现知识蒸馏来做中文文本分类是我在模型压缩时最常用的一套组合拳。适合两类人一类是人工智能课程大作业需要做完整项目实践的同学另一类是想把分类模型塞进低配服务器或边缘设备的工程师。反直觉的是学生模型拿到的不是教条般的0/1标签而是教师对每个类别的犹豫——这份犹豫才是知识。2. 先搞懂“蒸”的是什么教师-学生框架与软标签的前因后果我最早接触知识蒸馏时以为就是把大模型输出的概率当成新的标签去训练小模型。后来才发现真正值钱的是“怎么把概率变成标签”的过程。2.1 硬标签只告诉你答案软标签告诉你思考过程假设一个中文新闻分类任务三个类别体育、财经、娱乐。一条真实样本是“某俱乐部官宣新援加盟”真实标签是体育。硬标签是[0, 1, 0]学生模型从中学到的只有“这是体育”。但BERT这样的教师模型在softmax之后可能给出[0.65, 0.25, 0.10]——它认为这条样本有25%的财经味道10%的娱乐关联。这0.25和0.10不是噪音是类别边界的信息模型通过训练学到“俱乐部”“加盟”在财经语境里也常出现。用硬标签训练学生模型学生会把决策边界想象得很“陡峭”靠近边界的样本类别之间没有任何过渡。而软标签把边界的坡度原样搬了过来学生在学习时不仅知道答案还知道答案和邻近类别之间的距离。这也是为什么蒸馏后的学生模型往往比直接从硬标签训练的同结构模型更稳。教师模型在我这里更像一个“会解释的老师”而不是“给答案的考官”。在这个阶段其实不需要关心模型的内部结构。教师模型在推理时给出的类别分布是一个整体行为可以把它当成一个黑匣子来用我们要保留的正是这个黑匣子输出里的概率关系。2.2 温度T控制教师“犹豫”程度的旋钮softmax(Q/T) 里的温度 T 直接改变分布的形状。T 1 就是普通softmax教师的输出往往在正确类别上接近1错误类别的概率很小学生依然只能看到近乎one-hot的信号。T 越大分布越平缓类别差异被放大学生能看到教师更多的“思考痕迹”。但这个放大不是无代价的。我一般把 T 分成三个区间来调T在2到5之间适合大多数中文文本分类任务T低于1分布比原始softmax更尖锐等价于让教师更“自信”通常只在训练末尾使用T超过10分布过于均匀教师自己的错误也会被放大学生学到的更多是类别先验而不是知识。一个很典型的翻车现场是新手同学把T调到20学生模型迅速收敛到输出均匀分布。温度不是“越大越好”的超参它和教师模型本身的置信度强相关。如果教师模型在验证集上精度很高、输出的softmax非常自信T可以适当取到4~6如果教师本身中等水平、分布已经很平T取2~3更安全。调T的步骤我会固定为先固定T3训练一轮然后在验证集上把T 2、4、6都扫一遍其余参数不动看学生模型的F1变化。T取值范围分布形态适合场景T1原始softmax分布学生学到的软信号最少效果接近普通训练T2~5分布平滑类别关系清晰中文文本分类的常用区间从这里起步调T6~10分布接近均匀教师置信度低时用高置信度教师会引入过多噪声T10近乎均匀分布容易让学生学成“和事佬”一般不推荐2.3 KD Loss公式与Pytorch实现把知识变成梯度温度T是用来“软化”分布的真正让学生模型更新参数的是蒸馏损失。常见的KD Loss由两部分构成一部分是学生和真实硬标签之间的交叉熵另一部分是学生和教师软标签之间的KL散度。实践中第二个部分需要乘以 T^2原因在于KL散度里带有1/T的梯度缩放乘回去之后两部分损失的梯度量级才一致训练才不会一边倒。L_total alpha * CE(student_logits, y_true) (1 - alpha) * KL(softmax(student_logits / T), softmax(teacher_logits / T)) * T^2在Pytorch里我习惯把软化softmax抽成一个函数单独测一遍再放进训练循环里用。因为这里的log_softmax和KLDivLoss的配对非常容易写错。import torch import torch.nn.functional as F def log_softmax_with_temperature(logits, temperature): # 返回log概率方便直接接KLDivLoss return F.log_softmax(logits / temperature, dim-1) def kd_loss(student_logits, teacher_logits, temperature): # 学生分支用log_softmax教师分支用普通softmax student_log_probs log_softmax_with_temperature(student_logits, temperature) teacher_probs F.softmax(teacher_logits / temperature, dim-1) # KLDivLoss第一个参数要求是log概率第二个参数是普通概率 loss F.kl_div(student_log_probs, teacher_probs, reductionbatchmean) # 乘T^2保持梯度量级这个乘法对应公式里的补偿项 return loss * (temperature ** 2)参数说明temperature直接参与计算没有把dim写死保证logits形状是(batch, num_classes)时按最后一维算。reductionbatchmean在较新版本的PyTorch里语义明确它返回的是batch维度上的均值相比mean更适合蒸馏任务里衡量分布差异。另一个值得注意的点是教师logits和学生的logits要来自同一个label space一旦教师头换了类别数KL散度会直接在dim-1上报错。提示教师logits和学生logits的类别数必须一致否则KL散度会直接在dim-1上报错。3. 基于Pytorch搭建蒸馏训练模型选型、Loss与训练循环原理搞明白之后落地的时候就会遇到另一个问题模型代码从哪来。教师和学生模型在Pytorch生态里都有现成的实现关键是怎么把蒸馏逻辑接进训练循环。为什么用Pytorch而不是TensorFlow蒸馏需要同时维护教师和学生两条计算图PyTorch动态图的写法几乎是把公式直接翻译成代码调试时打印两个模型的logits也方便。如果你是从零开始搭这套环境pytorch环境搭建只剩一个关键点GPU版torch的CUDA版本要和驱动匹配装错了训练慢到像CPU还查不出原因。环境弄好之后先跑一个很小的张量运算确认设备可用再开始蒸馏。3.1 教师与学生模型的常见搭配BERT到TextCNN教师模型我用得最多的是bert-base-chinese它本身是transformers库里的标准模型在中文文本分类上稍微加一个分类头就能达到很稳的基线。学生模型我通常选TextCNN不是因为它新而是因为它在短文本分类上性价比极高卷积核并行扫描局部n-gram特征网络结构简单部署时不需要额外的优化就能跑得很快。学生模型不能选得太弱比如只有一个线性层那样即使蒸馏也学不到足够的表示。我一般用三组卷积核窗口大小分别取2、3、4每个窗口配128个卷积核后面接一个global max pooling把变长输入压成定长向量最后接分类层。这个配置在中英文句子分类上都能稳定复现。import torch import torch.nn as nn class TextCNNStudent(nn.Module): def __init__(self, vocab_size, embed_dim, num_classes, num_filters128, dropout0.3): super().__init__() # 词向量层负责把token id变成稠密向量 self.embedding nn.Embedding(vocab_size, embed_dim, padding_idx0) # 三组卷积核窗口大小对应2/3/4个字的局部组合 self.convs nn.ModuleDict({ 2: nn.Conv2d(1, num_filters, kernel_size(2, embed_dim)), 3: nn.Conv2d(1, num_filters, kernel_size(3, embed_dim)), 4: nn.Conv2d(1, num_filters, kernel_size(4, embed_dim)), }) self.dropout nn.Dropout(dropout) self.fc nn.Linear(num_filters * 3, num_classes) self._init_weights() def _init_weights(self): # 嵌入层和全连接层做合理的初始化避免蒸馏一开始就震荡 nn.init.kaiming_uniform_(self.fc.weight, a0.5) def forward(self, input_ids, attention_maskNone): x self.embedding(input_ids) # (batch, seq_len, embed_dim) x x.unsqueeze(1) # 卷积需要channel维度 pooled [] for conv in self.convs.values(): out conv(x).squeeze(3) # 去掉embed_dim维度 pooled.append(out.max(dim2).values) # global max pooling x torch.cat(pooled, dim1) x self.dropout(x) return self.fc(x) # 返回logits不在这里做softmax参数说明padding_idx0很重要因为HuggingFace的tokenizer会把[PAD]的token id设为0Embedding对id为0的位置不更新可以省一点显存和无效更新。forward里返回的是logits而不是概率蒸馏损失和交叉熵损失都需要在logits上作用softmax放外层处理。文本长度超过卷积核窗口时max pooling能确保输出尺寸固定这也是TextCNN对输入长度不敏感的原因。第一次做这个项目容易忽略一个点学生模型的vocab_size要和教师tokenizer的词表大小保持一致。因为整个训练的数据都来自同一个tokenizer学生的Embedding层输入是教师分词器产生的token id。如果这里对不上训练时会直接报idx越界。如果想让学生起步更快可以把BERT的embedding矩阵取出来按同样下标复制给学生能省不少早期训练时间。3.2 蒸馏损失函数怎么写CE与KL的加权组合有了两个模型和温度操作接着就是组装损失。第2章给过KD Loss的核心实现实际训练里还要把它和一个硬标签交叉熵组合。组合权重alpha我一般从0.3起步alpha是硬标签CE的权重1-alpha是蒸馏KL的权重。alpha太大学生只学到硬标签和普通训练没区别alpha太小学生被教师牵着走万一教师有系统性偏差错误会被完整继承。import torch import torch.nn as nn import torch.nn.functional as F class DistillLoss(nn.Module): def __init__(self, alpha0.3, temperature3.0): super().__init__() self.alpha alpha self.temperature temperature self.ce_loss nn.CrossEntropyLoss() def forward(self, student_logits, teacher_logits, labels): # 硬标签分支学生自己的分类能力 ce self.ce_loss(student_logits, labels) # 软标签分支学生对教师分布的拟合 student_log_probs F.log_softmax(student_logits / self.temperature, dim-1) teacher_probs F.softmax(teacher_logits.detach() / self.temperature, dim-1) # 教师分支要除以同样的温度否则两边的分布不对齐 kl F.kl_div(student_log_probs, teacher_probs, reductionbatchmean) kl kl * (self.temperature ** 2) total self.alpha * ce (1 - self.alpha) * kl return total, ce.item(), kl.item()参数说明teacher_logits.detach() 是整个蒸馏训练里最重要的一个操作。教师本身已经收敛蒸馏过程不想再更新它的参数detach切断了反向传播的计算图如果不加Pytorch会把教师和学生当成一个大网络整体回传不仅浪费显存还会把梯度噪声引入教师。返回的ce.item()和kl.item()是为了方便打印日志时观察两个损失各自的变化趋势一旦kl掉得很快而ce压不住就说明教师信号太强需要调小1-alpha。注意如果教师模型的forward返回的是transformers的output对象务必取.logits再参与损失计算不要把整个output对象传进DistillLoss。3.3 训练循环的三个关键细节eval模式、梯度裁剪、学习率训练循环表面上是标准的Pytorch写法但细节都藏在这几步里。第一个细节是教师模型必须切到eval模式否则Dropout会被错误地打开教师的输出变成随机采样蒸馏失去意义。第二个细节是梯度裁剪。学生模型的梯度在早期可能被教师logits的绝对值放大不裁剪的后果是loss一下飙升到几百。第三个细节是学习率。教师已经经过预训练学习率适合调低学生模型要重新学特征学习率可以给到比教师高一个数量级。optimizer torch.optim.AdamW(student_model.parameters(), lr2e-5, weight_decay0.01) scheduler torch.optim.lr_scheduler.LinearLR(optimizer, start_factor1.0, total_iters200) teacher_model.eval() student_model.train() for batch in train_loader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) labels batch[labels].to(device) # 教师前向全程不更新梯度也不会因为Dropout产生随机性 with torch.no_grad(): teacher_logits teacher_model(input_ids, attention_maskattention_mask).logits student_logits student_model(input_ids, attention_mask) loss, ce, kl distill_criterion(student_logits, teacher_logits, labels) optimizer.zero_grad() loss.backward() # 裁剪梯度的作用学生模型前期梯度方差大先限制范数再走优化器 nn.utils.clip_grad_norm_(student_model.parameters(), max_norm1.0) optimizer.step() scheduler.step()参数说明teacher模型的输出我直接取了.logits这是transformers库的约定如果你用的是自己定义的BERT分类器记得去掉最后面的softmax。学习率2e-5是BERT微调的常见起点但在蒸馏任务里学生模型的学习率不一定要和教师一样小。如果你把学生换成BiLSTM或者LSTM2e-5容易跑不动可以用1e-3级别。LinearLR做了简单的200步线性warmup前200步学习率从0逐步爬升这一步能有效缓解学生模型在训练初期被教师信号冲乱的风险。4. 中文文本分类的数据侧准备Tokenizer、标签与Dataset模型和损失定了数据侧才是中文文本分类最容易被问到的部分。BERT的输入是token id但中文文本要先经过tokenizer这里有几个和英文任务习惯不一样的地方需要单独处理。4.1 中文预训练模型按字切分Tokenizer和max_length的设定中文的bert-base-chinese不是按词切分的它内部按字切分整个词表是汉字级别的。这意味着分词器不需要jieba参与直接按字符切反而更稳。很多同学在中文文本分类里切出词边界再送进BERT多此一举不说还会稀疏掉本来连续的上下文。我一般用transformers的AutoTokenizer加载调用encode_plus一次性完成切分、加[CLS]/[SEP]、padding和截断。from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(bert-base-chinese) def encode_text(text, max_length64): # 中文短文本分类max_length取64能覆盖绝大多数单句样本 encoded tokenizer.encode_plus( text, max_lengthmax_length, paddingmax_length, truncationTrue, add_special_tokensTrue, return_tensorspt, return_attention_maskTrue, ) return encoded[input_ids].squeeze(0), encoded[attention_mask].squeeze(0)参数说明max_length64对中文单句分类足够了。中文字密度高一个句子平均20~40个字64个token把绝大多数样本完整保留。paddingmax_length配合truncationTrue会把超过64的样本截断并把不足64的样本补齐。两个返回值必须在同一个batch里保持形状一致DataLoader后续collate时才会正常。中文长文本分类另说如果样本是整段新闻、平均500字max_length需要提到128甚至256但此时显存占用会明显上升。4.2 类别不均衡怎么处理在蒸馏框架里给CE Loss加权重中文文本分类数据集往往存在类别不均衡比如财经类样本是娱乐类的10倍。直接从软标签训练学生模型会继承教师的概率分布而教师的分布本身就偏向高频类别所以学生也会跟着偏。我处理这种问题的顺序是先看教师模型的混淆矩阵确认教师在高频类上是否也偏如果教师已经偏就不能光靠蒸馏硬标签这一路的CE Loss需要加类别权重。假设你的训练集四个类别样本数如下class_counts [12000, 3000, 1000, 800] total sum(class_counts) # 类别权重和频率成反比低频类被放大 weights [total / (len(class_counts) * c) for c in class_counts] weights torch.tensor(weights, dtypetorch.float).to(device) criterion_ce_weighted nn.CrossEntropyLoss(weightweights)参数说明权重计算公式用的是“总样本数 / (类别数 * 该类样本数)”这个公式会把所有类别的权重均值控制在1附近不会让低频类权重爆炸。想简单一点也可以用1/c但高频类权重会过大。distill损失里CE部分替换成带权重的版本KL部分不要加权重因为教师的软分布本身就包含了类别先验强行加权会破坏分布关系。这里常见的错误是在CE和KL两路都加权重导致学生模型在低频类上过拟合。4.3 用Dataset缓存编码结果省掉重复分词的隐性开销在Pytorch里训练文本分类模型最容易忽视的开销是tokenizer。如果每次DataLoader取样本时才做分词训练一个epoch要重新对全部数据做一次字符切分在中文任务上尤其浪费。我的做法是先把所有文本编码成input_ids和attention_mask存进内存或磁盘缓存训练过程中只做张量搬运不再碰tokenizer。这个预处理只需要跑一次却能省下每轮训练里超过30%的时间。import torch from torch.utils.data import Dataset, DataLoader class TextDataset(Dataset): def __init__(self, texts, labels, max_length64): # 预处理阶段一次性完成tokenize并用list保存 self.input_ids [] self.attention_masks [] for text in texts: ids, mask encode_text(text, max_length) self.input_ids.append(ids) self.attention_masks.append(mask) self.labels labels def __len__(self): return len(self.labels) def __getitem__(self, idx): return { input_ids: self.input_ids[idx], attention_mask: self.attention_masks[idx], labels: torch.tensor(self.labels[idx], dtypetorch.long), } train_dataset TextDataset(train_texts, train_labels) train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers0, # Windows下默认0最稳Linux可以开到4 pin_memoryTrue, )参数说明DataLoader的batch_size32在BERT教师和TextCNN学生的组合下12GB显存通常能跑得动如果你的显卡只有6GB降到16或者干脆走离线蒸馏把教师的logits缓存在磁盘上。num_workers在Windows上的Pytorch容易踩多进程坑我一般直接设0Linux服务器再开2~4。pin_memoryTrue适合在GPU训练时启用它把数据页锁在内存里减少CPU到GPU的拷贝时间。预处理阶段如果发现内存吃紧可以把input_ids统一转成numpy数组存磁盘训练时用memory map加载。5. 蒸馏训练避坑指南5个翻车点的现象、原因与解决知识蒸馏在项目实践里并不总是顺利。我在这套流程里踩过的坑下面挑最典型的五个写出来每个都按现象、原因、解决三步拆开。5.1 学生效果上不去甚至不如直接用硬标签训练现象蒸馏跑了20个epoch学生模型精度和直接用硬标签训练差不多有时候还更低。看损失曲线KL散度掉得很慢。原因教师模型没有切到eval模式。PyTorch里模型默认是training模式BERT内部的Dropout层还在工作教师每次前向输出的logits都是带随机性的。学生学的是一个“抖动的教师”蒸馏信号质量很差。另外如果教师在蒸馏前没有被充分微调教师本身的精度不够学生学到的上限就被锁死了。解决在训练循环之前必须显式调用teacher_model.eval()并用with torch.no_grad()包住教师前向。如果你是离线缓存教师logits只需要确认生成缓存时模型是eval状态训练阶段完全不用管教师。还有一个容易被忽略的点要验证教师在自己数据集上的基线精度如果教师分类头都训练不到位先别急着蒸馏回去把教师微调好再来做知识迁移。5.2 学生模型输出变成“和事佬”每个类别概率都差不多现象训练中loss正常下降但验证时学生模型几乎对所有样本都输出均匀分布分类结果集中在高频类别上。原因温度T设得太大。T超过10以后教师softmax的输出接近均匀分布KL散度在学生看来变成了“教你做一个对所有类别都给1/3概率的人”。教师本身在错误类别上的概率也被放得很大学生学到的是“大家都有份”的平庸输出。解决把T降回到3左右。判断T是否合适的办法是打印教师的软标签取一批验证样本看教师模型的平均置信度。如果教师绝大多数样本的top-1概率都在0.8以上T可以取3~4如果教师本身top-1概率只有0.5左右T取2更合适。调T时要同步调alphaT变大意味着软标签贡献变强alpha适当调大能压住学生对均匀分布的过度拟合。5.3 蒸馏效果比直接训练还差学生的决策边界混乱现象学生模型在训练集上loss很低但验证集F1比从零训练还差且错误集中在padding位置附近。原因attention_mask没有正确使用。BERT教师需要它来区分真实token和padding但TextCNN学生模型的卷积对padding区域照样提取特征。更隐蔽的问题是padding部分不参与教师的池化但学生的卷积会把padding位置的“空字”特征也学进去导致学生学到“比padding残缺内容”这种伪特征。解决确认两个模型的forward都接收attention_mask。BERT侧必须显式传递TextCNN侧即使不参与计算也要在DataLoader返回的dict里保留这个字段。如果问题是padding占比过高直接把max_length从64降到实际覆盖95%样本的长度减少padding噪音。5.4 训练中loss震荡甚至出现nan现象loss从很小的值突然跳到几千继续训练又降回来间歇性nan。打印logits看到绝对值到了几十甚至上百。原因两个模型的logits量级不一致。教师模型的logits可能分布在[-5, 5]学生模型因为初始化问题可能到[-20, 20]。KL散度对logits的绝对量级很敏感量级不匹配时loss陡增。另一个原因是alpha和T没有组合好KL部分梯度被T^2放大后学生模型的Embedding层在短时间内剧烈更新导致梯度爆炸。解决在loss组装前先打印两个模型的logits标准差。量级差异过大就在KL分支入口先做一个归一化直接对teacher_logits做标准化或者乘一个scale系数让两边logits量级接近。梯度裁剪clip_grad_norm_加上max_norm1.0能减少大部分梯度爆炸。如果是nan且裁剪没用检查一下KLDivLoss之前是否出现softmax溢出logits上千以后softmax数值不稳定需要用log_softmax的数值稳定版本。5.5 显存不够教师和学生一起前向直接把显卡压爆现象batch_size32训练直接OOMbatch_size8能跑但速度慢得无法接受。原因教师BERT和学生TextCNN同时前向中间变量都留在计算图里显存占用等于两套模型的前向激活之和。这是在线蒸馏的固有成本。解决换成离线蒸馏。先用教师模型把所有训练样本推理一遍把logits和对应标签存到磁盘之后学生训练时只读文件不再加载教师模型。这一步是大项目里几乎必做的优化它还把常识蒸馏和学生训练解耦教师只需要跑一次后续调alpha、调T都不需要重新过教师。from pathlib import Path import numpy as np cache_path Path(./teacher_logits.npy) label_path Path(./train_labels.npy) if not cache_path.exists(): teacher_model.eval() all_logits, all_labels [], [] with torch.no_grad(): for batch in train_loader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) logits teacher_model(input_ids, attention_maskattention_mask).logits all_logits.append(logits.cpu().numpy()) all_labels.append(batch[labels].numpy()) cache_path.parent.mkdir(exist_okTrue) np.save(cache_path, np.concatenate(all_logits, axis0)) np.save(label_path, np.concatenate(all_labels, axis0))参数说明这段代码把教师的logits保存为numpy数组文件大小取决于样本量。10万条样本、10个类别时logits文件大约只有几十MB完全可以在显存不足的低配机器上先跑教师推理。离线蒸馏后学生训练循环里不再需要teacher_model直接把缓存文件通过Dataset加载学生的batch_size可以回升到64甚至128。要注意的是离线缓存的数据必须和训练数据严格同顺序否则学生样本和教师logits错位训练结果会非常诡异。6. 验证与进阶温度退火、效果对比与迁移到更多中文任务蒸馏训练收尾时有几个动作能让模型水平再往上走一点。6.1 温度退火训练后期把T降回1我在训练最后20%的epoch里会把温度从3线性降到1同时把alpha逐渐提高到0.8。这样做的目的是让学生先学教师的分布关系最后再回归到硬标签的精确决策边界避免推理阶段学生一直处在“犹豫”状态。实现上只需要在训练循环的step里根据总步数和当前步数算一个衰减系数。这个技巧不需要额外代码但收益很稳定。6.2 怎么量化蒸馏值不值一张表说话我习惯在蒸馏结束后把学生模型、教师模型、以及一个不加蒸馏训练的TextCNN放在一起对比。对比维度是精度、模型参数量、CPU单条推理延迟。如果学生模型精度接近教师但延迟只有教师的十分之一这个项目实践就可以收尾。模型Accuracy参数量CPU推理延迟BERT教师92.3%102M45msTextCNN 蒸馏90.1%1.2M3.2msTextCNN 无蒸馏86.8%1.2M3.1ms这是我做过的一个新闻四分类项目里的典型结果数值会随数据集变化但三条趋势是稳定的蒸馏帮助学生涨点明显教师依然精度最高但延迟和体积不可接受。如果你的学生模型蒸馏后和教师差距在2~3个点以内这个压缩就是值的。6.3 迁移到情感分析与NER以及转ONNX部署蒸馏的套路可以平移到其他中文任务。情感分析基本可以复用整套代码只需要把分类头改成2类或3类。序列标注任务会麻烦一些因为KL散度要作用到每个token位置并且需要处理Label Padding的对齐。部署时我一般会把学生模型转成ONNX在CPU上用ONNXRuntime跑TextCNN转ONNX非常顺这也是我选它的原因之一。我自己做蒸馏最大的教训是温度不是一个“越大越好”的旋钮它更像是教师的语调调错了学生就学歪。现在每次跑蒸馏我都会先打印一批教师的软标签再决定T和alpha先不急着堆训练时长先把教师信号的质量检查完。希望帮到你。本文还有配套的精品资源点击获取
返回列表