| 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) |