def get_prob_of_contradiction(logits: torch.Tensor) -> torch.Tensor: """ Returns probability of contradiction aka factual inconsistency. Args: logits (torch.Tensor): Tensor of shape (batch_size, 3). The second dimension represents the probabilities of contradiction, neutral, and entailment. Returns: torch.Tensor: Tensor of shape (batch_size,) with probability of contradiction. Note: This function assumes the probability of contradiction is in index 0 of logits. """ # Drop neutral logit (index=1), softmax, and get prob of contradiction (index=0) prob = F.softmax(logits[:, [0, 2]], dim=1)[:, 0] return prob