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

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