Метод expand
Метод expand применяется к тензору и возвращает новый тензор,
размерности которого расширены до указанных значений.
Важно, что метод не копирует данные, а создает новое представление
(вид) тензора, что позволяет экономить память.
Расширение возможно только для размерностей со значением 1.
Метод принимает в качестве аргументов целевые размеры тензора.
Синтаксис
t.expand(*sizes)
где *sizes - произвольное количество целых чисел,
задающих новые размерности тензора.
Пример
Давайте расширим одномерный тензор до двумерного:
import torch
t = torch.tensor([1, 2, 3])
res = t.expand(2, 3)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[1, 2, 3],
])
Как видно из примера, строки тензора были продублированы без создания новой копии данных.
Пример
Расширим двумерный тензор по первой размерности:
import torch
t = torch.tensor([[1, 2, 3]])
res = t.expand(3, -1)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[1, 2, 3],
[1, 2, 3],
])
Значение -1 во втором аргументе означает, что эта размерность останется без изменений.
Пример
Расширим тензор до трехмерного:
import torch
t = torch.tensor([[1, 2, 3]])
res = t.expand(2, 3, -1)
print(res)
Результат выполнения кода:
tensor([
[
[1, 2, 3],
[1, 2, 3],
[1, 2, 3],
],
[
[1, 2, 3],
[1, 2, 3],
[1, 2, 3],
],
])
В этом примере мы расширили тензор до формы (2, 3, 3).
Пример
Попытка расширить тензор по размерности, которая не равна 1, приведет к ошибке:
import torch
t = torch.tensor([[1, 2, 3], [4, 5, 6]])
res = t.expand(3, 3)
Результат выполнения кода:
RuntimeError: The expanded size of the tensor (3) must match the existing size (2) at non-singleton dimension 0.
Как видно из ошибки, метод expand может расширять только те размерности, размер которых равен 1.