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

Method cakld_loss

train/mytrainer.py:38–57  ·  view source on GitHub ↗
(self, labels, student_logits, teacher_logits, beta_prob)

Source from the content-addressed store, hash-verified

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)

Callers 1

compute_lossMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected