logits = torch.tensor([ [2.0, 1.0, 0.1, -1.0], # sample 0, correct class 0 [0.5, 2.5, 0.3, 0.2], # sample 1, correct class 1 [0.1, 0.2, 3.0, 0.1], # sample 2, correct class 2 ]) labels = torch.tensor([0, 1, 2]) print(nn.CrossEntropyLoss()(logits, labels)) # tensor(0.3015)