Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2"""3Wrappers around on some nn functions, mainly to support empty tensors.4 5Ideally, add support directly in PyTorch to empty tensors in those functions.6 7These can be removed once https://github.com/pytorch/pytorch/issues/120138is implemented9"""10 11import warnings12from typing import List, Optional13import torch14from torch.nn import functional as F15 16from detectron2.utils.env import TORCH_VERSION17 18 19def shapes_to_tensor(x: List[int], device: Optional[torch.device] = None) -> torch.Tensor:20 """21 Turn a list of integer scalars or integer Tensor scalars into a vector,22 in a way that's both traceable and scriptable.23 24 In tracing, `x` should be a list of scalar Tensor, so the output can trace to the inputs.25 In scripting or eager, `x` should be a list of int.26 """27 if torch.jit.is_scripting():28 return torch.as_tensor(x, device=device)29 if torch.jit.is_tracing():30 assert all(31 [isinstance(t, torch.Tensor) for t in x]32 ), "Shape should be tensor during tracing!"33 # as_tensor should not be used in tracing because it records a constant34 ret = torch.stack(x)35 if ret.device != device: # avoid recording a hard-coded device if not necessary36 ret = ret.to(device=device)37 return ret38 return torch.as_tensor(x, device=device)39 40 41def check_if_dynamo_compiling():42 if TORCH_VERSION >= (1, 14):43 from torch._dynamo import is_compiling44 45 return is_compiling()46 else:47 return False48 49 50def cat(tensors: List[torch.Tensor], dim: int = 0):51 """52 Efficient version of torch.cat that avoids a copy if there is only a single element in a list53 """54 assert isinstance(tensors, (list, tuple))55 if len(tensors) == 1:56 return tensors[0]57 return torch.cat(tensors, dim)58 59 60def empty_input_loss_func_wrapper(loss_func):61 def wrapped_loss_func(input, target, *, reduction="mean", **kwargs):62 """63 Same as `loss_func`, but returns 0 (instead of nan) for empty inputs.64 """65 if target.numel() == 0 and reduction == "mean":66 return input.sum() * 0.0 # connect the gradient67 return loss_func(input, target, reduction=reduction, **kwargs)68 69 return wrapped_loss_func70 71 72cross_entropy = empty_input_loss_func_wrapper(F.cross_entropy)73 74 75class _NewEmptyTensorOp(torch.autograd.Function):76 @staticmethod77 def forward(ctx, x, new_shape):78 ctx.shape = x.shape79 return x.new_empty(new_shape)80 81 @staticmethod82 def backward(ctx, grad):83 shape = ctx.shape84 return _NewEmptyTensorOp.apply(grad, shape), None85 86 87class Conv2d(torch.nn.Conv2d):88 """89 A wrapper around :class:`torch.nn.Conv2d` to support empty inputs and more features.90 """91 92 def __init__(self, *args, **kwargs):93 """94 Extra keyword arguments supported in addition to those in `torch.nn.Conv2d`:95 96 Args:97 norm (nn.Module, optional): a normalization layer98 activation (callable(Tensor) -> Tensor): a callable activation function99 100 It assumes that norm layer is used before activation.101 """102 norm = kwargs.pop("norm", None)103 activation = kwargs.pop("activation", None)104 super().__init__(*args, **kwargs)105 106 self.norm = norm107 self.activation = activation108 109 def forward(self, x):110 # torchscript does not support SyncBatchNorm yet111 # https://github.com/pytorch/pytorch/issues/40507112 # and we skip these codes in torchscript since:113 # 1. currently we only support torchscript in evaluation mode114 # 2. features needed by exporting module to torchscript are added in PyTorch 1.6 or115 # later version, `Conv2d` in these PyTorch versions has already supported empty inputs.116 if not torch.jit.is_scripting():117 # Dynamo doesn't support context managers yet118 is_dynamo_compiling = check_if_dynamo_compiling()119 if not is_dynamo_compiling:120 with warnings.catch_warnings(record=True):121 if x.numel() == 0 and self.training:122 # https://github.com/pytorch/pytorch/issues/12013123 assert not isinstance(124 self.norm, torch.nn.SyncBatchNorm125 ), "SyncBatchNorm does not support empty inputs!"126 127 x = F.conv2d(128 x, self.weight, self.bias, self.stride, self.padding, self.dilation, self.groups129 )130 if self.norm is not None:131 x = self.norm(x)132 if self.activation is not None:133 x = self.activation(x)134 return x135 136 137ConvTranspose2d = torch.nn.ConvTranspose2d138BatchNorm2d = torch.nn.BatchNorm2d139interpolate = F.interpolate140Linear = torch.nn.Linear141 142 143def nonzero_tuple(x):144 """145 A 'as_tuple=True' version of torch.nonzero to support torchscript.146 because of https://github.com/pytorch/pytorch/issues/38718147 """148 if torch.jit.is_scripting():149 if x.dim() == 0:150 return x.unsqueeze(0).nonzero().unbind(1)151 return x.nonzero().unbind(1)152 else:153 return x.nonzero(as_tuple=True)154 155 156@torch.jit.script_if_tracing157def move_device_like(src: torch.Tensor, dst: torch.Tensor) -> torch.Tensor:158 """159 Tracing friendly way to cast tensor to another tensor's device. Device will be treated160 as constant during tracing, scripting the casting process as whole can workaround this issue.161 """162 return src.to(dst.device)163 