РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
578 of 769 menu

Атрибут 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,
    который задаёт количество процессов для параллельной загрузки данных
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить