# Define a sampler to balance the classes # training dataset lbls = [dataset[idx][1] for idx in train_set.indices] bc = np.bincount(lbls) p_nOK = bc.sum()/bc[0] p_OK = bc.sum()/bc[1] lst_train = [p_nOK if lbl==0 else p_OK for lbl in lbls] train_sampler = WeightedRandomSampler(weights=lst_train, num_samples=len(lbls))