logits = model(**inputs).logits # neutral is already removed probs = torch.softmax(logits, dim=-1) prob_ = probs[0][1].item() # prob(contradiction)