PyTorch torch.logcumsumexp Function
PyTorch torch Reference Manual
torch.logcumsumexpis a function in PyTorch used to compute the logarithm of the cumulative sum of exponentials. It first exponentiates the elements, then performs a cumulative sum, and finally takes the logarithm, which can avoid numerical overflow.
Function Definition
torch.logcumsumexp(input, dim)
Parameter Description:
input: input tensordim: the dimension for cumulative summation
Usage Example
Example
import torch
# Create tensor
x = torch.tensor([1.0, 2.0, 3.0])
# Compute the logarithm of the cumulative sum of exponentials
y = torch.logcumsumexp(x, dim=0)
print(y)
# Create tensor
x = torch.tensor([1.0, 2.0, 3.0])
# Compute the logarithm of the cumulative sum of exponentials
y = torch.logcumsumexp(x, dim=0)
print(y)
The output result is:
tensor([1.0000, 2.3133, 3.1000])
Example
import torch
# Verify: log(exp(1) + exp(2)) = log(exp(1) + exp(2))
# log(e^1 + e^2) = log(e^1 + e^2) ≈ 2.3133
x = torch.tensor([1.0, 2.0])
y = torch.logcumsumexp(x, dim=0)
# Compare with direct calculation
import math
direct = math.log(math.exp(1) + math.exp(2))
print("logcumsumexp:", y[1].item())
print("direct:", direct)
# Verify: log(exp(1) + exp(2)) = log(exp(1) + exp(2))
# log(e^1 + e^2) = log(e^1 + e^2) ≈ 2.3133
x = torch.tensor([1.0, 2.0])
y = torch.logcumsumexp(x, dim=0)
# Compare with direct calculation
import math
direct = math.log(math.exp(1) + math.exp(2))
print("logcumsumexp:", y[1].item())
print("direct:", direct)
The output result is:
logcumsumexp: 2.313261505126953 direct: 2.313261505126953
Other Extensions