Функция broadcast_tensors
Функция broadcast_tensors приводит переданные тензоры к одинаковой форме в соответствии с правилами механизма broadcasting. Она принимает произвольное количество тензоров в качестве аргументов и возвращает кортеж тензоров, каждый из которых имеет одинаковую форму. Эта функция полезна, когда необходимо явно получить расширенные версии тензоров для дальнейших операций.
Синтаксис
torch.broadcast_tensors(*tensors)
Пример
Давайте приведём два тензора разных форм к общему размеру:
import torch
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([[1], [2], [3]])
res = torch.broadcast_tensors(t1, t2)
print(res[0])
print(res[1])
Результат выполнения кода:
tensor([
[1, 2, 3],
[1, 2, 3],
[1, 2, 3],
])
tensor([
[1, 1, 1],
[2, 2, 2],
[3, 3, 3],
])
Пример
Рассмотрим случай с тремя тензорами разных размерностей:
import torch
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([[1], [2]])
t3 = torch.tensor([[[1]]])
res = torch.broadcast_tensors(t1, t2, t3)
print(res[0].shape)
print(res[1].shape)
print(res[2].shape)
Результат выполнения кода:
torch.Size([2, 3, 1])
torch.Size([2, 3, 1])
torch.Size([2, 3, 1])
Пример
Попробуем привести тензоры, которые не могут быть расширены:
import torch
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([1, 2])
try:
res = torch.broadcast_tensors(t1, t2)
except RuntimeError as e:
print("Error:", str(e))
Результат выполнения кода:
"Error: The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 0"
Пример
Используем функцию для подготовки данных перед поэлементной операцией:
import torch
t1 = torch.tensor([1, 2, 3, 4])
t2 = torch.tensor([[10], [20]])
t1_b, t2_b = torch.broadcast_tensors(t1, t2)
res = t1_b + t2_b
print(res)
Результат выполнения кода:
tensor([
[11, 12, 13, 14],
[21, 22, 23, 24],
])
Смотрите также
-
функцию
broadcast_to,
которая приводит один тензор к указанной форме -
функцию
meshgrid,
которая создаёт координатные сетки из векторов -
функцию
reshape,
которая изменяет форму тензора -
функцию
atleast_1d,
которая приводит тензоры к как минимум одномерной форме