Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
vision_transformer.py482 linesDownload Raw Back to models
1import math2from collections import OrderedDict3from functools import partial4from typing import Any, Callable, List, NamedTuple, Optional5 6import torch7import torch.nn as nn8 9# from .._internally_replaced_utils import load_state_dict_from_url10from .vision_transformer_misc import ConvNormActivation11from .vision_transformer_utils import _log_api_usage_once12 13try:14    from torch.hub import load_state_dict_from_url15except ImportError:16    from torch.utils.model_zoo import load_url as load_state_dict_from_url17 18# __all__ = [19#     "VisionTransformer",20#     "vit_b_16",21#     "vit_b_32",22#     "vit_l_16",23#     "vit_l_32",24# ]25 26model_urls = {27    "vit_b_16": "https://download.pytorch.org/models/vit_b_16-c867db91.pth",28    "vit_b_32": "https://download.pytorch.org/models/vit_b_32-d86f8d99.pth",29    "vit_l_16": "https://download.pytorch.org/models/vit_l_16-852ce7e3.pth",30    "vit_l_32": "https://download.pytorch.org/models/vit_l_32-c7638314.pth",31}32 33 34class ConvStemConfig(NamedTuple):35    out_channels: int36    kernel_size: int37    stride: int38    norm_layer: Callable[..., nn.Module] = nn.BatchNorm2d39    activation_layer: Callable[..., nn.Module] = nn.ReLU40 41 42class MLPBlock(nn.Sequential):43    """Transformer MLP block."""44 45    def __init__(self, in_dim: int, mlp_dim: int, dropout: float):46        super().__init__()47        self.linear_1 = nn.Linear(in_dim, mlp_dim)48        self.act = nn.GELU()49        self.dropout_1 = nn.Dropout(dropout)50        self.linear_2 = nn.Linear(mlp_dim, in_dim)51        self.dropout_2 = nn.Dropout(dropout)52 53        nn.init.xavier_uniform_(self.linear_1.weight)54        nn.init.xavier_uniform_(self.linear_2.weight)55        nn.init.normal_(self.linear_1.bias, std=1e-6)56        nn.init.normal_(self.linear_2.bias, std=1e-6)57 58 59class EncoderBlock(nn.Module):60    """Transformer encoder block."""61 62    def __init__(63        self,64        num_heads: int,65        hidden_dim: int,66        mlp_dim: int,67        dropout: float,68        attention_dropout: float,69        norm_layer: Callable[..., torch.nn.Module] = partial(nn.LayerNorm, eps=1e-6),70    ):71        super().__init__()72        self.num_heads = num_heads73 74        # Attention block75        self.ln_1 = norm_layer(hidden_dim)76        self.self_attention = nn.MultiheadAttention(hidden_dim, num_heads, dropout=attention_dropout, batch_first=True)77        self.dropout = nn.Dropout(dropout)78 79        # MLP block80        self.ln_2 = norm_layer(hidden_dim)81        self.mlp = MLPBlock(hidden_dim, mlp_dim, dropout)82 83    def forward(self, input: torch.Tensor):84        torch._assert(input.dim() == 3, f"Expected (seq_length, batch_size, hidden_dim) got {input.shape}")85        x = self.ln_1(input)86        x, _ = self.self_attention(query=x, key=x, value=x, need_weights=False)87        x = self.dropout(x)88        x = x + input89 90        y = self.ln_2(x)91        y = self.mlp(y)92        return x + y93 94 95class Encoder(nn.Module):96    """Transformer Model Encoder for sequence to sequence translation."""97 98    def __init__(99        self,100        seq_length: int,101        num_layers: int,102        num_heads: int,103        hidden_dim: int,104        mlp_dim: int,105        dropout: float,106        attention_dropout: float,107        norm_layer: Callable[..., torch.nn.Module] = partial(nn.LayerNorm, eps=1e-6),108    ):109        super().__init__()110        # Note that batch_size is on the first dim because111        # we have batch_first=True in nn.MultiAttention() by default112        self.pos_embedding = nn.Parameter(torch.empty(1, seq_length, hidden_dim).normal_(std=0.02))  # from BERT113        self.dropout = nn.Dropout(dropout)114        layers: OrderedDict[str, nn.Module] = OrderedDict()115        for i in range(num_layers):116            layers[f"encoder_layer_{i}"] = EncoderBlock(117                num_heads,118                hidden_dim,119                mlp_dim,120                dropout,121                attention_dropout,122                norm_layer,123            )124        self.layers = nn.Sequential(layers)125        self.ln = norm_layer(hidden_dim)126 127    def forward(self, input: torch.Tensor):128        torch._assert(input.dim() == 3, f"Expected (batch_size, seq_length, hidden_dim) got {input.shape}")129        input = input + self.pos_embedding130        return self.ln(self.layers(self.dropout(input)))131 132 133class VisionTransformer(nn.Module):134    """Vision Transformer as per https://arxiv.org/abs/2010.11929."""135 136    def __init__(137        self,138        image_size: int,139        patch_size: int,140        num_layers: int,141        num_heads: int,142        hidden_dim: int,143        mlp_dim: int,144        dropout: float = 0.0,145        attention_dropout: float = 0.0,146        num_classes: int = 1000,147        representation_size: Optional[int] = None,148        norm_layer: Callable[..., torch.nn.Module] = partial(nn.LayerNorm, eps=1e-6),149        conv_stem_configs: Optional[List[ConvStemConfig]] = None,150    ):151        super().__init__()152        _log_api_usage_once(self)153        torch._assert(image_size % patch_size == 0, "Input shape indivisible by patch size!")154        self.image_size = image_size155        self.patch_size = patch_size156        self.hidden_dim = hidden_dim157        self.mlp_dim = mlp_dim158        self.attention_dropout = attention_dropout159        self.dropout = dropout160        self.num_classes = num_classes161        self.representation_size = representation_size162        self.norm_layer = norm_layer163 164        if conv_stem_configs is not None:165            # As per https://arxiv.org/abs/2106.14881166            seq_proj = nn.Sequential()167            prev_channels = 3168            for i, conv_stem_layer_config in enumerate(conv_stem_configs):169                seq_proj.add_module(170                    f"conv_bn_relu_{i}",171                    ConvNormActivation(172                        in_channels=prev_channels,173                        out_channels=conv_stem_layer_config.out_channels,174                        kernel_size=conv_stem_layer_config.kernel_size,175                        stride=conv_stem_layer_config.stride,176                        norm_layer=conv_stem_layer_config.norm_layer,177                        activation_layer=conv_stem_layer_config.activation_layer,178                    ),179                )180                prev_channels = conv_stem_layer_config.out_channels181            seq_proj.add_module(182                "conv_last", nn.Conv2d(in_channels=prev_channels, out_channels=hidden_dim, kernel_size=1)183            )184            self.conv_proj: nn.Module = seq_proj185        else:186            self.conv_proj = nn.Conv2d(187                in_channels=3, out_channels=hidden_dim, kernel_size=patch_size, stride=patch_size188            )189 190        seq_length = (image_size // patch_size) ** 2191 192        # Add a class token193        self.class_token = nn.Parameter(torch.zeros(1, 1, hidden_dim))194        seq_length += 1195 196        self.encoder = Encoder(197            seq_length,198            num_layers,199            num_heads,200            hidden_dim,201            mlp_dim,202            dropout,203            attention_dropout,204            norm_layer,205        )206        self.seq_length = seq_length207 208        heads_layers: OrderedDict[str, nn.Module] = OrderedDict()209        if representation_size is None:210            heads_layers["head"] = nn.Linear(hidden_dim, num_classes)211        else:212            heads_layers["pre_logits"] = nn.Linear(hidden_dim, representation_size)213            heads_layers["act"] = nn.Tanh()214            heads_layers["head"] = nn.Linear(representation_size, num_classes)215 216        self.heads = nn.Sequential(heads_layers)217 218        if isinstance(self.conv_proj, nn.Conv2d):219            # Init the patchify stem220            fan_in = self.conv_proj.in_channels * self.conv_proj.kernel_size[0] * self.conv_proj.kernel_size[1]221            nn.init.trunc_normal_(self.conv_proj.weight, std=math.sqrt(1 / fan_in))222            if self.conv_proj.bias is not None:223                nn.init.zeros_(self.conv_proj.bias)224        elif self.conv_proj.conv_last is not None and isinstance(self.conv_proj.conv_last, nn.Conv2d):225            # Init the last 1x1 conv of the conv stem226            nn.init.normal_(227                self.conv_proj.conv_last.weight, mean=0.0, std=math.sqrt(2.0 / self.conv_proj.conv_last.out_channels)228            )229            if self.conv_proj.conv_last.bias is not None:230                nn.init.zeros_(self.conv_proj.conv_last.bias)231 232        if hasattr(self.heads, "pre_logits") and isinstance(self.heads.pre_logits, nn.Linear):233            fan_in = self.heads.pre_logits.in_features234            nn.init.trunc_normal_(self.heads.pre_logits.weight, std=math.sqrt(1 / fan_in))235            nn.init.zeros_(self.heads.pre_logits.bias)236 237        if isinstance(self.heads.head, nn.Linear):238            nn.init.zeros_(self.heads.head.weight)239            nn.init.zeros_(self.heads.head.bias)240 241    def _process_input(self, x: torch.Tensor) -> torch.Tensor:242        n, c, h, w = x.shape243        p = self.patch_size244        torch._assert(h == self.image_size, "Wrong image height!")245        torch._assert(w == self.image_size, "Wrong image width!")246        n_h = h // p247        n_w = w // p248 249        # (n, c, h, w) -> (n, hidden_dim, n_h, n_w)250        x = self.conv_proj(x)251        # (n, hidden_dim, n_h, n_w) -> (n, hidden_dim, (n_h * n_w))252        x = x.reshape(n, self.hidden_dim, n_h * n_w)253 254        # (n, hidden_dim, (n_h * n_w)) -> (n, (n_h * n_w), hidden_dim)255        # The self attention layer expects inputs in the format (N, S, E)256        # where S is the source sequence length, N is the batch size, E is the257        # embedding dimension258        x = x.permute(0, 2, 1)259 260        return x261 262    def forward(self, x: torch.Tensor):263        out = {}264 265        # Reshape and permute the input tensor266        x = self._process_input(x)267        n = x.shape[0]268 269        # Expand the class token to the full batch270        batch_class_token = self.class_token.expand(n, -1, -1)271        x = torch.cat([batch_class_token, x], dim=1)272 273        274        x = self.encoder(x)275        img_feature = x[:,1:]276        H = W = int(self.image_size / self.patch_size)277        out['f4'] = img_feature.view(n, H, W, self.hidden_dim).permute(0,3,1,2)278 279        # Classifier "token" as used by standard language architectures280        x = x[:, 0]281        out['penultimate'] = x 282 283        x = self.heads(x) # I checked that for all pretrained ViT, this is just a fc 284        out['logits'] = x 285 286        return out287 288 289def _vision_transformer(290    arch: str,291    patch_size: int,292    num_layers: int,293    num_heads: int,294    hidden_dim: int,295    mlp_dim: int,296    pretrained: bool,297    progress: bool,298    **kwargs: Any,299) -> VisionTransformer:300    image_size = kwargs.pop("image_size", 224)301 302    model = VisionTransformer(303        image_size=image_size,304        patch_size=patch_size,305        num_layers=num_layers,306        num_heads=num_heads,307        hidden_dim=hidden_dim,308        mlp_dim=mlp_dim,309        **kwargs,310    )311 312    if pretrained:313        if arch not in model_urls:314            raise ValueError(f"No checkpoint is available for model type '{arch}'!")315        state_dict = load_state_dict_from_url(model_urls[arch], progress=progress)316        model.load_state_dict(state_dict)317 318    return model319 320 321def vit_b_16(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> VisionTransformer:322    """323    Constructs a vit_b_16 architecture from324    `"An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale" <https://arxiv.org/abs/2010.11929>`_.325 326    Args:327        pretrained (bool): If True, returns a model pre-trained on ImageNet328        progress (bool): If True, displays a progress bar of the download to stderr329    """330    return _vision_transformer(331        arch="vit_b_16",332        patch_size=16,333        num_layers=12,334        num_heads=12,335        hidden_dim=768,336        mlp_dim=3072,337        pretrained=pretrained,338        progress=progress,339        **kwargs,340    )341 342 343def vit_b_32(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> VisionTransformer:344    """345    Constructs a vit_b_32 architecture from346    `"An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale" <https://arxiv.org/abs/2010.11929>`_.347 348    Args:349        pretrained (bool): If True, returns a model pre-trained on ImageNet350        progress (bool): If True, displays a progress bar of the download to stderr351    """352    return _vision_transformer(353        arch="vit_b_32",354        patch_size=32,355        num_layers=12,356        num_heads=12,357        hidden_dim=768,358        mlp_dim=3072,359        pretrained=pretrained,360        progress=progress,361        **kwargs,362    )363 364 365def vit_l_16(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> VisionTransformer:366    """367    Constructs a vit_l_16 architecture from368    `"An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale" <https://arxiv.org/abs/2010.11929>`_.369 370    Args:371        pretrained (bool): If True, returns a model pre-trained on ImageNet372        progress (bool): If True, displays a progress bar of the download to stderr373    """374    return _vision_transformer(375        arch="vit_l_16",376        patch_size=16,377        num_layers=24,378        num_heads=16,379        hidden_dim=1024,380        mlp_dim=4096,381        pretrained=pretrained,382        progress=progress,383        **kwargs,384    )385 386 387def vit_l_32(pretrained: bool = False, progress: bool = True, **kwargs: Any) -> VisionTransformer:388    """389    Constructs a vit_l_32 architecture from390    `"An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale" <https://arxiv.org/abs/2010.11929>`_.391 392    Args:393        pretrained (bool): If True, returns a model pre-trained on ImageNet394        progress (bool): If True, displays a progress bar of the download to stderr395    """396    return _vision_transformer(397        arch="vit_l_32",398        patch_size=32,399        num_layers=24,400        num_heads=16,401        hidden_dim=1024,402        mlp_dim=4096,403        pretrained=pretrained,404        progress=progress,405        **kwargs,406    )407 408 409def interpolate_embeddings(410    image_size: int,411    patch_size: int,412    model_state: "OrderedDict[str, torch.Tensor]",413    interpolation_mode: str = "bicubic",414    reset_heads: bool = False,415) -> "OrderedDict[str, torch.Tensor]":416    """This function helps interpolating positional embeddings during checkpoint loading,417    especially when you want to apply a pre-trained model on images with different resolution.418 419    Args:420        image_size (int): Image size of the new model.421        patch_size (int): Patch size of the new model.422        model_state (OrderedDict[str, torch.Tensor]): State dict of the pre-trained model.423        interpolation_mode (str): The algorithm used for upsampling. Default: bicubic.424        reset_heads (bool): If true, not copying the state of heads. Default: False.425 426    Returns:427        OrderedDict[str, torch.Tensor]: A state dict which can be loaded into the new model.428    """429    # Shape of pos_embedding is (1, seq_length, hidden_dim)430    pos_embedding = model_state["encoder.pos_embedding"]431    n, seq_length, hidden_dim = pos_embedding.shape432    if n != 1:433        raise ValueError(f"Unexpected position embedding shape: {pos_embedding.shape}")434 435    new_seq_length = (image_size // patch_size) ** 2 + 1436 437    # Need to interpolate the weights for the position embedding.438    # We do this by reshaping the positions embeddings to a 2d grid, performing439    # an interpolation in the (h, w) space and then reshaping back to a 1d grid.440    if new_seq_length != seq_length:441        # The class token embedding shouldn't be interpolated so we split it up.442        seq_length -= 1443        new_seq_length -= 1444        pos_embedding_token = pos_embedding[:, :1, :]445        pos_embedding_img = pos_embedding[:, 1:, :]446 447        # (1, seq_length, hidden_dim) -> (1, hidden_dim, seq_length)448        pos_embedding_img = pos_embedding_img.permute(0, 2, 1)449        seq_length_1d = int(math.sqrt(seq_length))450        torch._assert(seq_length_1d * seq_length_1d == seq_length, "seq_length is not a perfect square!")451 452        # (1, hidden_dim, seq_length) -> (1, hidden_dim, seq_l_1d, seq_l_1d)453        pos_embedding_img = pos_embedding_img.reshape(1, hidden_dim, seq_length_1d, seq_length_1d)454        new_seq_length_1d = image_size // patch_size455 456        # Perform interpolation.457        # (1, hidden_dim, seq_l_1d, seq_l_1d) -> (1, hidden_dim, new_seq_l_1d, new_seq_l_1d)458        new_pos_embedding_img = nn.functional.interpolate(459            pos_embedding_img,460            size=new_seq_length_1d,461            mode=interpolation_mode,462            align_corners=True,463        )464 465        # (1, hidden_dim, new_seq_l_1d, new_seq_l_1d) -> (1, hidden_dim, new_seq_length)466        new_pos_embedding_img = new_pos_embedding_img.reshape(1, hidden_dim, new_seq_length)467 468        # (1, hidden_dim, new_seq_length) -> (1, new_seq_length, hidden_dim)469        new_pos_embedding_img = new_pos_embedding_img.permute(0, 2, 1)470        new_pos_embedding = torch.cat([pos_embedding_token, new_pos_embedding_img], dim=1)471 472        model_state["encoder.pos_embedding"] = new_pos_embedding473 474        if reset_heads:475            model_state_copy: "OrderedDict[str, torch.Tensor]" = OrderedDict()476            for k, v in model_state.items():477                if not k.startswith("heads"):478                    model_state_copy[k] = v479            model_state = model_state_copy480 481    return model_state482