import torch INPUT_SAMPLES = 10000 SUB_SAMPLE = INPUT_SAMPLES // 10 FEATURE_DIM = 16 def sample_data(input_array, labels): device = labels.device positive = torch.nonzero(labels == 1, as_tuple=True)[0] negative = torch.nonzero(labels == 0, as_tuple=True)[0] num_pos = min(positive.numel(), SUB_SAMPLE//2) num_neg = min(negative.numel(), SUB_SAMPLE//2) if num_neg < SUB_SAMPLE//2: num_pos = SUB_SAMPLE - num_neg elif num_pos < SUB_SAMPLE//2: num_neg = SUB_SAMPLE - num_pos # randomly select positive and negative examples perm1 = torch.randperm(positive.numel(), device=device)[:num_pos] perm2 = torch.randperm(negative.numel(), device=device)[:num_neg] pos_idxs = positive[perm1] neg_idxs = negative[perm2] sampled_idxs = torch.cat([pos_idxs, neg_idxs], dim=0) rand_perm = torch.randperm(SUB_SAMPLE, device=labels.device) sampled_idxs = sampled_idxs[rand_perm] return input_array[sampled_idxs], labels[sampled_idxs]