MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / KDTrainer

Class KDTrainer

train/mytrainer.py:26–166  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24mse_loss = MSELoss()
25
26class KDTrainer(Trainer):
27 def __init__(self, teacher_model, loss_type, mean_prob=0, *args, **kwargs):
28 super().__init__(*args, **kwargs)
29 # self.tlsd = tsld_loss
30 self.loss_fct_none = torch.nn.CrossEntropyLoss(reduction="none")
31 self.tmp = 1
32 self.teacher_model = teacher_model
33 # self.reverse_loss = reverse_loss
34 self.loss_type = loss_type
35 self.mean_prob = mean_prob
36 self.ce_loss_none = CrossEntropyLoss(reduction="none")
37
38 def cakld_loss(self, labels, student_logits, teacher_logits, beta_prob):
39 mask = (labels != -100)
40
41 # reverse
42 teacher_output_log_prob = F.log_softmax(teacher_logits, dim=2)
43 # Compute the softmax of the student's logits (approximate distribution)
44 student_output_soft = F.softmax(student_logits, dim=2)
45 # Calculate the reverse KL Divergence (KL(teacher_logits || student_logits))
46 reverse_kl = F.kl_div(teacher_output_log_prob, student_output_soft, reduction="none").sum(-1)
47
48 # forward
49 student_output_log_prob = F.log_softmax(student_logits, dim=2)
50 teacher_output_soft = F.softmax(teacher_logits, dim=2)
51 # Calculate the reverse KL Divergence (KL(teacher_logits || student_logits))
52 forward_kl = F.kl_div(student_output_log_prob, teacher_output_soft, reduction="none").sum(-1)
53
54 kl_loss = beta_prob * reverse_kl + (1 - beta_prob) * forward_kl
55 kl_loss *= mask
56 kl_loss = kl_loss.sum(-1).mean()
57 return kl_loss
58
59 def jsd_loss(self, labels, student_logits, teacher_logits, beta_prob):
60 mask = (labels != -100)
61 student_prob = F.softmax(student_logits, dim=2)
62 teacher_prob = F.softmax(teacher_logits, dim=2)
63
64 c_prob = beta_prob * teacher_prob + (1-beta_prob) * student_prob
65 c_log_prob = c_prob.log()
66
67
68 kl_loss_f = beta_prob * F.kl_div(c_log_prob, teacher_prob, reduction="none")
69 kl_loss_r = (1 - beta_prob) * F.kl_div(c_log_prob, student_prob, reduction="none")
70 kl_loss = kl_loss_f + kl_loss_r
71
72 kl_loss = kl_loss.sum(-1) * mask
73 kl_loss = kl_loss.sum(-1).mean()
74
75 return kl_loss
76
77 def ce_loss(self, labels, student_logits, teacher_logits):
78 mask = (labels != -100)
79
80 model_output_log_prob = F.log_softmax(student_logits, dim=2)
81 real_output_soft = F.softmax(teacher_logits / self.tmp, dim=2)
82
83 # loss = F.kl_div(model_output_log_prob, real_output_soft, reduction="batchmean")

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected