PyTorch torch.nn.LogSoftmax Function

PyTorch torch.nn 参考手册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())

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())

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)

Use Cases

  • Classification tasks: With NLLLoss
  • Numerical stability: Large logit values

Tip: LogSoftmax + NLLLoss is equivalent to CrossEntropyLoss.


PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual

Other Extensions