Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
wrappers.py163 linesDownload Raw Back to layers
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