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