我们在前文已经提及了教师模型和学生型的搭建及各自的前向传播函数,现在我们该根据之前的思路去构建总体损失函数。
其中,student_model 和 teacher_model 是教师模型和学生模型的实例化对象。
student_model.train()
teacher_model.eval()
输入:input_ids,attention_mask
在计算教师损失的时候需要把模型设置为eval()模式,并冻结参数关闭梯度计算。
T(温度)设置为 2 。下面是超参设置
import torch.nn.functional as F
optimizer = AdamW(student_model.parameters(), lr=conf.learning_rate) # 使用 AdamW 优化器
criterion = nn.CrossEntropyLoss() # 交叉熵损失,用于硬标签损失
T = 2.0
train_dataloader, test_dataloader, dev_dataloader = build_dataloader()
alpha = 0.7 # 软标签和硬标签损失的权重
step = 0 # 训练步数计数器
best_dev_f1 = 0.0 # 记录最佳验证 F1 分数
核心损失计算:
先得到教师模型前向传播结果:
输入进入教师模型后我们会拿到模型输出的逻辑值,如例:[[5, 1, 0], [2, 4, 1]] (不是真实内容,方便理解编造的)
teacher_logits = teacher_model(input_ids, attention_mask)
逻辑值/T 放缩后丢入softmax层计算概率得到 [[0.975559, 0.017868, 0.006573],[0.114195, 0.843795, 0.042010]] ,之后在最后一维做argmax拿到教师预测的硬标签 [0,1]
teacher_soft_probs = F.softmax(teacher_logits / T, dim=1)
teacher_preds = torch.argmax(teacher_soft_probs, dim=1)
再次强调,以上所有内容在with torch.no_grad(): 里进行
再拿到学生模型前向传播结果:
student_logits = student_model(input_ids, attention_mask)
# 分支1:高温soft分支(用于KL蒸馏)
# student_log_probs 必须用log_softmax,适配kl_div输入要求,T 为蒸馏温度,师生 logits 都除以 T 平滑分布
student_log_soft_probs = F.log_softmax(student_logits / T, dim=1)
# 分支2:常温T=1硬预测分支(用于CE硬损失)
student_hard_probs = F.softmax(student_logits, dim=1)
在计算软概率时外套了一层 log,用于平滑学生概率分布,使其不要过于尖锐(也就是某个类特别大,其他类特别小)
在计算学生模型的概率时做了两个分支处理,分别送入KL散度损失计算和CE损失计算。
最后在给每种损失赋予权重alpha。
# 损失计算
## KL蒸馏损失(soft loss,师生软分布对齐)
# 乘 T² 抵消温度缩放带来的梯度衰减,reduction统一均值
soft_loss = F.kl_div(input=student_log_soft_probs,
target=teacher_soft_probs,
reduction="batchmean"
) * (T ** 2)
## CE硬损失(hard loss,学生硬输出对齐真实标签y,和图完全匹配)
# criterion = nn.CrossEntropyLoss(),输入logits或概率均可
hard_loss = criterion(student_hard_probs, teacher_preds)
## 加权总损失 alpha为权重(示例图alpha=0.8)
loss = alpha * soft_loss + (1 – alpha) * hard_loss
这样我们就完成了最重要的损失计算部分。最后乘 T² 抵消温度缩放带来的梯度衰减
下面是损失计算的完整代码:
with torch.no_grad():
# 教师模型输出原始logits
teacher_logits = teacher_model(input_ids, attention_mask)
# 教师高温softmax得到soft labels(图中soft labels)
# [[5, 1, 0],
# [2, 4, 1]]
#
# to
#
# [[0.975559, 0.017868, 0.006573],
# [0.114195, 0.843795, 0.042010]]
teacher_soft_probs = F.softmax(teacher_logits / T, dim=1)
# teacher_preds:教师硬类别(argmax取最大概率)
teacher_preds = torch.argmax(teacher_soft_probs, dim=1)
# 学生模型前向,更新梯度
student_logits = student_model(input_ids, attention_mask)
# 分支1:高温soft分支(用于KL蒸馏)
# student_log_probs 必须用log_softmax,适配kl_div输入要求,T 为蒸馏温度,师生 logits 都除以 T 平滑分布
student_log_soft_probs = F.log_softmax(student_logits / T, dim=1)
# 分支2:常温T=1硬预测分支(用于CE硬损失)
student_hard_probs = F.softmax(student_logits, dim=1)
# 损失计算
## KL蒸馏损失(soft loss,师生软分布对齐)
# 乘 T² 抵消温度缩放带来的梯度衰减,reduction统一均值
soft_loss = F.kl_div(
input=student_log_soft_probs,
target=teacher_soft_probs,
reduction="batchmean"
) * (T ** 2)
## CE硬损失(hard loss,学生硬输出对齐真实标签y,和图完全匹配)
# criterion = nn.CrossEntropyLoss(),输入logits或概率均可
hard_loss = criterion(student_hard_probs, teacher_preds)
## 加权总损失 alpha为权重(示例图alpha=0.8)
loss = alpha * soft_loss + (1 – alpha) * hard_loss

