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

Function kl_loss_evaluate

utils/cal_loss.py:19–47  ·  view source on GitHub ↗
(logits, batch, tokenizer, args, knn_label, knn_prob)

Source from the content-addressed store, hash-verified

17 return interpolated
18
19def kl_loss_evaluate(logits, batch, tokenizer, args, knn_label, knn_prob):
20 label_probs = knn_prob
21
22 shift_logits = logits[:, :-1].contiguous() # (batch, seq_len-1, vocab_size)
23 shift_labels = batch['labels'][:, 1:].contiguous() # (batch, seq_len-1)
24
25 nonpad_mask = shift_labels != -100
26 shift_logits = shift_logits[nonpad_mask] # (nonpad b*t, vocab_size)
27 shift_labels = shift_labels[nonpad_mask] # (nonpad b*t)
28 label_probs = label_probs / label_probs.sum(dim=-1, keepdim=True) # Normalize label_probs
29
30 # Ensure that the dimensions match
31 assert shift_logits.shape == label_probs.shape, f"shift_logits.shape = {shift_logits.shape}, label_probs.shape = {label_probs.shape}"
32 assert torch.all(shift_labels == knn_label), f"shift_labels and knn_label are not the same"
33 assert torch.allclose(label_probs.sum(dim=-1), torch.ones_like(label_probs.sum(dim=-1))), f"label_probs does not sum to 1"
34
35 # Compute the label_probs
36 shift_probs = F.softmax(shift_logits, dim=-1)
37
38 # Calculate PPL
39 label_log_probs = label_probs.log()
40 label_log_probs = torch.nan_to_num(label_log_probs, nan=None, neginf=-10000.0)
41 lm_log_probs = F.log_softmax(shift_logits, dim=-1)
42 interpolate_log_probs = interpolate(label_log_probs, lm_log_probs, lmbda=args.lmbda)
43 nll_loss = F.nll_loss(interpolate_log_probs, shift_labels, reduction='sum')
44 lm_loss = F.nll_loss(lm_log_probs, shift_labels, reduction='sum')
45 token_num = shift_labels.shape[0]
46
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

Callers 1

mainFunction · 0.90

Calls 1

interpolateFunction · 0.70

Tested by

no test coverage detected