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

Method ce_loss

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

Source from the content-addressed store, hash-verified

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)

Callers 1

compute_lossMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected