# wrong class Net(nn.Module): def forward(self, x): return torch.softmax(self.fc(x), dim=1) # then CrossEntropyLoss -> bug # right class Net(nn.Module): def forward(self, x): return self.fc(x) # raw logits, that is all # and at inference time, when you want probabilities to show a user: with torch.no_grad(): probs = torch.softmax(model(x), dim=1)