Функция round
Функция torch.round применяется к тензору и возвращает новый тензор, в котором каждый элемент округлён до ближайшего целого числа. Если дробная часть числа равна ровно 0.5, то округление выполняется до ближайшего чётного целого числа (банковское округление). Функция принимает тензор в качестве первого аргумента. Вторым параметром можно указать размерность, по которой выполняется операция.
Синтаксис
torch.round(tensor, [decimals])
Пример
Давайте округлим элементы тензора до ближайшего целого числа:
import torch
t = torch.tensor([0.3, 1.7, 2.5, 3.1, 4.5])
res = torch.round(t)
print(res)
Результат выполнения кода:
tensor([0., 2., 2., 3., 4.])
Пример
Продемонстрируем правило округления «половина к чётному» для чисел с дробной частью ровно 0.5:
import torch
t = torch.tensor([1.5, 2.5, 3.5, 4.5])
res = torch.round(t)
print(res)
Результат выполнения кода:
tensor([2., 2., 4., 4.])
Пример
Функцию также можно вызывать как метод тензора для двумерных данных:
import torch
t = torch.tensor([
[0.6, 1.3, 2.5],
[3.4, 4.8, 5.5]
])
res = t.round()
print(res)
Результат выполнения кода:
tensor([
[1., 1., 2.],
[3., 5., 6.]
])