Атрибут batch_size
Атрибут batch_size класса DataLoader
возвращает размер батча, переданный в конструктор
при создании загрузчика данных. Этот атрибут
является целочисленным значением и определяет,
сколько образцов данных будет содержаться в
одной итерации загрузчика. Атрибут доступен
только для чтения и не может быть изменен
после создания объекта.
Синтаксис
dataloader.batch_size
Пример
Давайте создадим загрузчик данных с размером батча 32 и получим значение атрибута:
import torch
from torch.utils.data import DataLoader, TensorDataset
train_data = torch.randn(100, 10)
train_labels = torch.randint(0, 2, (100,))
dataset = TensorDataset(train_data, train_labels)
dataloader = DataLoader(dataset, batch_size=32)
res = dataloader.batch_size
print(res)
Результат выполнения кода:
32
Пример
Атрибут работает с любыми значениями размера батча, включая 1 и None (если батч не используется):
import torch
from torch.utils.data import DataLoader, TensorDataset
data = torch.randn(50, 5)
dataset = TensorDataset(data)
dataloader_1 = DataLoader(dataset, batch_size=1)
print(dataloader_1.batch_size)
dataloader_none = DataLoader(dataset, batch_size=None)
print(dataloader_none.batch_size)
Результат выполнения кода:
1
None
Пример
Атрибут часто используется в циклах обучения для динамического определения размера батча:
import torch
from torch.utils.data import DataLoader, TensorDataset
train_data = torch.randn(200, 20)
train_labels = torch.randint(0, 10, (200,))
dataset = TensorDataset(train_data, train_labels)
dataloader = DataLoader(dataset, batch_size=64, shuffle=True)
for epoch in range(2):
print(f"Epoch {epoch + 1}")
for batch_idx, (data, labels) in enumerate(dataloader):
current_batch = len(data)
print(f"Batch {batch_idx}: size {current_batch} of {dataloader.batch_size}")
if batch_idx >= 2:
break
if epoch >= 0:
break
Результат выполнения кода:
Epoch 1
Batch 0: size 64 of 64
Batch 1: size 64 of 64
Batch 2: size 64 of 64
Смотрите также
-
класс
DataLoader,
который используется для загрузки данных в PyTorch -
атрибут
dataset,
который возвращает набор данных, используемый загрузчиком -
атрибут
sampler,
который определяет стратегию выборки образцов из датасета -
атрибут
drop_last,
который определяет, отбрасывать ли последний неполный батч