File size: 3,287 Bytes
8e9877b
 
 
 
 
 
 
 
e64d6cd
4a75ba1
 
e64d6cd
8e9877b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e89f54f
 
6938545
 
 
e89f54f
 
6938545
 
 
 
 
 
 
e89f54f
8e9877b
 
 
 
 
 
 
e89f54f
 
 
e64d6cd
8e9877b
 
e89f54f
8e9877b
 
e89f54f
 
4a75ba1
 
 
 
 
 
 
8e9877b
 
 
e64d6cd
8e9877b
 
 
e89f54f
8e9877b
 
e64d6cd
e89f54f
8e9877b
 
 
e64d6cd
8e9877b
 
e89f54f
e64d6cd
 
 
 
e89f54f
 
8e9877b
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
import gradio as gr
import torch
from torch import nn
from torchvision import models, transforms
import cv2
import numpy as np
import os

# Load OpenCV face detector
face_cascade = cv2.CascadeClassifier(cv2.data.haarcascades + 'haarcascade_frontalface_default.xml')

# --- Model Architecture ---
class DeepfakeModel(nn.Module):
    def __init__(self, num_classes=2, latent_dim=2048, lstm_layers=1, hidden_dim=2048, bidirectional=False):
        super(DeepfakeModel, self).__init__()
        resnext = models.resnext50_32x4d(pretrained=True)
        self.model = nn.Sequential(*list(resnext.children())[:-2])
        self.lstm = nn.LSTM(latent_dim, hidden_dim, lstm_layers, bidirectional)
        self.dp = nn.Dropout(0.4)
        self.linear1 = nn.Linear(2048, num_classes)
        self.avgpool = nn.AdaptiveAvgPool2d(1)

    def forward(self, x):
        batch_size, seq_length, c, h, w = x.shape
        x = x.view(batch_size * seq_length, c, h, w)
        fmap = self.model(x)
        x = self.avgpool(fmap)
        x = x.view(batch_size, seq_length, 2048)
        x_lstm, _ = self.lstm(x, None)
        return fmap, self.dp(self.linear1(x_lstm[:, -1, :]))

device = "cuda" if torch.cuda.is_available() else "cpu"

# Load Model Weights
model = DeepfakeModel().to(device)

MODEL_PATH = "model_93_acc_100_frames_celeb_FF_data.pt"

model_found = False

try:
    model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
    model.eval()
    model_found = True
    print(f"✅ Loaded model: {MODEL_PATH}")
except Exception as e:
    print(f"❌ Error loading model: {e}")

transform = transforms.Compose([
    transforms.ToPILImage(),
    transforms.Resize((112, 112)),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])

def predict(video):
    if not model_found:
        return "❌ Error: No .pt model file found in Space.", 0.0

    cap = cv2.VideoCapture(video)
    frames = []
    # Extract up to 20 frames with faces
    while cap.isOpened() and len(frames) < 20:
        ret, frame = cap.read()
        if not ret: break
        
        gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
        faces = face_cascade.detectMultiScale(gray, 1.1, 4)
        if len(faces) > 0:
            (x, y, w, h) = faces[0]
            face_crop = frame[y:y+h, x:x+w]
            rgb_face = cv2.cvtColor(face_crop, cv2.COLOR_BGR2RGB)
            frames.append(transform(rgb_face))
    cap.release()

    if len(frames) < 5:
        return "⚠️ No faces detected in video.", 0.0

    input_tensor = torch.stack(frames).unsqueeze(0).to(device)
    with torch.no_grad():
        _, outputs = model(input_tensor)
        probabilities = torch.softmax(outputs, dim=1)
        confidence, prediction = torch.max(probabilities, 1)

    label = "REAL" if prediction.item() == 1 else "FAKE"
    conf = round(float(confidence.item()) * 100, 2)
    return label, conf

# --- Gradio UI ---
demo = gr.Interface(
    fn=predict,
    inputs=gr.Video(label="Upload Video"),
    outputs=[
        gr.Textbox(label="Result"),
        gr.Number(label="Confidence (%)")
    ],
    title="Deepfake Detection API",
    description="Upload a video to analyze it for deepfakes using our ResNext-LSTM model."
)

if __name__ == "__main__":
    demo.launch()