Атрибут collate_fn
Атрибут collate_fn класса DataLoader определяет функцию,
которая используется для объединения списка объектов, полученных из датасета,
в один мини-батч. По умолчанию эта функция преобразует список тензоров
в один тензор с дополнительным измерением (батч-измерением).
Однако, если ваши данные имеют нестандартную структуру, например,
содержат последовательности разной длины или сложные объекты,
вы можете передать собственную функцию для гибкой обработки.
Синтаксис
torch.utils.data.DataLoader(
dataset,
batch_size=1,
collate_fn=None,
...
)
Параметр collate_fn принимает вызываемый объект (функцию),
который получает на вход список выборок и должен вернуть батч.
Если параметр не задан (None), используется стандартная функция
default_collate.
Описание работы
Стандартная функция collate_fn работает корректно, если каждый
элемент датасета имеет одинаковую структуру: тензоры одинакового размера,
числовые значения или простые кортежи из тензоров. В сложных случаях
(например, предложения разной длины для задачи NLP) требуется
написать свою логику, чтобы правильно сгруппировать данные.
Пример с кастомной функцией
Рассмотрим датасет, который возвращает кортежи (тензор, метка, длина). Напишем функцию, которая объединяет их в батч:
import torch
from torch.utils.data import DataLoader, Dataset
class CustomDataset(Dataset):
def __init__(self):
self.data = [
(torch.tensor([1, 2, 3]), 0),
(torch.tensor([4, 5]), 1),
(torch.tensor([6, 7, 8, 9]), 0),
]
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
def my_collate(batch):
# batch - список из кортежей (tensor, label)
tensors = [item[0] for item in batch]
labels = [item[1] for item in batch]
# Преобразуем метки в тензор
labels = torch.tensor(labels)
# Паддинг тензоров до максимальной длины
max_len = max(t.size(0) for t in tensors)
padded_tensors = []
for t in tensors:
pad_len = max_len - t.size(0)
if pad_len > 0:
t = torch.cat([t, torch.zeros(pad_len, dtype=t.dtype)])
padded_tensors.append(t)
# Стек тензоров в батч
batch_tensors = torch.stack(padded_tensors)
return batch_tensors, labels
dataset = CustomDataset()
dataloader = DataLoader(dataset, batch_size=2, collate_fn=my_collate)
for batch_tensors, labels in dataloader:
print("Batch tensors:", batch_tensors)
print("Labels:", labels)
Результат выполнения кода:
Batch tensors: tensor([
[1., 2., 3., 0.],
[4., 5., 0., 0.]
])
Labels: tensor([0, 1])
Batch tensors: tensor([
[6., 7., 8., 9.]
])
Labels: tensor([0])
Пример использования default_collate
В этом примере мы не передаём collate_fn, и данные,
состоящие из тензоров одинаковой формы, автоматически
объединяются в батч:
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12]
])
labels = torch.tensor([0, 1, 0, 1])
dataset = TensorDataset(data, labels)
dataloader = DataLoader(dataset, batch_size=2)
for batch_data, batch_labels in dataloader:
print("Data:", batch_data)
print("Labels:", batch_labels)
Результат выполнения кода:
Data: tensor([
[1, 2, 3],
[4, 5, 6]
])
Labels: tensor([0, 1])
Data: tensor([
[7, 8, 9],
[10, 11, 12]
])
Labels: tensor([0, 1])
Смотрите также
-
класс
DataLoader,
который является основным инструментом для загрузки данных -
атрибут
dataset,
который хранит ссылку на используемый набор данных -
атрибут
batch_size,
который определяет размер мини-батча -
атрибут
num_workers,
который задаёт количество процессов для параллельной загрузки данных