Функция random_split
Функция random_split случайным образом разделяет переданный набор данных на несколько непересекающихся поднаборов заданной длины. Первым параметром функция принимает исходный набор данных (Dataset), вторым параметром - список или кортеж с размерами поднаборов. Опционально можно передать генератор случайных чисел для воспроизводимости результатов.
Синтаксис
torch.utils.data.random_split(dataset, lengths, generator)
Пример
Давайте создадим набор данных из 10 элементов и разделим его на две части:
import torch
from torch.utils.data import Dataset, random_split
class MyDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
t = [i for i in range(10)]
dataset = MyDataset(t)
train_size = 7
val_size = 3
train_dataset, val_dataset = random_split(
dataset, [train_size, val_size]
)
print("Train indices:", train_dataset.indices)
print("Val indices:", val_dataset.indices)
Результат выполнения кода:
"Train indices: [8, 4, 7, 9, 2, 0, 6]"
"Val indices: [1, 5, 3]"
Пример
Используем фиксированное зерно для воспроизводимого разделения:
import torch
from torch.utils.data import Dataset, random_split
class MyDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
torch.manual_seed(42)
dataset = MyDataset([10, 20, 30, 40, 50, 60, 70, 80])
train_size = 5
val_size = 3
train_dataset, val_dataset = random_split(
dataset, [train_size, val_size]
)
print("Train:", [dataset[i] for i in train_dataset.indices])
print("Val:", [dataset[i] for i in val_dataset.indices])
Результат выполнения кода:
"Train: [60, 80, 30, 40, 70]"
"Val: [10, 20, 50]"
Пример
Разделим набор данных на три части для обучения, валидации и тестирования:
import torch
from torch.utils.data import Dataset, random_split
class MyDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
torch.manual_seed(0)
dataset = MyDataset([1, 2, 3, 4, 5, 6, 7, 8, 9, 10])
train_size = 6
val_size = 2
test_size = 2
train_ds, val_ds, test_ds = random_split(
dataset, [train_size, val_size, test_size]
)
print("Train:", train_ds.indices)
print("Val:", val_ds.indices)
print("Test:", test_ds.indices)
Результат выполнения кода:
"Train: [6, 9, 7, 3, 2, 5]"
"Val: [8, 10]"
"Test: [1, 4]"
Смотрите также
-
класс
Subset,
который создаёт поднабор данных по заданным индексам -
класс
TensorDataset,
который оборачивает тензоры в набор данных -
класс
ConcatDataset,
который объединяет несколько наборов данных в один -
класс
SubsetRandomSampler,
который сэмплирует элементы из поднабора случайным образом