Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
vision_transformer_misc.py164 linesDownload Raw Back to models
1from typing import Callable, List, Optional2 3import torch4from torch import Tensor5 6from .vision_transformer_utils import _log_api_usage_once7 8 9interpolate = torch.nn.functional.interpolate10 11 12# This is not in nn13class FrozenBatchNorm2d(torch.nn.Module):14    """15    BatchNorm2d where the batch statistics and the affine parameters are fixed16 17    Args:18        num_features (int): Number of features ``C`` from an expected input of size ``(N, C, H, W)``19        eps (float): a value added to the denominator for numerical stability. Default: 1e-520    """21 22    def __init__(23        self,24        num_features: int,25        eps: float = 1e-5,26    ):27        super().__init__()28        _log_api_usage_once(self)29        self.eps = eps30        self.register_buffer("weight", torch.ones(num_features))31        self.register_buffer("bias", torch.zeros(num_features))32        self.register_buffer("running_mean", torch.zeros(num_features))33        self.register_buffer("running_var", torch.ones(num_features))34 35    def _load_from_state_dict(36        self,37        state_dict: dict,38        prefix: str,39        local_metadata: dict,40        strict: bool,41        missing_keys: List[str],42        unexpected_keys: List[str],43        error_msgs: List[str],44    ):45        num_batches_tracked_key = prefix + "num_batches_tracked"46        if num_batches_tracked_key in state_dict:47            del state_dict[num_batches_tracked_key]48 49        super()._load_from_state_dict(50            state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs51        )52 53    def forward(self, x: Tensor) -> Tensor:54        # move reshapes to the beginning55        # to make it fuser-friendly56        w = self.weight.reshape(1, -1, 1, 1)57        b = self.bias.reshape(1, -1, 1, 1)58        rv = self.running_var.reshape(1, -1, 1, 1)59        rm = self.running_mean.reshape(1, -1, 1, 1)60        scale = w * (rv + self.eps).rsqrt()61        bias = b - rm * scale62        return x * scale + bias63 64    def __repr__(self) -> str:65        return f"{self.__class__.__name__}({self.weight.shape[0]}, eps={self.eps})"66 67 68class ConvNormActivation(torch.nn.Sequential):69    """70    Configurable block used for Convolution-Normalzation-Activation blocks.71 72    Args:73        in_channels (int): Number of channels in the input image74        out_channels (int): Number of channels produced by the Convolution-Normalzation-Activation block75        kernel_size: (int, optional): Size of the convolving kernel. Default: 376        stride (int, optional): Stride of the convolution. Default: 177        padding (int, tuple or str, optional): Padding added to all four sides of the input. Default: None, in wich case it will calculated as ``padding = (kernel_size - 1) // 2 * dilation``78        groups (int, optional): Number of blocked connections from input channels to output channels. Default: 179        norm_layer (Callable[..., torch.nn.Module], optional): Norm layer that will be stacked on top of the convolutiuon layer. If ``None`` this layer wont be used. Default: ``torch.nn.BatchNorm2d``80        activation_layer (Callable[..., torch.nn.Module], optinal): Activation function which will be stacked on top of the normalization layer (if not None), otherwise on top of the conv layer. If ``None`` this layer wont be used. Default: ``torch.nn.ReLU``81        dilation (int): Spacing between kernel elements. Default: 182        inplace (bool): Parameter for the activation layer, which can optionally do the operation in-place. Default ``True``83        bias (bool, optional): Whether to use bias in the convolution layer. By default, biases are included if ``norm_layer is None``.84 85    """86 87    def __init__(88        self,89        in_channels: int,90        out_channels: int,91        kernel_size: int = 3,92        stride: int = 1,93        padding: Optional[int] = None,94        groups: int = 1,95        norm_layer: Optional[Callable[..., torch.nn.Module]] = torch.nn.BatchNorm2d,96        activation_layer: Optional[Callable[..., torch.nn.Module]] = torch.nn.ReLU,97        dilation: int = 1,98        inplace: Optional[bool] = True,99        bias: Optional[bool] = None,100    ) -> None:101        if padding is None:102            padding = (kernel_size - 1) // 2 * dilation103        if bias is None:104            bias = norm_layer is None105        layers = [106            torch.nn.Conv2d(107                in_channels,108                out_channels,109                kernel_size,110                stride,111                padding,112                dilation=dilation,113                groups=groups,114                bias=bias,115            )116        ]117        if norm_layer is not None:118            layers.append(norm_layer(out_channels))119        if activation_layer is not None:120            params = {} if inplace is None else {"inplace": inplace}121            layers.append(activation_layer(**params))122        super().__init__(*layers)123        _log_api_usage_once(self)124        self.out_channels = out_channels125 126 127class SqueezeExcitation(torch.nn.Module):128    """129    This block implements the Squeeze-and-Excitation block from https://arxiv.org/abs/1709.01507 (see Fig. 1).130    Parameters ``activation``, and ``scale_activation`` correspond to ``delta`` and ``sigma`` in in eq. 3.131 132    Args:133        input_channels (int): Number of channels in the input image134        squeeze_channels (int): Number of squeeze channels135        activation (Callable[..., torch.nn.Module], optional): ``delta`` activation. Default: ``torch.nn.ReLU``136        scale_activation (Callable[..., torch.nn.Module]): ``sigma`` activation. Default: ``torch.nn.Sigmoid``137    """138 139    def __init__(140        self,141        input_channels: int,142        squeeze_channels: int,143        activation: Callable[..., torch.nn.Module] = torch.nn.ReLU,144        scale_activation: Callable[..., torch.nn.Module] = torch.nn.Sigmoid,145    ) -> None:146        super().__init__()147        _log_api_usage_once(self)148        self.avgpool = torch.nn.AdaptiveAvgPool2d(1)149        self.fc1 = torch.nn.Conv2d(input_channels, squeeze_channels, 1)150        self.fc2 = torch.nn.Conv2d(squeeze_channels, input_channels, 1)151        self.activation = activation()152        self.scale_activation = scale_activation()153 154    def _scale(self, input: Tensor) -> Tensor:155        scale = self.avgpool(input)156        scale = self.fc1(scale)157        scale = self.activation(scale)158        scale = self.fc2(scale)159        return self.scale_activation(scale)160 161    def forward(self, input: Tensor) -> Tensor:162        scale = self._scale(input)163        return scale * input164