Функция chunk
Функция chunk разделяет тензор на несколько подтензоров вдоль указанной размерности.
Первым параметром функция принимает входной тензор.
Вторым параметром указывается количество частей, на которые нужно разбить тензор.
Третьим параметром можно указать размерность, по которой выполняется разделение (по умолчанию dim=0).
Функция возвращает кортеж из подтензоров.
Синтаксис
torch.chunk(tensor, chunks, dim=0)
Пример
Давайте разделим одномерный тензор из 5 элементов на 2 части:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = torch.chunk(t, 2)
print(res)
Результат выполнения кода:
(tensor([1, 2, 3]), tensor([4, 5]))
Пример
Разделим двумерный тензор на 3 части по строкам (размерность dim=0):
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12],
])
res = torch.chunk(t, 3, dim=0)
for i, part in enumerate(res):
print(f"Part {i}: {part}")
Результат выполнения кода:
Part 0: tensor([[1, 2, 3],
[4, 5, 6]])
Part 1: tensor([[7, 8, 9]])
Part 2: tensor([[10, 11, 12]])
Пример
Разделим тензор по столбцам (размерность dim=1):
import torch
t = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
res = torch.chunk(t, 2, dim=1)
for i, part in enumerate(res):
print(f"Part {i}:\n{part}")
Результат выполнения кода:
Part 0:
tensor([[1, 2],
[5, 6]])
Part 1:
tensor([[3, 4],
[7, 8]])
Пример
Если размерность не делится нацело на количество частей, последняя часть будет меньше:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6, 7])
res = torch.chunk(t, 3)
for i, part in enumerate(res):
print(f"Part {i}: {part}")
Результат выполнения кода:
Part 0: tensor([1, 2, 3])
Part 1: tensor([4, 5])
Part 2: tensor([6, 7])