PyTorch torch.nn.NLLLoss Function
PyTorch torch.nn Reference Manual
torch.nn.NLLLossIs the negative log-likelihood loss in PyTorch.
It is used for multi-class classification tasks and needs to be used with LogSoftmax.
Function Definition
torch.nn.NLLLoss(weight=None, ignore_index=-100, reduction='mean')
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# First LogSoftmax, then NLLLoss
log_softmax = nn.LogSoftmax(dim=1)
nll = nn.NLLLoss()
# Model output (logits)
logits = torch.randn(4, 10)
# After log softmax
log_probs = log_softmax(logits)
# True labels
targets = torch.tensor([2, 5, 1, 7])
loss = nll(log_probs, targets)
print("NLL Loss:", loss.item())
import torch.nn as nn
# First LogSoftmax, then NLLLoss
log_softmax = nn.LogSoftmax(dim=1)
nll = nn.NLLLoss()
# Model output (logits)
logits = torch.randn(4, 10)
# After log softmax
log_probs = log_softmax(logits)
# True labels
targets = torch.tensor([2, 5, 1, 7])
loss = nll(log_probs, targets)
print("NLL Loss:", loss.item())
Example 2: Equivalent to CrossEntropyLoss
Example
import torch
import torch.nn as nn
logits = torch.randn(4, 10)
targets = torch.tensor([2, 5, 1, 7])
# Method 1: Use CrossEntropyLoss directly
loss1 = nn.CrossEntropyLoss()(logits, targets)
# Method 2: LogSoftmax + NLLLoss
loss2 = nn.NLLLoss()(nn.LogSoftmax(dim=1)(logits), targets)
print("CrossEntropyLoss:", loss1.item())
print("LogSoftmax + NLLLoss:", loss2.item())
print("Results are the same:", abs(loss1 - loss2) < 1e-5)
import torch.nn as nn
logits = torch.randn(4, 10)
targets = torch.tensor([2, 5, 1, 7])
# Method 1: Use CrossEntropyLoss directly
loss1 = nn.CrossEntropyLoss()(logits, targets)
# Method 2: LogSoftmax + NLLLoss
loss2 = nn.NLLLoss()(nn.LogSoftmax(dim=1)(logits), targets)
print("CrossEntropyLoss:", loss1.item())
print("LogSoftmax + NLLLoss:", loss2.item())
print("Results are the same:", abs(loss1 - loss2) < 1e-5)
Use Cases
- Multi-class classification: Used with LogSoftmax
- Custom loss: Special requirements
Note: CrossEntropyLoss already includes LogSoftmax, so it is usually used directly.
Other Extensions