| 24 | mse_loss = MSELoss() |
| 25 | |
| 26 | class 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") |