Класс BatchSampler
Класс BatchSampler предназначен для
формирования пакетов индексов из данных,
предоставляемых базовым сэмплером. Первым
параметром конструктор принимает объект
сэмплера (например, SequentialSampler
или RandomSampler). Вторым параметром
передаётся размер пакета. Третьим параметром
можно указать, нужно ли отбрасывать последний
неполный пакет.
Синтаксис
torch.utils.data.BatchSampler(
sampler,
batch_size,
drop_last=False
)
Пример
Давайте создадим пакетный сэмплер для
последовательной выборки индексов из
диапазона 10 элементов размером
пакета 3:
import torch
from torch.utils.data import BatchSampler
from torch.utils.data import SequentialSampler
sampler = SequentialSampler(range(10))
batch_sampler = BatchSampler(
sampler,
batch_size=3,
drop_last=False
)
for batch in batch_sampler:
print(batch)
Результат выполнения кода:
[0, 1, 2]
[3, 4, 5]
[6, 7, 8]
[9]
Пример
Давайте используем BatchSampler с
отбрасыванием последнего неполного пакета:
import torch
from torch.utils.data import BatchSampler
from torch.utils.data import SequentialSampler
sampler = SequentialSampler(range(10))
batch_sampler = BatchSampler(
sampler,
batch_size=3,
drop_last=True
)
for batch in batch_sampler:
print(batch)
Результат выполнения кода:
[0, 1, 2]
[3, 4, 5]
[6, 7, 8]
Пример
Давайте используем BatchSampler с
перемешиванием данных с помощью
RandomSampler:
import torch
from torch.utils.data import BatchSampler
from torch.utils.data import RandomSampler
torch.manual_seed(0)
sampler = RandomSampler(range(10))
batch_sampler = BatchSampler(
sampler,
batch_size=3,
drop_last=False
)
for batch in batch_sampler:
print(batch)
Результат выполнения кода:
[6, 9, 3]
[4, 0, 5]
[2, 7, 1]
[8]
Пример
Давайте используем BatchSampler вместе
с загрузчиком данных DataLoader для
загрузки мини-пакетов:
import torch
from torch.utils.data import DataLoader
from torch.utils.data import TensorDataset
from torch.utils.data import BatchSampler
from torch.utils.data import SequentialSampler
data = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12],
[13, 14, 15],
[16, 17, 18],
])
dataset = TensorDataset(data)
sampler = SequentialSampler(dataset)
batch_sampler = BatchSampler(
sampler,
batch_size=2,
drop_last=True
)
dataloader = DataLoader(
dataset,
batch_sampler=batch_sampler
)
for batch in dataloader:
print(batch)
Результат выполнения кода:
[tensor([
[1, 2, 3],
[4, 5, 6],
])]
[tensor([
[7, 8, 9],
[10, 11, 12],
])]
[tensor([
[13, 14, 15],
[16, 17, 18],
])]
Смотрите также
-
класс
Sampler,
который является базовым классом для всех сэмплеров -
класс
SequentialSampler,
который возвращает индексы в последовательном порядке -
класс
RandomSampler,
который возвращает индексы в случайном порядке -
класс
SubsetRandomSampler,
который возвращает случайные индексы из подмножества данных