Случайное разбиение набора в PyTorch
Удобно часть примеров оставить
для проверки, а остальные - для
обучения. Функция
random_split из
torch.utils.data делит
набор на несколько частей по
заданным длинам; сумма длин
должна совпадать с размером
исходного набора.
Чтобы разбиение повторялось,
передают объект Generator
с зафиксированным зерном.
Разделим пять элементов на
три и два, затем выведем длины
обеих частей:
import torch
from torch.utils.data import TensorDataset, random_split
features = torch.tensor([[1.0], [2.0], [3.0], [4.0], [5.0]])
targets = torch.tensor([0, 1, 0, 1, 0])
ds = TensorDataset(features, targets)
gen = torch.Generator().manual_seed(42)
train_part, test_part = random_split(ds, [3, 2], generator=gen)
print(len(train_part), len(test_part)) # выведет 3 2
Каждая часть ведёт себя как отдельный набор: к ней можно обращаться по индексу и подключать загрузчик с пакетами.
Соберите набор из 10
строк по одному числу от 0.0
до 9.0 и такого же числа
нулевых меток.
Разделите его на части длиной
7 и 3 с зерном
0. Выведите длины
обеих частей.
Для набора из шести пар
признак-метка разделите данные
на 4 и 2 элемента,
зафиксировав зерно 100.
Выведите, сколько элементов
попало во вторую часть.
Из восьми строк-признаков
и восьми меток получите набор
и разделите его поровну на две
части по 4 элемента.
С зерном 1 выведите
длину первой части.