PyTorch torch.nn.Softmax Function
PyTorch torch.nn Reference Manual
torch.nn.Softmaxis the Softmax activation function in PyTorch.
It converts input into a probability distribution, with all outputs summing to 1.
Function Definition
torch.nn.Softmax(dim=None)
Parameters:
dim: The dimension along which Softmax is performed
Formula
Softmax(x_i) = exp(x_i) / sum(exp(x_j))
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
softmax = nn.Softmax(dim=1)
logits = torch.tensor([[2.0, 1.0, 0.1], [1.0, 2.0, 0.5]])
probs = softmax(logits)
print("Logits:", logits.tolist())
print("Probabilities:", probs.tolist())
print("Row sums:", probs.sum(dim=1).tolist())
import torch.nn as nn
softmax = nn.Softmax(dim=1)
logits = torch.tensor([[2.0, 1.0, 0.1], [1.0, 2.0, 0.5]])
probs = softmax(logits)
print("Logits:", logits.tolist())
print("Probabilities:", probs.tolist())
print("Row sums:", probs.sum(dim=1).tolist())
Example 2: dim Parameter
Example
import torch
import torch.nn as nn
# 3D input
x = torch.randn(2, 3, 4)
# softmax along different dimensions
print("dim=1:", nn.Softmax(dim=1)(x).sum(dim=1)[:1])
print("dim=2:", nn.Softmax(dim=2)(x).sum(dim=2)[:1])
print("dim=-1:", nn.Softmax(dim=-1)(x).sum(dim=-1)[:1])
import torch.nn as nn
# 3D input
x = torch.randn(2, 3, 4)
# softmax along different dimensions
print("dim=1:", nn.Softmax(dim=1)(x).sum(dim=1)[:1])
print("dim=2:", nn.Softmax(dim=2)(x).sum(dim=2)[:1])
print("dim=-1:", nn.Softmax(dim=-1)(x).sum(dim=-1)[:1])
Example 3: Classification Output
Example
import torch
import torch.nn as nn
model = nn.Linear(784, 10)
# Output logits
logits = model(torch.randn(4, 784))
# Convert to probabilities
probs = nn.Softmax(dim=1)(logits)
print("Predicted class:", probs.argmax(dim=1).tolist())
print("Highest probability:", probs.max(dim=1).values.tolist())
import torch.nn as nn
model = nn.Linear(784, 10)
# Output logits
logits = model(torch.randn(4, 784))
# Convert to probabilities
probs = nn.Softmax(dim=1)(logits)
print("Predicted class:", probs.argmax(dim=1).tolist())
print("Highest probability:", probs.max(dim=1).values.tolist())
Use Cases
- Multi-class classification output: Probability distribution
- Attention mechanism
- Probabilistic models
Note: The values after Softmax are all positive and sum to 1.
Other extensions