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

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