Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2from torch import nn3from torchvision.ops import roi_align4 5 6# NOTE: torchvision's RoIAlign has a different default aligned=False7class ROIAlign(nn.Module):8 def __init__(self, output_size, spatial_scale, sampling_ratio, aligned=True):9 """10 Args:11 output_size (tuple): h, w12 spatial_scale (float): scale the input boxes by this number13 sampling_ratio (int): number of inputs samples to take for each output14 sample. 0 to take samples densely.15 aligned (bool): if False, use the legacy implementation in16 Detectron. If True, align the results more perfectly.17 18 Note:19 The meaning of aligned=True:20 21 Given a continuous coordinate c, its two neighboring pixel indices (in our22 pixel model) are computed by floor(c - 0.5) and ceil(c - 0.5). For example,23 c=1.3 has pixel neighbors with discrete indices [0] and [1] (which are sampled24 from the underlying signal at continuous coordinates 0.5 and 1.5). But the original25 roi_align (aligned=False) does not subtract the 0.5 when computing neighboring26 pixel indices and therefore it uses pixels with a slightly incorrect alignment27 (relative to our pixel model) when performing bilinear interpolation.28 29 With `aligned=True`,30 we first appropriately scale the ROI and then shift it by -0.531 prior to calling roi_align. This produces the correct neighbors; see32 detectron2/tests/test_roi_align.py for verification.33 34 The difference does not make a difference to the model's performance if35 ROIAlign is used together with conv layers.36 """37 super().__init__()38 self.output_size = output_size39 self.spatial_scale = spatial_scale40 self.sampling_ratio = sampling_ratio41 self.aligned = aligned42 43 from torchvision import __version__44 45 version = tuple(int(x) for x in __version__.split(".")[:2])46 # https://github.com/pytorch/vision/pull/243847 assert version >= (0, 7), "Require torchvision >= 0.7"48 49 def forward(self, input, rois):50 """51 Args:52 input: NCHW images53 rois: Bx5 boxes. First column is the index into N. The other 4 columns are xyxy.54 """55 assert rois.dim() == 2 and rois.size(1) == 556 if input.is_quantized:57 input = input.dequantize()58 return roi_align(59 input,60 rois.to(dtype=input.dtype),61 self.output_size,62 self.spatial_scale,63 self.sampling_ratio,64 self.aligned,65 )66 67 def __repr__(self):68 tmpstr = self.__class__.__name__ + "("69 tmpstr += "output_size=" + str(self.output_size)70 tmpstr += ", spatial_scale=" + str(self.spatial_scale)71 tmpstr += ", sampling_ratio=" + str(self.sampling_ratio)72 tmpstr += ", aligned=" + str(self.aligned)73 tmpstr += ")"74 return tmpstr75 