Emilyl613's picture
Update app.py
d1dc6dc verified
Raw
History Blame Contribute Delete
7.95 kB
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()