Deepfake / inference.py
Prabhsimran09's picture
Update inference.py
866ee65 verified
Raw
History Blame Contribute Delete
4.78 kB
import torch
import torch.nn.functional as F
import numpy as np
import cv2
import io
import matplotlib.pyplot as plt
from PIL import Image
from torchvision import transforms
from model import DeepfakeClassifier
# Device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Load model
model = DeepfakeClassifier()
state_dict = torch.load("model_weights/model_87_acc_20_frames_final_data.pt", map_location=device)
model.load_state_dict(state_dict, strict=False)
model.to(device)
model.eval()
# ✅ Image transform (ImageNet normalization restored)
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406],
[0.229, 0.224, 0.225])
])
# Softmax with temperature scaling
def softmax_with_temperature(logits, T=0.5):
return torch.softmax(logits / T, dim=1)
# Grad-CAM logic
def apply_gradcam(image_tensor, target_layer):
gradients, activations = [], []
def backward_hook(module, grad_input, grad_output):
gradients.append(grad_output[0])
def forward_hook(module, input, output):
activations.append(output)
handle_fwd = target_layer.register_forward_hook(forward_hook)
handle_bwd = target_layer.register_backward_hook(backward_hook)
output = model(image_tensor)
class_idx = output.argmax().item()
score = output[0, class_idx]
model.zero_grad()
score.backward()
grads = gradients[0]
acts = activations[0]
weights = grads.mean(dim=[2, 3], keepdim=True)
cam = (weights * acts).sum(dim=1).squeeze().cpu().detach().numpy()
cam = np.maximum(cam, 0)
cam = cv2.resize(cam, (224, 224))
cam -= cam.min()
cam /= (cam.max() + 1e-8)
handle_fwd.remove()
handle_bwd.remove()
return cam, class_idx
# Grad-CAM overlay
def overlay_cam(image, cam):
heatmap = cv2.applyColorMap(np.uint8(255 * cam), cv2.COLORMAP_JET)
image = np.array(image.resize((224, 224)).convert("RGB"))
overlay = cv2.addWeighted(image, 0.6, heatmap, 0.4, 0)
return overlay
# ✅ Predict image (with normalization and T=1.0)
def predict_image(pil_image):
image_tensor = transform(pil_image).unsqueeze(0).to(device)
cam, _ = apply_gradcam(image_tensor, model.model[6][2].conv3)
prob = softmax_with_temperature(model(image_tensor), T=0.5)[0]
print("Image Probabilities:", prob.tolist()) # Optional debug
fake_confidence = prob[0].item()
label = "Fake" if fake_confidence > 0.7 else "Real"
confidence = fake_confidence if label == "Fake" else prob[1].item()
overlay = overlay_cam(pil_image, cam)
return f"{label} ({confidence:.2f} confidence)", Image.fromarray(overlay)
# Predict video
def predict_video(video_path):
cap = cv2.VideoCapture(video_path)
frames = []
while len(frames) < 20:
ret, frame = cap.read()
if not ret:
break
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
pil_img = Image.fromarray(frame)
frames.append(transform(pil_img))
cap.release()
if len(frames) == 0:
return "No valid frames found.", [], None
# Per-frame prediction + Grad-CAM
frame_probs = []
cam_images = []
for frame in frames:
frame_tensor = frame.unsqueeze(0).to(device)
with torch.no_grad():
out = model(frame_tensor)
prob = softmax_with_temperature(out, T=2.0)[0][0].item()
frame_probs.append(prob)
cam, _ = apply_gradcam(frame_tensor, model.model[6][2].conv3)
frame_np = (frame.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)
frame_np = np.clip(frame_np, 0, 255)
pil_frame = Image.fromarray(frame_np)
overlay = overlay_cam(pil_frame, cam)
cam_images.append(Image.fromarray(overlay))
# Sequence prediction (LSTM)
if len(frames) < 20:
frames += [frames[-1]] * (20 - len(frames))
video_tensor = torch.stack(frames).unsqueeze(0).to(device)
with torch.no_grad():
output = model(video_tensor)
seq_prob = softmax_with_temperature(output, T=2.0)[0][0].item()
seq_label = "Fake" if seq_prob > 0.8 else "Real"
# Confidence chart
fig, ax = plt.subplots()
ax.plot(frame_probs, marker='o', label='Fake Probability')
ax.axhline(0.5, color='red', linestyle='--', label='Threshold')
ax.set_title('Per-frame Fake Confidence')
ax.set_xlabel('Frame Index')
ax.set_ylabel('Probability')
ax.legend()
ax.grid(True)
buf = io.BytesIO()
plt.savefig(buf, format='png')
plt.close(fig)
buf.seek(0)
chart_img = Image.open(buf)
# Final decision
label_text = f"Final Decision (Sequence): {seq_label} ({seq_prob:.2f})"
return label_text, cam_images, chart_img