David310/Detect_AI-generated_Image
4
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()