PyTorch torch.sum Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.sumis a function in PyTorch used to compute the sum of tensor elements. It can compute the sum of all elements, or along specified dimensions.

This is a commonly used reduction operation in deep learning, used in scenarios such as loss calculation and statistics.

Function Definition

torch.sum(input, dim, keepdim, dtype, out)

Parameters:

  • input(Tensor): The input tensor.
  • dim(int or tuple of int, optional): The dimension(s) to compute. IfNone, then compute the sum of all elements.
  • keepdim(bool, optional): Whether to keep the dimension. Default isFalse。
  • dtype(torch.dtype, optional): The data type of the output tensor.
  • out(Tensor, optional): The output tensor.

Return Value:

  • torch.Tensor: Returns the computed tensor.

Usage Examples

Example 1: Sum of All Elements

Example

import torch

# Create tensor
x = torch.tensor([[1, 2, 3], [4, 5, 6]])

# Compute sum of all elements
total = torch.sum(x)

print("Tensor:")
print(x)
print("Sum of elements:", total)

The output result is:

张量:
tensor([[1, 2, 3],
        [4, 5, 6]])
元素之和: tensor(21)

Example 2: Sum Along a Specified Dimension

Example

import torch

x = torch.tensor([[1, 2, 3], [4, 5, 6]])

# Sum along dim=0 (columns)
sum_dim0 = torch.sum(x, dim=0)
print("Sum along dim=0:", sum_dim0)

# Sum along dim=1 (rows)
sum_dim1 = torch.sum(x, dim=1)
print("Sum along dim=1:", sum_dim1)

The output result is:

沿 dim=0 求和: tensor([5, 7, 9])
沿 dim=1 求和: tensor([ 6, 15])

Example 3: Using keepdim to Preserve Dimensions

Example

import torch

x = torch.tensor([[1, 2, 3], [4, 5, 6]])

# Without keeping dimension
sum1 = torch.sum(x, dim=0)
print("Without keeping dimension:", sum1.shape)

# Keep dimension
sum2 = torch.sum(x, dim=0, keepdim=True)
print("Keeping dimension:", sum2.shape)
print(sum2)

The output result is:

不保持维度: torch.Size([3])
保持维度: torch.Size([1, 3])
tensor([[5, 7, 9]])

Example 4: Computing Loss in Neural Networks

Example
import torch

# Simulate predicted and true values
predictions = torch.tensor([0.1, 0.9, 0.8, 0.3])
targets = torch.tensor([0.0, 1.0, 1.0, 0.0])

# Compute mean squared error loss
loss = torch.sum((predictions - targets) ** 2) / len(predictions)

print("MSE Loss:", loss.item())

The output result is:

MSE 损失: 0.07499999690771103

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions