Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
model.py487 linesDownload Raw Back to clip
1from collections import OrderedDict2from typing import Tuple, Union3 4import numpy as np5import torch6import torch.nn.functional as F7from torch import nn8 9 10class Bottleneck(nn.Module):11    expansion = 412 13    def __init__(self, inplanes, planes, stride=1):14        super().__init__()15 16        # all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 117        self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False)18        self.bn1 = nn.BatchNorm2d(planes)19        self.relu1 = nn.ReLU(inplace=True)20 21        self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False)22        self.bn2 = nn.BatchNorm2d(planes)23        self.relu2 = nn.ReLU(inplace=True)24 25        self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity()26 27        self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False)28        self.bn3 = nn.BatchNorm2d(planes * self.expansion)29        self.relu3 = nn.ReLU(inplace=True)30 31        self.downsample = None32        self.stride = stride33 34        if stride > 1 or inplanes != planes * Bottleneck.expansion:35            # downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 136            self.downsample = nn.Sequential(OrderedDict([37                ("-1", nn.AvgPool2d(stride)),38                ("0", nn.Conv2d(inplanes, planes * self.expansion, 1, stride=1, bias=False)),39                ("1", nn.BatchNorm2d(planes * self.expansion))40            ]))41 42    def forward(self, x: torch.Tensor):43        identity = x44 45        out = self.relu1(self.bn1(self.conv1(x)))46        out = self.relu2(self.bn2(self.conv2(out)))47        out = self.avgpool(out)48        out = self.bn3(self.conv3(out))49 50        if self.downsample is not None:51            identity = self.downsample(x)52 53        out += identity54        out = self.relu3(out)55        return out56 57 58class AttentionPool2d(nn.Module):59    def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):60        super().__init__()61        self.positional_embedding = nn.Parameter(torch.randn(spacial_dim ** 2 + 1, embed_dim) / embed_dim ** 0.5)62        self.k_proj = nn.Linear(embed_dim, embed_dim)63        self.q_proj = nn.Linear(embed_dim, embed_dim)64        self.v_proj = nn.Linear(embed_dim, embed_dim)65        self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)66        self.num_heads = num_heads67 68    def forward(self, x):69        x = x.flatten(start_dim=2).permute(2, 0, 1)  # NCHW -> (HW)NC70        x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0)  # (HW+1)NC71        x = x + self.positional_embedding[:, None, :].to(x.dtype)  # (HW+1)NC72        x, _ = F.multi_head_attention_forward(73            query=x[:1], key=x, value=x,74            embed_dim_to_check=x.shape[-1],75            num_heads=self.num_heads,76            q_proj_weight=self.q_proj.weight,77            k_proj_weight=self.k_proj.weight,78            v_proj_weight=self.v_proj.weight,79            in_proj_weight=None,80            in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),81            bias_k=None,82            bias_v=None,83            add_zero_attn=False,84            dropout_p=0,85            out_proj_weight=self.c_proj.weight,86            out_proj_bias=self.c_proj.bias,87            use_separate_proj_weight=True,88            training=self.training,89            need_weights=False90        )91        return x.squeeze(0)92 93 94class ModifiedResNet(nn.Module):95    """96    A ResNet class that is similar to torchvision's but contains the following changes:97    - There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool.98    - Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 199    - The final pooling layer is a QKV attention instead of an average pool100    """101 102    def __init__(self, layers, output_dim, heads, input_resolution=224, width=64):103        super().__init__()104        self.output_dim = output_dim105        self.input_resolution = input_resolution106 107        # the 3-layer stem108        self.conv1 = nn.Conv2d(3, width // 2, kernel_size=3, stride=2, padding=1, bias=False)109        self.bn1 = nn.BatchNorm2d(width // 2)110        self.relu1 = nn.ReLU(inplace=True)111        self.conv2 = nn.Conv2d(width // 2, width // 2, kernel_size=3, padding=1, bias=False)112        self.bn2 = nn.BatchNorm2d(width // 2)113        self.relu2 = nn.ReLU(inplace=True)114        self.conv3 = nn.Conv2d(width // 2, width, kernel_size=3, padding=1, bias=False)115        self.bn3 = nn.BatchNorm2d(width)116        self.relu3 = nn.ReLU(inplace=True)117        self.avgpool = nn.AvgPool2d(2)118 119        # residual layers120        self._inplanes = width  # this is a *mutable* variable used during construction121        self.layer1 = self._make_layer(width, layers[0])122        self.layer2 = self._make_layer(width * 2, layers[1], stride=2)123        self.layer3 = self._make_layer(width * 4, layers[2], stride=2)124        self.layer4 = self._make_layer(width * 8, layers[3], stride=2)125 126        embed_dim = width * 32  # the ResNet feature dimension127        self.attnpool = AttentionPool2d(input_resolution // 32, embed_dim, heads, output_dim)128 129    def _make_layer(self, planes, blocks, stride=1):130        layers = [Bottleneck(self._inplanes, planes, stride)]131 132        self._inplanes = planes * Bottleneck.expansion133        for _ in range(1, blocks):134            layers.append(Bottleneck(self._inplanes, planes))135 136        return nn.Sequential(*layers)137 138    def forward(self, x):139        def stem(x):140            x = self.relu1(self.bn1(self.conv1(x)))141            x = self.relu2(self.bn2(self.conv2(x)))142            x = self.relu3(self.bn3(self.conv3(x)))143            x = self.avgpool(x)144            return x145 146        x = x.type(self.conv1.weight.dtype)147        x = stem(x)148        x = self.layer1(x)149        x = self.layer2(x)150        x = self.layer3(x)151        x = self.layer4(x)152        x = self.attnpool(x)153 154        return x155 156 157class LayerNorm(nn.LayerNorm):158    """Subclass torch's LayerNorm to handle fp16."""159 160    def forward(self, x: torch.Tensor):161        orig_type = x.dtype162        ret = super().forward(x.type(torch.float32))163        return ret.type(orig_type)164 165 166class QuickGELU(nn.Module):167    def forward(self, x: torch.Tensor):168        return x * torch.sigmoid(1.702 * x)169 170 171class ResidualAttentionBlock(nn.Module):172    def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):173        super().__init__()174 175        self.attn = nn.MultiheadAttention(d_model, n_head)176        self.ln_1 = LayerNorm(d_model)177        self.mlp = nn.Sequential(OrderedDict([178            ("c_fc", nn.Linear(d_model, d_model * 4)),179            ("gelu", QuickGELU()),180            ("c_proj", nn.Linear(d_model * 4, d_model))181        ]))182        self.ln_2 = LayerNorm(d_model)183        self.attn_mask = attn_mask184 185    def attention(self, x: torch.Tensor):186        self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None187        return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]188 189    def forward(self, x: torch.Tensor):190        x = x + self.attention(self.ln_1(x))191        x = x + self.mlp(self.ln_2(x))192        return x193 194 195class Transformer(nn.Module):196    def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None):197        super().__init__()198        self.width = width199        self.layers = layers200        self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])201 202    def forward(self, x: torch.Tensor):203        out = {}204        for idx, layer in enumerate(self.resblocks.children()):205            x = layer(x)206            out['layer'+str(idx)] = x[0] # shape:LND. choose cls token feature   207        return out, x 208 209        # return self.resblocks(x)  # This is the original code 210 211 212class VisionTransformer(nn.Module):213    def __init__(self, input_resolution: int, patch_size: int, width: int, layers: int, heads: int, output_dim: int):214        super().__init__()215        self.input_resolution = input_resolution216        self.output_dim = output_dim217        self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)218 219        scale = width ** -0.5220        self.class_embedding = nn.Parameter(scale * torch.randn(width))221        self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))222        self.ln_pre = LayerNorm(width)223 224        self.transformer = Transformer(width, layers, heads)225 226        self.ln_post = LayerNorm(width)227        self.proj = nn.Parameter(scale * torch.randn(width, output_dim))228 229 230 231    def forward(self, x: torch.Tensor):232        """233        原代码这里的x是4个dimension,即batchsize*RGBchannels*224*224234        若只输入一张图片,因为没有batchsize维度,需要在最前面加一个维度,见下面第一行代码235        """236        x = x.reshape(-1,x.shape[-3],x.shape[-2],x.shape[-1])237        # print(x.shape)238        x = self.conv1(x)  # shape = [*, width, grid, grid]239        # print(x.shape)240        x = x.reshape(x.shape[0], x.shape[1], -1)  # shape = [*, width, grid ** 2]241        # print(x.shape)242        x = x.permute(0, 2, 1)  # shape = [*, grid ** 2, width]243        # print(self.class_embedding.to(x.dtype).shape)244        # print(torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device).shape)245        # print(x.shape)246        x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1)  # shape = [*, grid ** 2 + 1, width]247        x = x + self.positional_embedding.to(x.dtype)248        x = self.ln_pre(x)249 250        x = x.permute(1, 0, 2)  # NLD -> LND251        out, x = self.transformer(x)252        x = x.permute(1, 0, 2)  # LND -> NLD253 254        x = self.ln_post(x[:, 0, :])255 256 257        out['before_projection'] = x  258 259        if self.proj is not None:260            x = x @ self.proj261        out['after_projection'] = x 262 263        """264        将ViT-Large中第20,22,24层(或16,20,24)的[cls]feature做加权平均, 经过projection后输出265        """266        out['res_output'] = torch.zeros_like(out['before_projection'])267        # for layer_output in [[0.2, out['layer15']], [0.3, out['layer19']], [0.5, out['layer23']]]:268        for layer_output in [[0.2, out['layer19']], [0.3, out['layer21']], [0.5, out['layer23']]]:269        # for layer_output in [[0.5, out['layer15']], [0.5, out['layer21']]]:270            # layer_output[1] = layer_output[1].permute(1, 0, 2)  # LND -> NLD271            layer_output[1] = self.ln_post(layer_output[1])272            out['res_output'] += layer_output[0]*layer_output[1]273        out['res_output'] = out['res_output'] @ self.proj274 275        """276        将ViT每一层Encoder的[cls]feature都输出277        278        形式e.g.279        out['layer0'] = ...280        out['layer1'] = ...281        """282        # Return both intermediate features and final clip feature 283        return out284        285        # This only returns CLIP features 286        # return x 287 288 289class CLIP(nn.Module):290    def __init__(self,291                 embed_dim: int,292                 # vision293                 image_resolution: int,294                 vision_layers: Union[Tuple[int, int, int, int], int],295                 vision_width: int,296                 vision_patch_size: int,297                 # text298                 context_length: int,299                 vocab_size: int,300                 transformer_width: int,301                 transformer_heads: int,302                 transformer_layers: int303                 ):304        super().__init__()305 306        self.context_length = context_length307 308        if isinstance(vision_layers, (tuple, list)):309            vision_heads = vision_width * 32 // 64310            self.visual = ModifiedResNet(311                layers=vision_layers,312                output_dim=embed_dim,313                heads=vision_heads,314                input_resolution=image_resolution,315                width=vision_width316            )317        else:318            vision_heads = vision_width // 64319            self.visual = VisionTransformer(320                input_resolution=image_resolution,321                patch_size=vision_patch_size,322                width=vision_width,323                layers=vision_layers,324                heads=vision_heads,325                output_dim=embed_dim326            )327 328        self.transformer = Transformer(329            width=transformer_width,330            layers=transformer_layers,331            heads=transformer_heads,332            attn_mask=self.build_attention_mask()333        )334 335        self.vocab_size = vocab_size336        self.token_embedding = nn.Embedding(vocab_size, transformer_width)337        self.positional_embedding = nn.Parameter(torch.empty(self.context_length, transformer_width))338        self.ln_final = LayerNorm(transformer_width)339 340        self.text_projection = nn.Parameter(torch.empty(transformer_width, embed_dim))341        self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))342 343        self.initialize_parameters()344 345    def initialize_parameters(self):346        nn.init.normal_(self.token_embedding.weight, std=0.02)347        nn.init.normal_(self.positional_embedding, std=0.01)348 349        if isinstance(self.visual, ModifiedResNet):350            if self.visual.attnpool is not None:351                std = self.visual.attnpool.c_proj.in_features ** -0.5352                nn.init.normal_(self.visual.attnpool.q_proj.weight, std=std)353                nn.init.normal_(self.visual.attnpool.k_proj.weight, std=std)354                nn.init.normal_(self.visual.attnpool.v_proj.weight, std=std)355                nn.init.normal_(self.visual.attnpool.c_proj.weight, std=std)356 357            for resnet_block in [self.visual.layer1, self.visual.layer2, self.visual.layer3, self.visual.layer4]:358                for name, param in resnet_block.named_parameters():359                    if name.endswith("bn3.weight"):360                        nn.init.zeros_(param)361 362        proj_std = (self.transformer.width ** -0.5) * ((2 * self.transformer.layers) ** -0.5)363        attn_std = self.transformer.width ** -0.5364        fc_std = (2 * self.transformer.width) ** -0.5365        for block in self.transformer.resblocks:366            nn.init.normal_(block.attn.in_proj_weight, std=attn_std)367            nn.init.normal_(block.attn.out_proj.weight, std=proj_std)368            nn.init.normal_(block.mlp.c_fc.weight, std=fc_std)369            nn.init.normal_(block.mlp.c_proj.weight, std=proj_std)370 371        if self.text_projection is not None:372            nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)373 374    def build_attention_mask(self):375        # lazily create causal attention mask, with full attention between the vision tokens376        # pytorch uses additive attention mask; fill with -inf377        mask = torch.empty(self.context_length, self.context_length)378        mask.fill_(float("-inf"))379        mask.triu_(1)  # zero out the lower diagonal380        return mask381 382    @property383    def dtype(self):384        return self.visual.conv1.weight.dtype385 386    def encode_image(self, image):387        return self.visual(image.type(self.dtype))388 389    def encode_text(self, text):390        x = self.token_embedding(text).type(self.dtype)  # [batch_size, n_ctx, d_model]391 392        x = x + self.positional_embedding.type(self.dtype)393        x = x.permute(1, 0, 2)  # NLD -> LND394        x = self.transformer(x)395        x = x.permute(1, 0, 2)  # LND -> NLD396        x = self.ln_final(x).type(self.dtype)397 398        # x.shape = [batch_size, n_ctx, transformer.width]399        # take features from the eot embedding (eot_token is the highest number in each sequence)400        x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection401 402        return x403 404    def forward(self, image, text):405        image_features = self.encode_image(image)        # 经过修改,self.encode_image(image)输出的是每一层Encoder的[cls]feature406        text_features = self.encode_text(text)407 408        # 对倒数3层的[cls]feature做平均409        image_features = (image_features['layer'+str(self.vision_layers-1)]+image_features['layer'+str(self.vision_layers-2)]+image_features['layer'+str(self.vision_layers-3)])/3410        # 对倒数3层的[cls]feature做加权平均411        # image_features = 0.5*image_features['layer'+str(self.vision_layers-1)] + 0.3*image_features['layer'+str(self.vision_layers-2)] + 0.2*image_features['layer'+str(self.vision_layers-3)]412        413        # normalized features414        image_features = image_features / image_features.norm(dim=1, keepdim=True)415        text_features = text_features / text_features.norm(dim=1, keepdim=True)416 417        # cosine similarity as logits418        logit_scale = self.logit_scale.exp()419        logits_per_image = logit_scale * image_features @ text_features.t()420        logits_per_text = logits_per_image.t()421 422        # shape = [global_batch_size, global_batch_size]423        return logits_per_image, logits_per_text424 425 426def convert_weights(model: nn.Module):427    """Convert applicable model parameters to fp16"""428 429    def _convert_weights_to_fp16(l):430        if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)):431            l.weight.data = l.weight.data.half()432            if l.bias is not None:433                l.bias.data = l.bias.data.half()434 435        if isinstance(l, nn.MultiheadAttention):436            for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]:437                tensor = getattr(l, attr)438                if tensor is not None:439                    tensor.data = tensor.data.half()440 441        for name in ["text_projection", "proj"]:442            if hasattr(l, name):443                attr = getattr(l, name)444                if attr is not None:445                    attr.data = attr.data.half()446 447    model.apply(_convert_weights_to_fp16)448 449 450def build_model(state_dict: dict):451    vit = "visual.proj" in state_dict452 453    if vit:454        vision_width = state_dict["visual.conv1.weight"].shape[0]455        vision_layers = len([k for k in state_dict.keys() if k.startswith("visual.") and k.endswith(".attn.in_proj_weight")])456        vision_patch_size = state_dict["visual.conv1.weight"].shape[-1]457        grid_size = round((state_dict["visual.positional_embedding"].shape[0] - 1) ** 0.5)458        image_resolution = vision_patch_size * grid_size459    else:460        counts: list = [len(set(k.split(".")[2] for k in state_dict if k.startswith(f"visual.layer{b}"))) for b in [1, 2, 3, 4]]461        vision_layers = tuple(counts)462        vision_width = state_dict["visual.layer1.0.conv1.weight"].shape[0]463        output_width = round((state_dict["visual.attnpool.positional_embedding"].shape[0] - 1) ** 0.5)464        vision_patch_size = None465        assert output_width ** 2 + 1 == state_dict["visual.attnpool.positional_embedding"].shape[0]466        image_resolution = output_width * 32467 468    embed_dim = state_dict["text_projection"].shape[1]469    context_length = state_dict["positional_embedding"].shape[0]470    vocab_size = state_dict["token_embedding.weight"].shape[0]471    transformer_width = state_dict["ln_final.weight"].shape[0]472    transformer_heads = transformer_width // 64473    transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith("transformer.resblocks")))474 475    model = CLIP(476        embed_dim,477        image_resolution, vision_layers, vision_width, vision_patch_size,478        context_length, vocab_size, transformer_width, transformer_heads, transformer_layers479    )480 481    for key in ["input_resolution", "context_length", "vocab_size"]:482        if key in state_dict:483            del state_dict[key]484 485    convert_weights(model)486    model.load_state_dict(state_dict)487    return model.eval()