def get_pointwise_loss(self, inputs: List[str], tok): self.model.eval() all_probs = [] with torch.inference_mode(): for input in inputs: ids_list: List[int] = tok.encode(input).ids # ids has shape (1, len(ids_list)) ids: Torch.Tensor = torch.tensor(ids_list, device=self.device).unsqueeze(0) # probs below is the probability that the token at that location # completes the sentence (in ids) so far. y = self.model(ids) criterion = nn.CrossEntropyLoss(reduction='none', ignore_index=0) # Compute the loss starting from the 2nd token in the model's output. loss = criterion(y[:,:,:-1], ids[:,1:]) # To compute the probability of each token, we need to compute the # negative of the log loss and exponentiate it. loss = loss * -1.0 # Set the probabilities that we are not interested in to -inf. # This is done to make the softmax set these values to 0.0 loss[loss == 0.0] = float("-inf") # probs holds the probability of each token's prediction # starting from the 2nd token since we don't want to include # the probability of the model predicting the beginning of # a sentence given no existing sentence context. # # To compute perplexity, we should probably ignore the first # handful of predictions of the model since there's insufficient # context. We don't do that here, though. probs = loss.exp() all_probs.append(probs) # # return all_probs #