import gradio as gr import torch import torch.nn as nn import torch.nn.functional as F import numpy as np import cv2 from PIL import Image import requests import io import os from torchvision import transforms from transformers import pipeline # ── Model Definition ──────────────────────────────────────── class CIFAKECNN(nn.Module): def __init__(self, labelnum=1): super(CIFAKECNN, self).__init__() self.conv1 = nn.Conv2d(in_channels=3, out_channels=32, kernel_size=3, padding=0) self.pool = nn.MaxPool2d(kernel_size=2, stride=2) self.conv2 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, padding=0) self.conv3 = nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=0) self.fc1 = nn.Linear(128*2*2, 256) self.dropout = nn.Dropout(0.5) self.fc2 = nn.Linear(256, labelnum) def forward(self, x): x = self.pool(F.relu(self.conv1(x))) x = self.pool(F.relu(self.conv2(x))) x = self.pool(F.relu(self.conv3(x))) x = x.view(x.size(0), -1) x = F.relu(self.fc1(x)) x = self.dropout(x) x = torch.sigmoid(self.fc2(x)) return x # ── Load our CNN ───────────────────────────────────────────── device = torch.device('cpu') model = CIFAKECNN().to(device) model.load_state_dict(torch.load('cnn_model.pth', map_location=device)) model.eval() # ── Load pretrained detector ───────────────────────────────── pretrained_detector = pipeline("image-classification", model="Organika/sdxl-detector") # ── Transform ──────────────────────────────────────────────── transform = transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]) ]) # ── GradCAM ────────────────────────────────────────────────── class GradCAM: def __init__(self, model, target_layer): self.model = model self.gradients = None self.activations = None target_layer.register_forward_hook(self._save_activation) target_layer.register_backward_hook(self._save_gradient) def _save_activation(self, module, input, output): self.activations = output.detach() def _save_gradient(self, module, grad_input, grad_output): self.gradients = grad_output[0].detach() def generate(self, image_tensor): self.model.zero_grad() image_tensor = image_tensor.unsqueeze(0).to(device) image_tensor.requires_grad = True output = self.model(image_tensor) output.backward() weights = self.gradients.mean(dim=[2, 3], keepdim=True) cam = (weights * self.activations).sum(dim=1, keepdim=True) cam = torch.relu(cam) cam = cam.squeeze().cpu().numpy() cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8) return cam grad_cam = GradCAM(model, model.conv3) def get_gradcam_overlay(image_pil, image_tensor): cam = grad_cam.generate(image_tensor) img = np.array(image_pil.resize((32, 32))) cam_resized = cv2.resize(cam, (32, 32)) heatmap = cv2.applyColorMap((cam_resized * 255).astype(np.uint8), cv2.COLORMAP_JET) heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB) overlaid = (0.5 * img + 0.5 * heatmap).astype(np.uint8) overlaid_large = Image.fromarray(overlaid).resize((224, 224), Image.LANCZOS) return overlaid_large # ── Detection Function ─────────────────────────────────────── def detect_image(image): if image is None: return "Please upload an image.", "Please upload an image.", None image_pil = Image.fromarray(image).convert('RGB') image_tensor = transform(image_pil) # Our CNN with torch.no_grad(): output = model(image_tensor.unsqueeze(0)) confidence = output.item() label = "FAKE (AI-Generated)" if confidence > 0.5 else "REAL" confidence_pct = confidence if confidence > 0.5 else 1 - confidence our_result = f"**{label}**\nConfidence: {confidence_pct:.1%}" # Pretrained detector pretrained_result = pretrained_detector(image_pil) top = pretrained_result[0] pretrained_label = top['label'] pretrained_conf = top['score'] pretrained_out = f"**{pretrained_label}**\nConfidence: {pretrained_conf:.1%}" gradcam_img = get_gradcam_overlay(image_pil, image_tensor) return our_result, pretrained_out, gradcam_img # ── Generate & Detect Function ─────────────────────────────── HF_TOKEN = os.environ.get("HF_TOKEN") def generate_and_detect(prompt): if not prompt: return None, "Please enter a prompt.", "Please enter a prompt.", None response = requests.post( "https://router.huggingface.co/hf-inference/models/black-forest-labs/FLUX.1-schnell", headers={"Authorization": f"Bearer {HF_TOKEN}"}, json={"inputs": prompt} ) if response.status_code != 200: return None, f"Generation failed: {response.text}", "", None image_pil = Image.open(io.BytesIO(response.content)).convert('RGB') our_result, pretrained_out, gradcam_img = detect_image(np.array(image_pil)) return image_pil, our_result, pretrained_out, gradcam_img # ── Gradio Interface ───────────────────────────────────────── with gr.Blocks(title="AI Image Detector") as demo: gr.Markdown("# 🔍 Truth in the Noise: AI Image Detector") gr.Markdown("Detect whether an image is real or AI-generated. We compare our custom CNN (trained on CIFAKE) with a pretrained detector.") with gr.Tabs(): with gr.Tab("📤 Upload & Detect"): gr.Markdown("Upload any image to compare both models.") with gr.Row(): upload_input = gr.Image(label="Upload Image") with gr.Column(): gr.Markdown("**🔵 Our CNN (trained on CIFAKE 32×32)**") upload_our = gr.Markdown() gr.Markdown("**🟡 Pretrained Detector (Organika/sdxl-detector)**") upload_pretrained = gr.Markdown() upload_gradcam = gr.Image(label="GradCAM (Our CNN)") upload_btn = gr.Button("Detect", variant="primary") upload_btn.click(detect_image, inputs=upload_input, outputs=[upload_our, upload_pretrained, upload_gradcam]) with gr.Tab("🎨 Generate & Detect"): gr.Markdown("Generate an AI image and detect it with both models.") prompt_input = gr.Textbox(label="Prompt", placeholder="e.g. a cat sitting on a chair") generate_btn = gr.Button("Generate & Detect", variant="primary") with gr.Row(): generated_img = gr.Image(label="Generated Image") with gr.Column(): gr.Markdown("**🔵 Our CNN (trained on CIFAKE 32×32)**") generate_our = gr.Markdown() gr.Markdown("**🟡 Pretrained Detector (Organika/sdxl-detector)**") generate_pretrained = gr.Markdown() generate_gradcam = gr.Image(label="GradCAM (Our CNN)") generate_btn.click(generate_and_detect, inputs=prompt_input, outputs=[generated_img, generate_our, generate_pretrained, generate_gradcam]) demo.launch()