PyTorch torch.mean Function
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))
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])
Other Extensions