PyTorch torch.nn.LogSoftmax Function
PyTorch torch.nn Reference Manual
torch.nn.LogSoftmaxIt is the Log Softmax activation function in PyTorch.
It is the logarithmic form of Softmax, which is numerically more stable and often used with NLLLoss.
Function Definition
torch.nn.LogSoftmax(dim=None)
Formula
LogSoftmax(x_i) = log(exp(x_i) / sum(exp(x_j)))
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
log_softmax = nn.LogSoftmax(dim=1)
logits = torch.tensor([[2.0, 1.0, 0.1]])
log_probs = log_softmax(logits)
print("Logits:", logits.tolist())
print("Log Softmax:", log_probs.tolist())
print("After exp:", log_probs.exp().tolist())
import torch.nn as nn
log_softmax = nn.LogSoftmax(dim=1)
logits = torch.tensor([[2.0, 1.0, 0.1]])
log_probs = log_softmax(logits)
print("Logits:", logits.tolist())
print("Log Softmax:", log_probs.tolist())
print("After exp:", log_probs.exp().tolist())
Example 2: With NLLLoss
Example
import torch
import torch.nn as nn
# Classification task
logits = torch.randn(4, 10)
targets = torch.tensor([2, 5, 1, 7])
# LogSoftmax + NLLLoss = CrossEntropyLoss
loss = nn.NLLLoss()(nn.LogSoftmax(dim=1)(logits), targets)
print("NLL Loss:", loss.item())
# Equivalent to
loss2 = nn.CrossEntropyLoss()(logits, targets)
print("CrossEntropyLoss:", loss2.item())
import torch.nn as nn
# Classification task
logits = torch.randn(4, 10)
targets = torch.tensor([2, 5, 1, 7])
# LogSoftmax + NLLLoss = CrossEntropyLoss
loss = nn.NLLLoss()(nn.LogSoftmax(dim=1)(logits), targets)
print("NLL Loss:", loss.item())
# Equivalent to
loss2 = nn.CrossEntropyLoss()(logits, targets)
print("CrossEntropyLoss:", loss2.item())
Example 3: Numerical Stability
Example
import torch
import torch.nn as nn
# Large logits
logits = torch.tensor([[1000, 1001, 1002]])
# Softmax may overflow
try:
sm = nn.Softmax(dim=1)(logits)
print("Softmax:", sm)
except:
print("Softmax overflow")
# LogSoftmax is numerically stable
lsm = nn.LogSoftmax(dim=1)(logits)
print("LogSoftmax:", lsm)
import torch.nn as nn
# Large logits
logits = torch.tensor([[1000, 1001, 1002]])
# Softmax may overflow
try:
sm = nn.Softmax(dim=1)(logits)
print("Softmax:", sm)
except:
print("Softmax overflow")
# LogSoftmax is numerically stable
lsm = nn.LogSoftmax(dim=1)(logits)
print("LogSoftmax:", lsm)
Use Cases
- Classification tasks: With NLLLoss
- Numerical stability: Large logit values
Tip: LogSoftmax + NLLLoss is equivalent to CrossEntropyLoss.
Other Extensions