| """ |
| Model definition for the AI-vs-Real image detector. |
| |
| This MUST ship alongside best_model.pt — the checkpoint only stores weights |
| (a state_dict), so the class definitions here are required to reconstruct the |
| network before loading those weights. |
| |
| Architecture: 3-branch ensemble |
| - CLIP ViT-L-14 (frozen) + trainable adapter -> 768 (semantic) |
| - EfficientNet-B3 (fine-tuned) -> 1536 (texture) |
| - FFT-CNN (custom, on Fourier spectrum) -> 512 (frequency) |
| concatenated -> fusion MLP -> 2 logits (0=real, 1=ai) |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| import torchvision.models as tv_models |
| import open_clip |
|
|
|
|
| class FFTBranch(nn.Module): |
| """Detects frequency-domain fingerprints left by generator upsampling.""" |
|
|
| def __init__(self, out_dim=512): |
| super().__init__() |
| self.cnn = nn.Sequential( |
| nn.Conv2d(1, 32, 3, padding=1, bias=False), nn.BatchNorm2d(32), nn.ReLU(True), nn.MaxPool2d(2), |
| nn.Conv2d(32, 64, 3, padding=1, bias=False), nn.BatchNorm2d(64), nn.ReLU(True), nn.MaxPool2d(2), |
| nn.Conv2d(64, 128, 3, padding=1, bias=False), nn.BatchNorm2d(128), nn.ReLU(True), nn.MaxPool2d(2), |
| nn.Conv2d(128, 256, 3, padding=1, bias=False), nn.BatchNorm2d(256), nn.ReLU(True), |
| nn.AdaptiveAvgPool2d(1), nn.Flatten(), |
| ) |
| self.proj = nn.Sequential(nn.Linear(256, out_dim), nn.LayerNorm(out_dim), nn.GELU()) |
|
|
| def fft(self, x): |
| gray = 0.299 * x[:, 0] + 0.587 * x[:, 1] + 0.114 * x[:, 2] |
| mag = torch.log1p(torch.abs(torch.fft.fft2(gray))) |
| mag = torch.fft.fftshift(mag, dim=(-2, -1)) |
| B = mag.shape[0] |
| mn = mag.view(B, -1).min(1).values.view(B, 1, 1) |
| mx = mag.view(B, -1).max(1).values.view(B, 1, 1) |
| return ((mag - mn) / (mx - mn + 1e-8)).unsqueeze(1) |
|
|
| def forward(self, x): |
| return self.proj(self.cnn(self.fft(x))) |
|
|
|
|
| class CLIPBranch(nn.Module): |
| """Frozen CLIP visual encoder + small trainable adapter (semantic view).""" |
|
|
| def __init__(self, model_name="ViT-L-14", pretrained="openai"): |
| super().__init__() |
| clip_model, _, _ = open_clip.create_model_and_transforms(model_name, pretrained=pretrained) |
| self.visual = clip_model.visual |
| self.visual.eval() |
| for p in self.visual.parameters(): |
| p.requires_grad = False |
| self.register_buffer("mean", torch.tensor([0.48145466, 0.4578275, 0.40821073]).view(1, 3, 1, 1)) |
| self.register_buffer("std", torch.tensor([0.26862954, 0.26130258, 0.27577711]).view(1, 3, 1, 1)) |
| d = 768 |
| self.adapter = nn.Sequential( |
| nn.Linear(d, d), nn.LayerNorm(d), nn.GELU(), nn.Dropout(0.1), |
| nn.Linear(d, d), nn.LayerNorm(d), |
| ) |
|
|
| def forward(self, x): |
| with torch.no_grad(): |
| f = self.visual((x - self.mean) / self.std) |
| return self.adapter(f) |
|
|
|
|
| class EfficientNetBranch(nn.Module): |
| """Fine-tuned EfficientNet-B3 (local texture / spatial artifacts).""" |
|
|
| def __init__(self, dropout=0.4): |
| super().__init__() |
| m = tv_models.efficientnet_b3(weights=tv_models.EfficientNet_B3_Weights.IMAGENET1K_V1) |
| self.features = m.features |
| self.avgpool = m.avgpool |
| self.drop = nn.Dropout(dropout) |
| self.register_buffer("mean", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)) |
| self.register_buffer("std", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)) |
|
|
| def forward(self, x): |
| x = (x - self.mean) / self.std |
| return self.drop(torch.flatten(self.avgpool(self.features(x)), 1)) |
|
|
|
|
| class ArtifactDetector(nn.Module): |
| """Full 3-branch ensemble with a fusion head.""" |
|
|
| def __init__(self, cfg): |
| super().__init__() |
| self.clip = CLIPBranch(cfg["clip_model"], cfg["clip_pretrain"]) |
| self.effnet = EfficientNetBranch(cfg["effnet_dropout"]) |
| self.fft = FFTBranch(cfg["fft_out_dim"]) |
| total = 768 + 1536 + cfg["fft_out_dim"] |
| self.fusion = nn.Sequential( |
| nn.Linear(total, 512), nn.BatchNorm1d(512), nn.GELU(), nn.Dropout(0.4), |
| nn.Linear(512, 128), nn.BatchNorm1d(128), nn.GELU(), nn.Dropout(0.2), |
| nn.Linear(128, 2), |
| ) |
|
|
| def forward(self, x): |
| return self.fusion(torch.cat([self.clip(x), self.effnet(x), self.fft(x)], dim=1)) |
|
|
|
|
| |
| DEFAULT_CFG = { |
| "clip_model": "ViT-L-14", |
| "clip_pretrain": "openai", |
| "effnet_dropout": 0.4, |
| "fft_out_dim": 512, |
| "image_size": 224, |
| } |
|
|
|
|
| def load_detector(checkpoint_path, device="cpu"): |
| """Rebuild the model from its embedded cfg and load trained weights.""" |
| ckpt = torch.load(checkpoint_path, map_location=device) |
| cfg = ckpt.get("cfg", DEFAULT_CFG) |
| model = ArtifactDetector(cfg).to(device) |
| model.load_state_dict(ckpt["model_state"]) |
| model.eval() |
| return model, cfg |
|
|