Функция sum
Функция sum вычисляет сумму элементов тензора.
Первым параметром она принимает тензор, для которого нужно вычислить сумму.
Вторым необязательным параметром можно указать размерность dim, вдоль которой будет производиться суммирование.
Третьим параметром можно передать keepdim, который сохраняет размерность тензора в результате.
Также функция поддерживает параметр dtype для указания типа данных результата.
Синтаксис
torch.sum(input, [dim], [keepdim], [dtype])
Пример
Давайте вычислим сумму всех элементов тензора чисел 1, 2, 3, 4, 5:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = torch.sum(t)
print(res)
Результат выполнения кода:
tensor(15)
Пример
Теперь вычислим сумму элементов двумерного тензора:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = torch.sum(t)
print(res)
Результат выполнения кода:
tensor(21)
Пример
Вычислим сумму элементов вдоль строк (размерность 0) с сохранением размерности:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = torch.sum(t, dim=0, keepdim=True)
print(res)
Результат выполнения кода:
tensor([
[5, 7, 9],
])
Пример
Вычислим сумму элементов вдоль столбцов (размерность 1):
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = torch.sum(t, dim=1)
print(res)
Результат выполнения кода:
tensor([6, 15])
Пример
Вычислим сумму с изменением типа данных на float:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = torch.sum(t, dtype=torch.float)
print(res)
Результат выполнения кода:
tensor(15.)