Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
sampling.py55 linesDownload Raw Back to modeling
1# Copyright (c) Facebook, Inc. and its affiliates.2import torch3 4from detectron2.layers import nonzero_tuple5 6__all__ = ["subsample_labels"]7 8 9def subsample_labels(10    labels: torch.Tensor, num_samples: int, positive_fraction: float, bg_label: int11):12    """13    Return `num_samples` (or fewer, if not enough found)14    random samples from `labels` which is a mixture of positives & negatives.15    It will try to return as many positives as possible without16    exceeding `positive_fraction * num_samples`, and then try to17    fill the remaining slots with negatives.18 19    Args:20        labels (Tensor): (N, ) label vector with values:21            * -1: ignore22            * bg_label: background ("negative") class23            * otherwise: one or more foreground ("positive") classes24        num_samples (int): The total number of labels with value >= 0 to return.25            Values that are not sampled will be filled with -1 (ignore).26        positive_fraction (float): The number of subsampled labels with values > 027            is `min(num_positives, int(positive_fraction * num_samples))`. The number28            of negatives sampled is `min(num_negatives, num_samples - num_positives_sampled)`.29            In order words, if there are not enough positives, the sample is filled with30            negatives. If there are also not enough negatives, then as many elements are31            sampled as is possible.32        bg_label (int): label index of background ("negative") class.33 34    Returns:35        pos_idx, neg_idx (Tensor):36            1D vector of indices. The total length of both is `num_samples` or fewer.37    """38    positive = nonzero_tuple((labels != -1) & (labels != bg_label))[0]39    negative = nonzero_tuple(labels == bg_label)[0]40 41    num_pos = int(num_samples * positive_fraction)42    # protect against not enough positive examples43    num_pos = min(positive.numel(), num_pos)44    num_neg = num_samples - num_pos45    # protect against not enough negative examples46    num_neg = min(negative.numel(), num_neg)47 48    # randomly select positive and negative examples49    perm1 = torch.randperm(positive.numel(), device=positive.device)[:num_pos]50    perm2 = torch.randperm(negative.numel(), device=negative.device)[:num_neg]51 52    pos_idx = positive[perm1]53    neg_idx = negative[perm2]54    return pos_idx, neg_idx55