Функция default_collate
Функция default_collate используется в загрузчике данных
DataLoader для объединения отдельных образцов в батч.
Она принимает список образцов и преобразует их в тензоры,
сохраняя структуру данных. Если образцы являются тензорами,
числами, строками или вложенными структурами, функция
рекурсивно обрабатывает каждый элемент.
Синтаксис
torch.utils.data.default_collate(batch)
Параметр batch - список образцов, полученных из датасета.
Функция возвращает объединённый батч, где каждый элемент имеет
дополнительное измерение (размер батча).
Пример
Давайте создадим простой список тензоров и объединим их в батч
с помощью default_collate:
import torch
from torch.utils.data import default_collate
batch = [
torch.tensor([1, 2, 3]),
torch.tensor([4, 5, 6]),
torch.tensor([7, 8, 9]),
]
res = default_collate(batch)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
Пример
Функция умеет обрабатывать вложенные структуры, например, кортежи из тензоров и чисел:
import torch
from torch.utils.data import default_collate
batch = [
(torch.tensor([1, 2]), 10),
(torch.tensor([3, 4]), 20),
(torch.tensor([5, 6]), 30),
]
res = default_collate(batch)
print(res[0])
print(res[1])
Результат выполнения кода:
tensor([
[1, 2],
[3, 4],
[5, 6],
])
tensor([10, 20, 30])
Пример
Если в списке образцов присутствуют строки, функция объединяет их в список, а не в тензор:
import torch
from torch.utils.data import default_collate
batch = [
{'text': 'abc', 'label': 0},
{'text': 'def', 'label': 1},
{'text': 'ghi', 'label': 2},
]
res = default_collate(batch)
print(res['text'])
print(res['label'])
Результат выполнения кода:
['abc', 'def', 'ghi']
tensor([0, 1, 2])
Пример
Если образцы имеют разные размеры, но являются тензорами,
default_collate попытается объединить их, если это
возможно. В противном случае возникнет ошибка:
import torch
from torch.utils.data import default_collate
batch = [
torch.tensor([1, 2, 3]),
torch.tensor([4, 5]),
torch.tensor([6]),
]
try:
res = default_collate(batch)
except RuntimeError as e:
print("Error:", str(e))
Результат выполнения кода:
"Error: each element in list of batch should be of equal size"
Смотрите также
-
класс
BatchSampler,
который создаёт батчи индексов для выборки -
класс
TensorDataset,
который упаковывает тензоры в датасет -
функцию
default_convert,
которая преобразует отдельный образец в тензор -
класс
DataLoader,
который используетdefault_collateдля сборки батчей