| 47 | return nll_loss, lm_loss, token_num |
| 48 | |
| 49 | def kl_loss_token(logits, batch, tokenizer, args, knn_label, knn_prob, alpha=0.5): |
| 50 | label_probs = knn_prob |
| 51 | |
| 52 | shift_logits = logits[:, :-1].contiguous() # (batch, seq_len-1, vocab_size) |
| 53 | shift_labels = batch['labels'][:, 1:].contiguous() # (batch, seq_len-1) |
| 54 | |
| 55 | nonpad_mask = shift_labels != -100 |
| 56 | shift_logits = shift_logits[nonpad_mask] # (nonpad b*t, vocab_size) |
| 57 | shift_labels = shift_labels[nonpad_mask] # (nonpad b*t) |
| 58 | label_probs = label_probs / label_probs.sum(dim=-1, keepdim=True) # Normalize label_probs |
| 59 | |
| 60 | # Ensure that the dimensions match |
| 61 | assert shift_logits.shape == label_probs.shape, f"shift_logits.shape = {shift_logits.shape}, label_probs.shape = {label_probs.shape}" |
| 62 | assert torch.all(shift_labels == knn_label), f"shift_labels and knn_label are not the same" |
| 63 | assert torch.allclose(label_probs.sum(dim=-1), torch.ones_like(label_probs.sum(dim=-1))), f"label_probs does not sum to 1" |
| 64 | |
| 65 | # MemDec loss |
| 66 | kl_loss = F.kl_div(F.log_softmax(shift_logits, dim=-1), label_probs, reduction='batchmean') |
| 67 | |
| 68 | loss_fct = nn.CrossEntropyLoss() |
| 69 | lm_loss = loss_fct(shift_logits, shift_labels) |
| 70 | |
| 71 | total_loss = alpha * kl_loss + (1 - alpha) * lm_loss |
| 72 | |
| 73 | logger.info(f"KL loss: {kl_loss} LM loss: {lm_loss} Total loss: {total_loss}") |
| 74 | |
| 75 | return total_loss, kl_loss, lm_loss |