1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
| # 知识蒸馏
import torch
import torch.nn as nn
class DistillationTrainer:
def __init__(self, teacher_model, student_model, temperature=3.0):
self.teacher = teacher_model
self.student = student_model
self.temperature = temperature
self.teacher.eval()
def distillation_loss(
self,
student_output,
teacher_output,
labels,
alpha=0.5
):
"""蒸馏损失函数"""
# 软标签损失(来自教师)
soft_loss = nn.KLDivLoss(reduction='batchmean')(
nn.functional.log_softmax(student_output / self.temperature, dim=1),
nn.functional.softmax(teacher_output / self.temperature, dim=1)
) * (self.temperature ** 2)
# 硬标签损失(来自真实标签)
hard_loss = nn.CrossEntropyLoss()(student_output, labels)
# 组合损失
return alpha * soft_loss + (1 - alpha) * hard_loss
def train_step(self, inputs, labels, optimizer):
"""训练步骤"""
# 教师模型推理(不计算梯度)
with torch.no_grad():
teacher_output = self.teacher(inputs)
# 学生模型前向传播
student_output = self.student(inputs)
# 计算损失
loss = self.distillation_loss(
student_output,
teacher_output,
labels
)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss.item()
def train(self, train_loader, epochs, learning_rate):
"""训练学生模型"""
optimizer = torch.optim.Adam(
self.student.parameters(),
lr=learning_rate
)
for epoch in range(epochs):
total_loss = 0
for inputs, labels in train_loader:
loss = self.train_step(inputs, labels, optimizer)
total_loss += loss
avg_loss = total_loss / len(train_loader)
print(f"Epoch {epoch + 1}/{epochs}, Loss: {avg_loss:.4f}")
|