Метод repeat_interleave
Метод repeat_interleave повторяет элементы тензора указанное число раз.
В отличие от метода repeat, который дублирует весь тензор целиком,
repeat_interleave работает с каждым элементом индивидуально.
Метод принимает количество повторений для каждого элемента, а также
может принимать размерность, вдоль которой производится повторение.
Синтаксис
t.repeat_interleave(repeats, [dim])
Параметры
Метод принимает два параметра:
-
repeats- количество повторений для каждого элемента. Может быть целым числом (все элементы повторяются одинаково) или тензором (для каждого элемента своё количество повторений); -
dim(необязательный) - размерность, вдоль которой производится повторение. Если не указан, тензор сплющивается.
Метод возвращает новый тензор с повторёнными элементами.
Пример с одинаковым количеством повторений
Давайте повторим каждый элемент тензора по два раза:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = t.repeat_interleave(2)
print(res)
Результат выполнения кода:
tensor([1, 1, 2, 2, 3, 3, 4, 4, 5, 5])
Пример с разным количеством повторений
Давайте зададим для каждого элемента своё количество повторений с помощью тензора:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
repeats = torch.tensor([1, 2, 3, 2, 1])
res = t.repeat_interleave(repeats)
print(res)
Результат выполнения кода:
tensor([1, 2, 2, 3, 3, 3, 4, 4, 5])
Пример с указанием размерности
Давайте повторим элементы вдоль строк двумерного тензора:
import torch
t = torch.tensor([[1, 2, 3], [4, 5, 6]])
res = t.repeat_interleave(2, dim=0)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[1, 2, 3],
[4, 5, 6],
[4, 5, 6],
])
Пример с разными повторениями вдоль размерности
Давайте повторим строки двумерного тензора разное количество раз:
import torch
t = torch.tensor([[1, 2, 3], [4, 5, 6]])
repeats = torch.tensor([2, 1])
res = t.repeat_interleave(repeats, dim=0)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[1, 2, 3],
[4, 5, 6],
])
Пример использования в нейронных сетях
Давайте используем repeat_interleave для расширения
скрытого состояния перед применением свёртки:
import torch
import torch.nn as nn
torch.manual_seed(0)
t = torch.randn(2, 3)
print("Исходный тензор:", t)
res = t.repeat_interleave(2, dim=0)
print("После повторения:", res)
Результат выполнения кода:
"Исходный тензор: tensor([[ 1.5410, -0.2934, -2.1788], [ 0.5684, -1.0845, -1.3986]])"
"После повторения: tensor([[ 1.5410, -0.2934, -2.1788], [ 1.5410, -0.2934, -2.1788], [ 0.5684, -1.0845, -1.3986], [ 0.5684, -1.0845, -1.3986]])"