Spaces:
Sleeping
Sleeping
| 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() |