| 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") |
| 84 | kl_loss = F.kl_div(model_output_log_prob, real_output_soft, reduction="none") |
| 85 | kl_loss = kl_loss.sum(-1) * mask |
| 86 | kl_loss = kl_loss.sum(-1).mean() |
| 87 | return kl_loss |
| 88 | |
| 89 | def re_loss(self, labels, student_logits, teacher_logits): |
| 90 | mask = (labels != -100) |