""" 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"] # 2816 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 config — used as a fallback if the checkpoint does not embed its own. 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