| 17 | return interpolated |
| 18 | |
| 19 | def 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 | |
| 49 | def kl_loss_token(logits, batch, tokenizer, args, knn_label, knn_prob, alpha=0.5): |
| 50 | label_probs = knn_prob |