模型压缩
在对模型的效果影响不大的前提下,尽可能减小模型的占用空间,减少模型参数量,让模型变得简单,提升运行速度。
分类:
- 模型量化
- 模型蒸馏
- 模型剪枝
- 低秩分解
四种方式简介
- Pruning 剪枝:深度神经网络中许多参数在训练期间贡献不大、属于冗余参数,训练后可将这些参数从网络中删除,对准确性影响很小。
- Quantization 量化:DNN 权重通常存储为 32 位浮点数(fp32),量化通过减少每个权重所需的位数(如 fp16、int8、int4)来压缩原始网络,可显著减小模型大小并加速推理。
- Knowledge distillation 知识蒸馏:在大型数据集上训练好大型复杂模型(教师模型)使其具备泛化与推理能力后,将其知识转移到较小的网络(学生模型)。
- Low-rank factorization 低秩因式分解:通过矩阵分解识别深度神经网络的冗余参数,将大型矩阵分解为较小的矩阵,从而减小模型大小。
模型蒸馏
模型压缩详细文章:模型蒸馏
相关概念:
- 教师模型:是一个大模型,参数量大,预测精确,但是运行速度慢,模型体积大
- 学生模型:是一个小模型,参数量少,预测相对精确,但是运行速度快,模型体积小
- 硬标签:模型预测结果的概率值,非0即1。学生模型只能学到有限的内容,只能学到一个最终的答案
- 软标签:模型预测结果的概率值,取值区间在[0,1]之间。学生模型能够学到预测结果的概率分布
- 中间层:模型的思考过程
KL散度介绍【面试问题】:
- 作用:衡量两个预测概率分布之间的差异性
- 规律:该值越小越好。最小为0,那么说明学生模型和教师模型的结果完全一样,蒸馏效果最好
- 扩展:硬标签中 KL散度值=交叉熵,因为信息熵带入预测分布概率进去H(P)结果是0
# 1- 教师模型的:嵌入层(词嵌入层、片段编码、位置编码)、Encoder编码器
with torch.no_grad():
bert_output = self.bert_model(input_ids=input_ids,attention_mask=attention_mask)
##--------------------------------------------------------------------------
# embedding_dim:由我们自己设置,与教师模型没有任何关系
self.embedding = nn.Embedding(num_embeddings=self.vocab_size,embedding_dim=self.embedding_dim)
#---------------------------------------------------------------------------
# 3.2- 循环神经网络层:BiLSTM
self.lstm = nn.LSTM(
input_size=self.embedding_dim,
hidden_size=self.hidden_size,
num_layers=self.num_layers,
batch_first=True,
bidirectional=True
)
#---------------------------------------------------------------------------
# 6.5- 计算KL散度值
p = torch.log_softmax(teacher_pred/T, dim=-1) # 老师软化后的知识
q = torch.log_softmax(student_pred/T, dim=-1) # 学生软化后的预测
kl_value = torch.nn.functional.kl_div(
input=q, target=p, reduction="batchmean", log_target=True
)
#---------------------------------------------------------------------------
# 6.7- 蒸馏的总损失值
distll_loss = (1 - alpha) * hard_label_loss + alpha * (T**2) * kl_value
剪枝
# 2.2- 进行全局非结构化剪枝
"""
参数解释:
parameters:剪枝范围。也就是规定对Bert模型中什么层的什么参数进行剪枝
pruning_method:剪枝方式
amount:剪枝的参数情况。有两种类的参数值,如下
整数:表示具体对多少个参数剪枝
小数:表示对Bert模型中多大比例的参数进行剪枝。推荐
"""
prune.global_unstructured(
parameters=parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=0.3
)