import torch from torch.utils.data import DataLoader, TensorDataset, WeightedRandomSampler labels = torch.tensor([0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1]) features = torch.arange(len(labels), dtype=torch.float32).unsqueeze(1) train_dataset = TensorDataset(features, labels)