PyTorch torch.mean Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.meanIt is a function in PyTorch used to compute the mean of a tensor.

Function Definition

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

Usage Examples

Example

import torch

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

# Calculate the mean of all elements
print("Global mean:", torch.mean(x))

# Calculate the mean along dim=0
print("dim=0 mean:", torch.mean(x, dim=0))

The output result is:

全局均值: tensor(3.5000)
dim=0 均值: tensor([2.5000, 3.5000, 4.5000])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions