MCPcopy Create free account
hub / github.com/LUMIA-Group/MemoryDecoder / kl_loss_token

Function kl_loss_token

utils/cal_loss.py:49–75  ·  view source on GitHub ↗
(logits, batch, tokenizer, args, knn_label, knn_prob, alpha=0.5)

Source from the content-addressed store, hash-verified

47 return nll_loss, lm_loss, token_num
48
49def 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

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected