PyTorch torch.logcumsumexp Function


Pytorch torch 参考手册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 tensor
  • dim: 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)

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)

The output result is:

logcumsumexp: 2.313261505126953
direct: 2.313261505126953

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions