import torch
import torch.nn as nn
import torchvision as tv
from torch.utils.data import DataLoader
from torch import optim
from PIL import Image

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

# Image transform
transform = tv.transforms.Compose([
    tv.transforms.Resize((224,224)),
    tv.transforms.ToTensor()
])

# Load dataset (folder must contain 'fake' and 'real')
dataset = tv.datasets.ImageFolder("/content/drive/MyDrive/MSCIT/SEM2/CF/content/train", transform=transform)
loader = DataLoader(dataset, batch_size=8, shuffle=True)

# Simple CNN model
class DeepFakeNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(3,8,3,padding=1), nn.ReLU(), nn.MaxPool2d(2),
            nn.Conv2d(8,16,3,padding=1), nn.ReLU(), nn.MaxPool2d(2)
        )
        self.fc = nn.Sequential(
            nn.Flatten(),
            nn.Linear(16*56*56,2)
        )

    def forward(self,x):
        x = self.conv(x)
        x = self.fc(x)
        return x

model = DeepFakeNet().to(device)

# Loss and optimizer
loss_fn = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# Training loop
for images, labels in loader:
    images, labels = images.to(device), labels.to(device)

    optimizer.zero_grad()
    outputs = model(images)
    loss = loss_fn(outputs, labels)
    loss.backward()
    optimizer.step()

print("Training Done!")

# Save model
torch.save(model.state_dict(), "deepfake_model.pth")

import torch
from PIL import Image
import torchvision.transforms as T

model = DeepFakeNet()
model.load_state_dict(torch.load("deepfake_model.pth"))
model.eval()

transform = T.Compose([
    T.Resize((224,224)),
    T.ToTensor()
])

img = Image.open("/content/drive/MyDrive/MSCIT/SEM2/CF/content/train/real/real1.jpg").convert("RGB")
img = transform(img).unsqueeze(0)

output = model(img)
pred = torch.argmax(output)

print("Prediction:", "Fake" if pred==0 else "Real")