def opt_sample_data(input, labels): pos_mask = labels == 1 neg_mask = labels == 0 num_pos_idxs = torch.count_nonzero(pos_mask, dim=-1) num_neg_idxs = torch.count_nonzero(neg_mask, dim=-1) half_samples = labels.new_full((), SUB_SAMPLE // 2) num_pos = torch.minimum(num_pos_idxs, half_samples) num_neg = torch.minimum(num_neg_idxs, half_samples) num_pos = torch.where( num_neg < SUB_SAMPLE // 2, SUB_SAMPLE - num_neg, num_pos ) num_neg = SUB_SAMPLE - num_pos # create random ordering on pos and neg entries rand = torch.rand_like(labels, dtype=torch.float32) pos_rand = torch.where(pos_mask, rand, -1) neg_rand = torch.where(neg_mask, rand, -1) # select top pos entries and invalidate others # since CPU doesn't know num_pos, we assume maximum to avoid sync top_pos_rand, top_pos_idx = torch.topk(pos_rand, k=SUB_SAMPLE) arange = torch.arange(SUB_SAMPLE, device=labels.device) if num_pos.numel() > 1: # unsqueeze to support batched input arange = arange.unsqueeze(0) num_pos = num_pos.unsqueeze(-1) num_neg = num_neg.unsqueeze(-1) top_pos_rand = torch.where(arange >= num_pos, -1, top_pos_rand) # repeat for neg entries top_neg_rand, top_neg_idx = torch.topk(neg_rand, k=SUB_SAMPLE) top_neg_rand = torch.where(arange >= num_neg, -1, top_neg_rand) # combine and mix together positive and negative idxs cat_rand = torch.cat([top_pos_rand, top_neg_rand], dim=-1) cat_idx = torch.cat([top_pos_idx, top_neg_idx], dim=-1) topk_rand_idx = torch.topk(cat_rand, k=SUB_SAMPLE)[1] sampled_idxs = torch.gather(cat_idx, dim=-1, index=topk_rand_idx) sampled_input = torch.gather(input, dim=-2, index=sampled_idxs.unsqueeze(-1)) sampled_labels = torch.gather(labels, dim=-1, index=sampled_idxs) return sampled_input, sampled_labels