{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":16880,"databundleVersionId":858837,"sourceType":"competition"},{"sourceId":924245,"sourceType":"datasetVersion","datasetId":464091},{"sourceId":1010500,"sourceType":"datasetVersion","datasetId":536104}],"dockerImageVersionId":29844,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **About the dataset**","metadata":{}},{"cell_type":"markdown","source":"This dataset contains faces extracted from deepfake-detection-challenge. All images were of size 224x224. \n\nDue to memory issue we will only use a sample of the entire dataset for prediction.","metadata":{}},{"cell_type":"markdown","source":"# **Import packages**","metadata":{}},{"cell_type":"code","source":"!pip install --quiet lime","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:08.896406Z","iopub.execute_input":"2024-11-01T14:28:08.896748Z","iopub.status.idle":"2024-11-01T14:28:15.040146Z","shell.execute_reply.started":"2024-11-01T14:28:08.896686Z","shell.execute_reply":"2024-11-01T14:28:15.039237Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import statements\nimport numpy as np \nimport pandas as pd \nimport cv2\n\nimport os\nimport urllib.request \nimport sys\nimport glob\nimport json\n\nimport sklearn\nfrom sklearn.model_selection import train_test_split\n\nimport random\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport plotly.graph_objs as go","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-11-01T14:28:15.042848Z","iopub.execute_input":"2024-11-01T14:28:15.043217Z","iopub.status.idle":"2024-11-01T14:28:15.049792Z","shell.execute_reply.started":"2024-11-01T14:28:15.043145Z","shell.execute_reply":"2024-11-01T14:28:15.048879Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import Neural Network and PyTorch Libraries\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as f\nimport torch.optim as optim\n\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision import models\nimport torch.optim.lr_scheduler as lr_scheduler\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.051362Z","iopub.execute_input":"2024-11-01T14:28:15.051698Z","iopub.status.idle":"2024-11-01T14:28:15.060124Z","shell.execute_reply.started":"2024-11-01T14:28:15.051636Z","shell.execute_reply":"2024-11-01T14:28:15.059386Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.rc('font', size=14)\nplt.rc('axes', labelsize=14, titlesize=14)\nplt.rc('legend', fontsize=14)\nplt.rc('xtick', labelsize=10)\nplt.rc('ytick', labelsize=10)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.061494Z","iopub.execute_input":"2024-11-01T14:28:15.061736Z","iopub.status.idle":"2024-11-01T14:28:15.068207Z","shell.execute_reply.started":"2024-11-01T14:28:15.061692Z","shell.execute_reply":"2024-11-01T14:28:15.067159Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device('cuda')\n    print(\"CUDA is available. Using GPU.\")\nelse:\n    device = torch.device('cpu')\n    print(\"CUDA is not available. Using CPU.\")","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.071025Z","iopub.execute_input":"2024-11-01T14:28:15.071281Z","iopub.status.idle":"2024-11-01T14:28:15.077039Z","shell.execute_reply.started":"2024-11-01T14:28:15.071230Z","shell.execute_reply":"2024-11-01T14:28:15.076355Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Data Visualization**","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/deepfake-faces/metadata.csv\")\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.080873Z","iopub.execute_input":"2024-11-01T14:28:15.081237Z","iopub.status.idle":"2024-11-01T14:28:15.199296Z","shell.execute_reply.started":"2024-11-01T14:28:15.081170Z","shell.execute_reply":"2024-11-01T14:28:15.198411Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.shape","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.200541Z","iopub.execute_input":"2024-11-01T14:28:15.200769Z","iopub.status.idle":"2024-11-01T14:28:15.206728Z","shell.execute_reply.started":"2024-11-01T14:28:15.200731Z","shell.execute_reply":"2024-11-01T14:28:15.206037Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(df[df[\"label\"] == \"FAKE\"]), len(df[df[\"label\"] == \"REAL\"])","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.208828Z","iopub.execute_input":"2024-11-01T14:28:15.209306Z","iopub.status.idle":"2024-11-01T14:28:15.250802Z","shell.execute_reply.started":"2024-11-01T14:28:15.209100Z","shell.execute_reply":"2024-11-01T14:28:15.250111Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"I only use 16,000 images for train/validate/test","metadata":{}},{"cell_type":"code","source":"real_df = df[df[\"label\"] == \"FAKE\"]\nfake_df = df[df[\"label\"] == \"REAL\"]\nSAMPLE_SIZE = 8000\n\nreal_df = real_df.sample(SAMPLE_SIZE, random_state=42)\nfake_df = fake_df.sample(SAMPLE_SIZE, random_state=42)\n\nsample_df = pd.concat([real_df, fake_df])","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.252119Z","iopub.execute_input":"2024-11-01T14:28:15.252337Z","iopub.status.idle":"2024-11-01T14:28:15.304917Z","shell.execute_reply.started":"2024-11-01T14:28:15.252300Z","shell.execute_reply":"2024-11-01T14:28:15.304291Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data, test_data = train_test_split(\n    sample_df,\n    test_size=0.2,\n    random_state=42,\n    stratify=sample_df[\"label\"]\n)\n\ntest_data, val_data = train_test_split(\n    test_data,\n    test_size=0.5,\n    random_state=42,\n    stratify=test_data[\"label\"]\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.306106Z","iopub.execute_input":"2024-11-01T14:28:15.306325Z","iopub.status.idle":"2024-11-01T14:28:15.354803Z","shell.execute_reply.started":"2024-11-01T14:28:15.306288Z","shell.execute_reply":"2024-11-01T14:28:15.354061Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.shape, val_data.shape, test_data.shape","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.356137Z","iopub.execute_input":"2024-11-01T14:28:15.356416Z","iopub.status.idle":"2024-11-01T14:28:15.362578Z","shell.execute_reply.started":"2024-11-01T14:28:15.356368Z","shell.execute_reply":"2024-11-01T14:28:15.361575Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\n\n# Assuming train_data, val_data, and test_data are pandas DataFrames\ny = dict()\n\ny[0] = []\ny[1] = []\n\n# Iterating through each label set\nfor set_name in (np.array(train_data['label']), np.array(val_data['label']), np.array(test_data['label'])):\n    y[0].append(np.sum(set_name == 'REAL'))  # Count the 'REAL' labels\n    y[1].append(np.sum(set_name == 'FAKE'))  # Count the 'FAKE' labels\n\n# Define the bar positions and width\nlabels = ['Train Set', 'Validation Set', 'Test Set']\nx = np.arange(len(labels))  # Label locations\nwidth = 0.35  # Width of the bars\n\n# Create the bar chart\nfig, ax = plt.subplots(figsize=(8, 6))\n\n# Create the bars for 'REAL' and 'FAKE' categories\nrects1 = ax.bar(x - width/2, y[0], width, label='REAL', color='#33cc33', alpha=0.7)\nrects2 = ax.bar(x + width/2, y[1], width, label='FAKE', color='#ff3300', alpha=0.7)\n\n# Add labels, title, and custom ticks on the x-axis\nax.set_xlabel('Set')\nax.set_ylabel('Count')\nax.set_title('Count of classes in each set')\nax.set_xticks(x)\nax.set_xticklabels(labels)\nax.legend()\n\n# Display the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.364022Z","iopub.execute_input":"2024-11-01T14:28:15.364386Z","iopub.status.idle":"2024-11-01T14:28:15.650356Z","shell.execute_reply.started":"2024-11-01T14:28:15.364329Z","shell.execute_reply":"2024-11-01T14:28:15.649281Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The original image dataset were biased with more fake images than real since we are taking a sample of it its better to take equal proportion of real and fake images.","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor cur, i in enumerate(train_data.index[25:50]):\n    plt.subplot(5, 5, cur + 1)\n    plt.xticks([])\n    plt.yticks([])\n    plt.grid(False)\n    \n    image_bgr = cv2.imread('/kaggle/input/deepfake-faces/faces_224/' + train_data.loc[i, 'videoname'][:-4] + '.jpg')\n    image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)  # Convert BGR to RGB\n    \n    # Display the image\n    plt.imshow(image_rgb)\n    \n    if(train_data.loc[i, 'label'] == 'FAKE'):\n        plt.xlabel('FAKE Image')\n    else:\n        plt.xlabel('REAL Image')\n        \nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:15.652184Z","iopub.execute_input":"2024-11-01T14:28:15.652794Z","iopub.status.idle":"2024-11-01T14:28:17.117127Z","shell.execute_reply.started":"2024-11-01T14:28:15.652707Z","shell.execute_reply":"2024-11-01T14:28:17.116188Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Create dataset**","metadata":{}},{"cell_type":"markdown","source":"FAKE faces are labeled `1`, while REAL faces are labeled `0`","metadata":{}},{"cell_type":"code","source":"class DeepfakeDataset(Dataset):\n    def __init__(self, data_name, is_training=True):\n        \"\"\"\n        Args:\n            data_name (pd.DataFrame): DataFrame containing 'videoname' and 'label' columns.\n        \"\"\"\n        self.data_name = data_name\n        self.is_training = is_training\n        \n        # Initialize the images and labels using the retrieve_data function logic\n        self.images, self.labels = self.retrieve_data()\n        \n        # Training augmentations\n        if is_training:\n            self.transform_func = A.Compose([\n                A.Resize(224, 224),\n                A.RandomRotate90(p=0.5),\n                A.HorizontalFlip(p=0.5),\n                A.RandomBrightnessContrast(p=0.2),\n                A.OneOf([\n                    A.GaussNoise(p=1),\n                    A.GaussianBlur(p=1),\n                ], p=0.2),\n                A.Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2()\n            ])\n        else:\n            # Validation/test transforms\n            self.transform_func = A.Compose([\n                A.Resize(224, 224),\n                A.Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2()\n            ])\n    \n    def retrieve_data(self):\n        images, labels = [], []\n        \n        for img, label in zip(self.data_name[\"videoname\"], self.data_name[\"label\"]):\n            image = cv2.imread(\"/kaggle/input/deepfake-faces/faces_224/\" + img[:-4] + \".jpg\")\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  # Convert from BGR to RGB format\n            images.append(image)\n            \n            # Convert labels: FAKE -> 1, REAL -> 0\n            if label == \"FAKE\":\n                labels.append(1)\n            else:\n                labels.append(0)\n        \n        return np.array(images), np.array(labels)\n    \n    def __len__(self):\n        \"\"\"Returns the size of the dataset\"\"\"\n        return len(self.labels)\n    \n    def __getitem__(self, idx):\n        \"\"\"Fetch a single sample from the dataset\"\"\"\n        # Get image and label\n        image = self.images[idx]\n        label = self.labels[idx]\n        \n        # Apply transforms - note the dictionary format required by Albumentations\n        transformed = self.transform_func(image=image)\n        image = transformed[\"image\"]  # Get the transformed image from the result dictionary\n        \n        # Convert the label to a tensor\n        label = torch.tensor(label, dtype=torch.long)\n        \n        return image, label","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:17.118512Z","iopub.execute_input":"2024-11-01T14:28:17.118787Z","iopub.status.idle":"2024-11-01T14:28:17.137074Z","shell.execute_reply.started":"2024-11-01T14:28:17.118724Z","shell.execute_reply":"2024-11-01T14:28:17.136296Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = DeepfakeDataset(data_name=train_data)\nval_data = DeepfakeDataset(data_name=val_data)\ntest_data = DeepfakeDataset(data_name=test_data)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:28:17.138338Z","iopub.execute_input":"2024-11-01T14:28:17.138608Z","iopub.status.idle":"2024-11-01T14:29:02.998080Z","shell.execute_reply.started":"2024-11-01T14:28:17.138566Z","shell.execute_reply":"2024-11-01T14:29:02.997408Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(train_data, batch_size=16, shuffle=True)\nval_loader = DataLoader(val_data, batch_size=16, shuffle=True)\ntest_loader = DataLoader(test_data, batch_size=16, shuffle=True)\n\n# Verify if dataset is created accurately\nimages, labels = next(iter(train_loader))\nprint(images.shape, labels.shape)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:29:02.999449Z","iopub.execute_input":"2024-11-01T14:29:02.999671Z","iopub.status.idle":"2024-11-01T14:29:03.050139Z","shell.execute_reply.started":"2024-11-01T14:29:02.999634Z","shell.execute_reply":"2024-11-01T14:29:03.049407Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Modeling**","metadata":{}},{"cell_type":"markdown","source":"## **Create train and eval function**","metadata":{}},{"cell_type":"markdown","source":"- **Loss function:** `Weighted Cross Entropy Loss` with softmax activation function to return final probability.\n    - Apply the `compute_class_weight` to calculate percentage values per class for a weighted CrossEntropy. Since some classes have considerably more samples than others, all classes are weighted and taken as input into the loss calculation according to their respective number of samples\n- **Early stopping:** Stop training if loss is increasing.\n- **Optimizer:** Adam\n- **Learning rate:** 1e-4\n- **Epochs:** 10","metadata":{}},{"cell_type":"code","source":"class EarlyStopping:\n    def __init__(self, patience=4, verbose=False, delta=0):\n        self.patience = patience\n        self.verbose = verbose\n        self.counter = 0\n        self.best_loss = None\n        self.early_stop = False\n        self.delta = delta\n\n    def __call__(self, val_loss):\n        if self.best_loss is None:\n            self.best_loss = val_loss\n        elif val_loss > self.best_loss + self.delta:\n            self.counter += 1\n            if self.verbose:\n                print(f\"EarlyStopping counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_loss = val_loss\n            self.counter = 0","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:29:03.051604Z","iopub.execute_input":"2024-11-01T14:29:03.052110Z","iopub.status.idle":"2024-11-01T14:29:03.061742Z","shell.execute_reply.started":"2024-11-01T14:29:03.052047Z","shell.execute_reply":"2024-11-01T14:29:03.061069Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_losses = []\nval_losses = []\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs):\n    early_stopping = EarlyStopping(patience=5, verbose=True)\n\n    # Train loop\n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0.0\n\n        # Progress bar for training batches\n        train_bar = tqdm(train_loader, desc=f'Training Epoch {epoch + 1}/{num_epochs}', unit='batch')\n\n        for imgs, labels in train_bar:\n            imgs = imgs.to(device)\n            labels = labels.to(device)\n\n            optimizer.zero_grad()\n\n            # Forward\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n            # Backward\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item() * imgs.size(0)\n            train_losses.append(train_loss)\n\n            # Update progress bar with the current loss\n            train_bar.set_postfix(loss=loss.item())\n\n        # Validate\n        model.eval()\n        val_loss = 0.0\n        correct_preds = 0\n        with torch.no_grad():\n            val_bar = tqdm(val_loader, desc=f'Validating Epoch {epoch + 1}/{num_epochs}', unit='batch')\n\n            for imgs, labels in val_bar:\n                imgs = imgs.to(device)\n                labels = labels.to(device)\n\n                outputs = model(imgs)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item() * imgs.size(0)\n\n                _, preds = torch.max(outputs, dim=1)\n                correct_preds += torch.sum(preds == labels)\n\n                # Update validation progress bar with the current loss\n                val_bar.set_postfix(loss=loss.item())\n\n        val_loss = val_loss / len(val_loader.dataset)\n        val_losses.append(val_loss)\n        accuracy = correct_preds.double() / len(val_loader.dataset)\n        scheduler.step(val_loss)\n\n        print(f'Epoch {epoch + 1}/{num_epochs}, Training Loss: {train_loss / len(train_loader.dataset):.4f}, '\n              f'Validation Loss: {val_loss:.4f}, Accuracy: {accuracy:.4f}')\n\n        # Early stopping\n        early_stopping(val_loss)\n        if early_stopping.early_stop:\n            print(\"Early stopping triggered. Stopping training.\")\n            break\n","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:29:03.063138Z","iopub.execute_input":"2024-11-01T14:29:03.063430Z","iopub.status.idle":"2024-11-01T14:29:03.080168Z","shell.execute_reply.started":"2024-11-01T14:29:03.063387Z","shell.execute_reply":"2024-11-01T14:29:03.079291Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import precision_score, recall_score\n\ndef evaluate_model(model, test_loader, criterion):\n    model.eval()\n    test_losses = []\n    correct_preds = 0\n    \n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for imgs, labels in test_loader:\n            imgs = imgs.to(device)\n            labels = labels.to(device)\n\n            outputs = model(imgs)\n            test_loss = criterion(outputs, labels)\n            \n            _, preds = torch.max(outputs, dim=1)\n            correct_preds += torch.sum(preds == labels)\n\n            test_losses.append(test_loss.item())\n            all_preds.extend(preds.cpu().numpy())  \n            all_labels.extend(labels.cpu().numpy())  \n\n    accuracy = float((correct_preds.double() / len(test_loader.dataset)) * 100)\n    precision = precision_score(all_labels, all_preds)\n    recall = recall_score(all_labels, all_preds)\n\n    print(\"\\nAccuracy: \", accuracy)\n    print(\"Precision: \", precision)\n    print(\"Recall: \", recall)\n","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:29:03.081293Z","iopub.execute_input":"2024-11-01T14:29:03.081510Z","iopub.status.idle":"2024-11-01T14:29:03.093158Z","shell.execute_reply.started":"2024-11-01T14:29:03.081473Z","shell.execute_reply":"2024-11-01T14:29:03.092347Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Build ResNet50 model**","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass ResNet50Classifier(nn.Module):\n    def __init__(self, dropout_rate=0.3):\n        super(ResNet50Classifier, self).__init__()\n        \n        # Load pretrained ResNet50\n        self.model = models.resnet50(pretrained=True)\n        \n        # Remove the original fully connected layer\n        num_features = self.model.fc.in_features\n        self.model.fc = nn.Identity()  # Replace with Identity to get features\n        \n        # Create new classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(num_features, 256),  # num_features is 2048 for ResNet50\n            nn.ReLU(),\n            nn.Dropout(dropout_rate),\n            nn.Linear(256, 2)  # 2 classes: real and fake\n        )\n        \n    def forward(self, x):\n        if len(x.shape) == 5: \n            # Input shape: [batch, channels, sequence_length, height, width]\n            batch_size, channels, seq_len, height, width = x.shape\n\n            # Rearrange to [batch * sequence_length, channels, height, width]\n            x = x.view(batch_size * seq_len, channels, height, width)\n            \n            # Extract features using ResNet for each frame\n            features = self.model(x)  # Output shape: [batch * sequence_length, 2048]\n            \n            # Reshape to [batch_size, seq_len, 2048] for averaging over the sequence dimension\n            features = features.view(batch_size, seq_len, -1)\n            \n            # Average pool over frames in the sequence\n            features = features.mean(dim=1)  # Shape: [batch_size, 2048]\n        else:\n            # Single image input (4D tensor), directly extract features\n            features = self.model(x)  # Shape: [batch_size, 2048]\n        \n        # Apply classifier\n        x = self.classifier(features)\n        \n        return x\n\n    def get_features(self, x):\n        \"\"\"Extract features before classification layer\"\"\"\n        if len(x.shape) == 5:  # Video input\n            batch_size, channels, seq_len, height, width = x.shape\n            x = x.view(batch_size * seq_len, channels, height, width)\n            features = self.model(x)\n            features = features.view(batch_size, seq_len, -1).mean(dim=1)\n        else:  # Single image input\n            features = self.model(x)\n        return features\n\n    def freeze_backbone(self, freeze=True):\n        \"\"\"Freeze or unfreeze the feature extractor\"\"\"\n        for param in self.model.parameters():\n            param.requires_grad = not freeze\n\n    def unfreeze_last_layers(self, num_layers=2):\n        \"\"\"Unfreeze the last few layers of the backbone\"\"\"\n        # First freeze everything\n        self.freeze_backbone(True)\n        \n        # Then unfreeze the last layers\n        for layer in list(self.model.children())[-num_layers:]:\n            for param in layer.parameters():\n                param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:29:03.094407Z","iopub.execute_input":"2024-11-01T14:29:03.094847Z","iopub.status.idle":"2024-11-01T14:29:03.113051Z","shell.execute_reply.started":"2024-11-01T14:29:03.094770Z","shell.execute_reply":"2024-11-01T14:29:03.112309Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = ResNet50Classifier()\nmodel = model.to(device)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\nscheduler = lr_scheduler.ReduceLROnPlateau(optimizer, min_lr=1e-6, factor=0.5, patience=1, verbose=True)\nnum_epochs = 5\n\ntrain_model(model, train_loader,val_loader, criterion, optimizer, scheduler, num_epochs)\nprint(\"\\n\")\nevaluate_model(model, test_loader, criterion)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:29:03.114338Z","iopub.execute_input":"2024-11-01T14:29:03.114572Z","iopub.status.idle":"2024-11-01T14:43:30.421296Z","shell.execute_reply.started":"2024-11-01T14:29:03.114532Z","shell.execute_reply":"2024-11-01T14:43:30.420110Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nepochs = range(1, num_epochs + 1)\n\nplt.figure(figsize=(14, 5))\n\n# Plot Training and Validation Loss\nplt.plot(num_epochs, train_losses, label=\"Training Loss\")\nplt.plot(num_epochs, val_losses, label=\"Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training and Validation Loss\")\nplt.legend()\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Deepfake video classification**","metadata":{}},{"cell_type":"markdown","source":"## **About the dataset**","metadata":{}},{"cell_type":"markdown","source":"Files: \n- `train_sample_videos.zip` - A zip file containing a sample set of training videos and a metadata.json with labels. The full set of training videos is available through the links provided above.\n\n- `test_videos.zip` - A zip file containing a small set of videos to be used as a public validation set.","metadata":{}},{"cell_type":"markdown","source":"Metadata columns:\n- `filename` - The filename of the video\n- `label` - REAL or FAKE\n- `original` - In case a train set video is FAKE, the original video is listed here\n- `split` - This is always equal to \"train\"","metadata":{}},{"cell_type":"markdown","source":"## **Data visualization**","metadata":{}},{"cell_type":"code","source":"DATA_FOLDER = \"/kaggle/input/deepfake-detection-challenge\"\nTRAIN_SAMPLE_FOLDER = \"train_sample_videos\"\nTEST_FOLDER = \"test_videos\"\n\nprint(f\"Train samples: {len(os.listdir(os.path.join(DATA_FOLDER, TRAIN_SAMPLE_FOLDER)))}\")\nprint(f\"Test samples: {len(os.listdir(os.path.join(DATA_FOLDER, TEST_FOLDER)))}\")","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:30.422928Z","iopub.execute_input":"2024-11-01T14:43:30.423261Z","iopub.status.idle":"2024-11-01T14:43:30.433998Z","shell.execute_reply.started":"2024-11-01T14:43:30.423200Z","shell.execute_reply":"2024-11-01T14:43:30.433170Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sample_metadata = pd.read_json(\n    \"/kaggle/input/deepfake-detection-challenge/train_sample_videos/metadata.json\"\n).T\n\ntrain_sample_metadata.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:30.435539Z","iopub.execute_input":"2024-11-01T14:43:30.435880Z","iopub.status.idle":"2024-11-01T14:43:30.616942Z","shell.execute_reply.started":"2024-11-01T14:43:30.435820Z","shell.execute_reply":"2024-11-01T14:43:30.616110Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sample_metadata \\\n.groupby(\"label\")[\"label\"] \\\n.count() \\\n.plot(figsize=(8, 6), kind=\"bar\", title=\"Labels distribution in dataset\")","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:30.618181Z","iopub.execute_input":"2024-11-01T14:43:30.618401Z","iopub.status.idle":"2024-11-01T14:43:30.889132Z","shell.execute_reply.started":"2024-11-01T14:43:30.618363Z","shell.execute_reply":"2024-11-01T14:43:30.887972Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### **Display some FAKE videos**","metadata":{}},{"cell_type":"code","source":"fake_train_sample_videos = list(train_sample_metadata.loc[train_sample_metadata[\"label\"] == \"FAKE\"].sample(3).index)\n\nfake_train_sample_videos","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:30.891226Z","iopub.execute_input":"2024-11-01T14:43:30.891616Z","iopub.status.idle":"2024-11-01T14:43:30.905420Z","shell.execute_reply.started":"2024-11-01T14:43:30.891556Z","shell.execute_reply":"2024-11-01T14:43:30.904360Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def display_image_from_video(video_path):\n    '''\n    input: video_path - path for video\n    process:\n    1. perform a video capture from the video\n    2. read the image\n    3. display the image\n    '''\n    capture_image = cv2.VideoCapture(video_path) \n    ret, frame = capture_image.read()\n    fig = plt.figure(figsize=(10,10))\n    ax = fig.add_subplot(111)\n    frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n    ax.imshow(frame)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:30.906917Z","iopub.execute_input":"2024-11-01T14:43:30.907283Z","iopub.status.idle":"2024-11-01T14:43:30.917051Z","shell.execute_reply.started":"2024-11-01T14:43:30.907222Z","shell.execute_reply":"2024-11-01T14:43:30.915829Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for video_file in fake_train_sample_videos:\n    display_image_from_video(os.path.join(DATA_FOLDER, TRAIN_SAMPLE_FOLDER, video_file))","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:30.919224Z","iopub.execute_input":"2024-11-01T14:43:30.919837Z","iopub.status.idle":"2024-11-01T14:43:32.282353Z","shell.execute_reply.started":"2024-11-01T14:43:30.919540Z","shell.execute_reply":"2024-11-01T14:43:32.281661Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### **Display some REAL videos**","metadata":{}},{"cell_type":"code","source":"real_train_sample_videos = list(train_sample_metadata.loc[train_sample_metadata[\"label\"] == \"REAL\"].sample(3).index)\n\nreal_train_sample_videos","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:32.283606Z","iopub.execute_input":"2024-11-01T14:43:32.283845Z","iopub.status.idle":"2024-11-01T14:43:32.291505Z","shell.execute_reply.started":"2024-11-01T14:43:32.283805Z","shell.execute_reply":"2024-11-01T14:43:32.290820Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for video_file in real_train_sample_videos:\n    display_image_from_video(os.path.join(DATA_FOLDER, TRAIN_SAMPLE_FOLDER, video_file))","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:32.292878Z","iopub.execute_input":"2024-11-01T14:43:32.293180Z","iopub.status.idle":"2024-11-01T14:43:33.645157Z","shell.execute_reply.started":"2024-11-01T14:43:32.293112Z","shell.execute_reply":"2024-11-01T14:43:33.644183Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### **Videos with same originals**","metadata":{}},{"cell_type":"code","source":"train_sample_metadata[\"original\"].value_counts()[0:5]","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:33.646895Z","iopub.execute_input":"2024-11-01T14:43:33.647558Z","iopub.status.idle":"2024-11-01T14:43:33.658832Z","shell.execute_reply.started":"2024-11-01T14:43:33.647325Z","shell.execute_reply":"2024-11-01T14:43:33.657927Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Pick one of the original videos with largest number of duplicates and modify visualization function to work with multiple images.","metadata":{}},{"cell_type":"code","source":"def display_image_from_video_list(video_path_list, video_folder=TRAIN_SAMPLE_FOLDER):\n    '''\n    input: video_path_list - path for video\n    process:\n    0. for each video in the video path list\n        1. perform a video capture from the video\n        2. read the image\n        3. display the image\n    '''\n    plt.figure()\n    fig, ax = plt.subplots(2,3,figsize=(16,8))\n    # we only show images extracted from the first 6 videos\n    for i, video_file in enumerate(video_path_list[0:6]):\n        video_path = os.path.join(DATA_FOLDER, video_folder,video_file)\n        capture_image = cv2.VideoCapture(video_path) \n        ret, frame = capture_image.read()\n        frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n        ax[i//3, i%3].imshow(frame)\n        ax[i//3, i%3].set_title(f\"Video: {video_file}\")\n        ax[i//3, i%3].axis('on')","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:33.660684Z","iopub.execute_input":"2024-11-01T14:43:33.661191Z","iopub.status.idle":"2024-11-01T14:43:33.672725Z","shell.execute_reply.started":"2024-11-01T14:43:33.660984Z","shell.execute_reply":"2024-11-01T14:43:33.671985Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"same_original_fake_train_sample_video = list(train_sample_metadata.loc[train_sample_metadata.original=='atvmxvwyns.mp4'].index)\ndisplay_image_from_video_list(same_original_fake_train_sample_video)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:33.684713Z","iopub.execute_input":"2024-11-01T14:43:33.685208Z","iopub.status.idle":"2024-11-01T14:43:36.269470Z","shell.execute_reply.started":"2024-11-01T14:43:33.685155Z","shell.execute_reply":"2024-11-01T14:43:36.268697Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### **Test video files**","metadata":{}},{"cell_type":"code","source":"test_videos = pd.DataFrame(list(os.listdir(os.path.join(DATA_FOLDER, TEST_FOLDER))), columns=['video'])\n\ntest_videos.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.272006Z","iopub.execute_input":"2024-11-01T14:43:36.272298Z","iopub.status.idle":"2024-11-01T14:43:36.284185Z","shell.execute_reply.started":"2024-11-01T14:43:36.272246Z","shell.execute_reply":"2024-11-01T14:43:36.283484Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display_image_from_video(os.path.join(DATA_FOLDER, TEST_FOLDER, test_videos.iloc[1].video))","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.285521Z","iopub.execute_input":"2024-11-01T14:43:36.285742Z","iopub.status.idle":"2024-11-01T14:43:36.655887Z","shell.execute_reply.started":"2024-11-01T14:43:36.285704Z","shell.execute_reply":"2024-11-01T14:43:36.655112Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### **Play FAKE video file**","metadata":{}},{"cell_type":"code","source":"fake_videos = list(train_sample_metadata.loc[train_sample_metadata.label=='FAKE'].index)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.657329Z","iopub.execute_input":"2024-11-01T14:43:36.657601Z","iopub.status.idle":"2024-11-01T14:43:36.663328Z","shell.execute_reply.started":"2024-11-01T14:43:36.657550Z","shell.execute_reply":"2024-11-01T14:43:36.662415Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import HTML\nfrom base64 import b64encode\n\ndef play_video(video_file, subset=TRAIN_SAMPLE_FOLDER):\n    '''\n    Display video\n    param: video_file - the name of the video file to display\n    param: subset - the folder where the video file is located (can be TRAIN_SAMPLE_FOLDER or TEST_Folder)\n    '''\n    video_url = open(os.path.join(DATA_FOLDER, subset,video_file),'rb').read()\n    data_url = \"data:video/mp4;base64,\" + b64encode(video_url).decode()\n    return HTML(\"\"\"<video width=500 controls><source src=\"%s\" type=\"video/mp4\"></video>\"\"\" % data_url)\n\nplay_video(fake_videos[10])","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.664982Z","iopub.execute_input":"2024-11-01T14:43:36.665547Z","iopub.status.idle":"2024-11-01T14:43:36.833378Z","shell.execute_reply.started":"2024-11-01T14:43:36.665479Z","shell.execute_reply":"2024-11-01T14:43:36.832007Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"From visual inspection of these fakes videos, in some cases is very easy to spot the anomalies created when engineering the deep fake, in some cases is more difficult.","metadata":{}},{"cell_type":"markdown","source":"# **Modelling for Deepfake videos classification**","metadata":{}},{"cell_type":"markdown","source":"## **Build the dataset**","metadata":{}},{"cell_type":"code","source":"# IMG_SIZE = 224\n# BATCH_SIZE = 16\n# EPOCHS = 10\n\n# MAX_SEQ_LENGTH = 20\n# NUM_FEATURES = 2048","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.834702Z","iopub.execute_input":"2024-11-01T14:43:36.835122Z","iopub.status.idle":"2024-11-01T14:43:36.839385Z","shell.execute_reply.started":"2024-11-01T14:43:36.835077Z","shell.execute_reply":"2024-11-01T14:43:36.838172Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"In this example we will do the following:\n\n- Capture the frames of a video.\n\n- Extract frames from the videos until a maximum frame count is reached.\n\n- In the case, where a video's frame count is lesser than the maximum frame count we will pad the video with zeros.","metadata":{}},{"cell_type":"code","source":"# def crop_center_square(frame):\n#     y, x = frame.shape[0:2]\n#     min_dim = min(y, x)\n#     start_x = (x // 2) - (min_dim // 2)\n#     start_y = (y // 2) - (min_dim // 2)\n#     return frame[start_y : start_y + min_dim, start_x : start_x + min_dim]\n\n\n# def load_video(path, max_frames=0, resize=(IMG_SIZE, IMG_SIZE)):\n#     cap = cv2.VideoCapture(path)\n#     frames = []\n#     try:\n#         while True:\n#             ret, frame = cap.read()\n#             if not ret:\n#                 break\n#             frame = crop_center_square(frame)\n#             frame = cv2.resize(frame, resize)\n#             frame = frame[:, :, [2, 1, 0]]\n#             frames.append(frame)\n\n#             if len(frames) == max_frames:\n#                 break\n#     finally:\n#         cap.release()\n#     return np.array(frames)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.840755Z","iopub.execute_input":"2024-11-01T14:43:36.841090Z","iopub.status.idle":"2024-11-01T14:43:36.847960Z","shell.execute_reply.started":"2024-11-01T14:43:36.841040Z","shell.execute_reply":"2024-11-01T14:43:36.847146Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_all_videos(df, root_dir):\n    num_samples = len(df)\n    video_paths = list(df.index)\n    labels = df[\"label\"].values\n    labels = np.array(labels=='FAKE').astype(np.int)\n\n    # `frame_masks` and `frame_features` are what we will feed to our sequence model.\n    # `frame_masks` will contain a bunch of booleans denoting if a timestep is\n    # masked with padding or not.\n    frame_masks = np.zeros(shape=(num_samples, MAX_SEQ_LENGTH), dtype=\"bool\")\n    frame_features = np.zeros(\n        shape=(num_samples, MAX_SEQ_LENGTH, NUM_FEATURES), dtype=\"float32\"\n    )\n\n    # For each video.\n    for idx, path in enumerate(video_paths):\n        # Gather all its frames and add a batch dimension.\n        frames = load_video(os.path.join(root_dir, path))\n        frames = frames[None, ...]\n\n        # Initialize placeholders to store the masks and features of the current video.\n        temp_frame_mask = np.zeros(shape=(1, MAX_SEQ_LENGTH,), dtype=\"bool\")\n        temp_frame_features = np.zeros(\n            shape=(1, MAX_SEQ_LENGTH, NUM_FEATURES), dtype=\"float32\"\n        )\n\n        # Extract features from the frames of the current video.\n        for i, batch in enumerate(frames):\n            video_length = batch.shape[0]\n            length = min(MAX_SEQ_LENGTH, video_length)\n            for j in range(length):\n                temp_frame_features[i, j, :] = feature_extractor.predict(\n                    batch[None, j, :]\n                )\n            temp_frame_mask[i, :length] = 1  # 1 = not masked, 0 = masked\n\n        frame_features[idx,] = temp_frame_features.squeeze()\n        frame_masks[idx,] = temp_frame_mask.squeeze()\n\n    return (frame_features, frame_masks), labels","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.849069Z","iopub.execute_input":"2024-11-01T14:43:36.849420Z","iopub.status.idle":"2024-11-01T14:43:36.862490Z","shell.execute_reply.started":"2024-11-01T14:43:36.849375Z","shell.execute_reply":"2024-11-01T14:43:36.861850Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_set, test_set = train_test_split(train_sample_metadata,test_size=0.1,random_state=42,stratify=train_sample_metadata[\"label\"])\n\n# print(train_set.shape, test_set.shape)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.863654Z","iopub.execute_input":"2024-11-01T14:43:36.864047Z","iopub.status.idle":"2024-11-01T14:43:36.871176Z","shell.execute_reply.started":"2024-11-01T14:43:36.863909Z","shell.execute_reply":"2024-11-01T14:43:36.870329Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_data, train_labels = prepare_all_videos(train_set, \"train\")\n# test_data, test_labels = prepare_all_videos(test_set, \"test\")\n\n# print(f\"Frame features in train set: {train_data[0].shape}\")\n# print(f\"Frame masks in train set: {train_data[1].shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.872272Z","iopub.execute_input":"2024-11-01T14:43:36.872545Z","iopub.status.idle":"2024-11-01T14:43:36.877673Z","shell.execute_reply.started":"2024-11-01T14:43:36.872495Z","shell.execute_reply":"2024-11-01T14:43:36.877045Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metafiles = sorted(glob.glob('/kaggle/input/dfdc-video-faces/part*/*/*.json'))\ntrain_metafiles = metafiles[2:]\neval_metafiles = metafiles[:2]\nprint(train_metafiles)\nprint(eval_metafiles)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.878716Z","iopub.execute_input":"2024-11-01T14:43:36.878971Z","iopub.status.idle":"2024-11-01T14:43:36.910884Z","shell.execute_reply.started":"2024-11-01T14:43:36.878928Z","shell.execute_reply":"2024-11-01T14:43:36.910151Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VideoDataset(Dataset):\n\n    def __init__(self, dataset_type, meta_files, batch_size):\n        self.dataset_type = dataset_type\n        self.batch_size = batch_size\n        self.clip_shape = (16, 112, 112)\n        self.load_meta(meta_files)\n\n        print('%s dataset: %d real and %d fake samples' % (\n              dataset_type, self.n_reals, self.n_fakes))\n\n    def load_clip(self, filename, target):\n        depth, height, width = self.clip_shape\n        clip = np.zeros((3, depth, height, width), dtype=np.uint8)\n        reader = cv2.VideoCapture(filename)\n        if not reader.isOpened():\n            logging.warn('could not open %s' % filename)\n            return torch.from_numpy(clip).float()\n\n        # If training, use the same cropping parameters for an entire set\n        # of video clips\n        if target == 0 or self.dataset_type == 'eval':\n            nframes = int(reader.get(cv2.CAP_PROP_FRAME_COUNT))\n            frame_height = int(reader.get(cv2.CAP_PROP_FRAME_HEIGHT))\n            frame_width = int(reader.get(cv2.CAP_PROP_FRAME_WIDTH))\n            self.start_frame = random.randint(0, nframes - self.clip_shape[0])\n            self.start_row = random.randint(0, frame_height - height)\n            self.start_col = random.randint(0, frame_width - width)\n\n        for i in range(self.start_frame):\n            reader.grab()\n\n        for i in range(depth):\n            reader.grab()\n            success, frame = reader.retrieve()\n            if not success:\n                logging.warn('could not load frame %d in %s' % (\n                             start_frame + i, filename))\n                break\n            frame = frame[self.start_row:self.start_row + height,\n                          self.start_col:self.start_col + width]\n            clip[:, i] = frame.transpose((2, 0, 1))\n\n        reader.release()\n        return torch.from_numpy(clip).float()\n\n    def load_meta(self, meta_files):\n        meta = []\n        for meta_file in meta_files:\n            dirname = os.path.dirname(meta_file)\n            with open(meta_file) as meta_fd:\n                meta_dict = json.load(meta_fd)\n                new_dict = {}\n                # Expand filenames to their paths\n                for real in meta_dict:\n                    fakes = meta_dict[real]\n                    fakes = [os.path.join(dirname, fake) for fake in fakes]\n                    new_dict[os.path.join(dirname, real)] = fakes\n                meta_list = list(new_dict.items())\n                meta.extend(meta_list)\n\n        random.shuffle(meta)\n        self.clips = []\n        self.targets = []\n        for item in meta:\n            real = item[0]\n            fakes = item[1]\n            random.shuffle(fakes)\n            for fake in fakes:\n                # Oversample from non-fake videos\n                self.clips.append(real)\n                self.targets.append(0)\n                self.clips.append(fake)\n                self.targets.append(1)\n                # Use a small subset for evaluation\n                if self.dataset_type == 'eval':\n                    break\n\n        if self.dataset_type == 'eval':\n            # Make 4 copies to get random crops from\n            for _ in range(2):\n                self.clips.extend(self.clips)\n                self.targets.extend(self.targets)\n\n        self.len = len(self.clips)\n        self.n_fakes = np.sum(self.targets)\n        self.n_reals = self.len -  self.n_fakes\n\n    def __getitem__(self, index):\n        filename = self.clips[index]\n        target = self.targets[index]\n        clip = self.load_clip(filename, target)\n\n        return clip, target\n\n    def __len__(self):\n        return self.len","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.912118Z","iopub.execute_input":"2024-11-01T14:43:36.912369Z","iopub.status.idle":"2024-11-01T14:43:36.941225Z","shell.execute_reply.started":"2024-11-01T14:43:36.912321Z","shell.execute_reply":"2024-11-01T14:43:36.940293Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metafiles = sorted(glob.glob('/kaggle/input/dfdc-video-faces/part*/*/*.json'))\ntrain_metafiles = metafiles[2:]\neval_metafiles = metafiles[:2]\nprint(train_metafiles)\nprint(eval_metafiles)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T14:43:36.942617Z","iopub.execute_input":"2024-11-01T14:43:36.942889Z","iopub.status.idle":"2024-11-01T14:43:36.983674Z","shell.execute_reply.started":"2024-11-01T14:43:36.942840Z","shell.execute_reply":"2024-11-01T14:43:36.982747Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 5\nbatch_size = 16\n\ntrain_dataset = VideoDataset('train', train_metafiles, batch_size)\n\ntrain_loader = torch.utils.data.DataLoader(\n    train_dataset, batch_size=batch_size, shuffle=False,\n    num_workers=1, pin_memory=True, sampler=None)\n\neval_dataset = VideoDataset('eval', eval_metafiles, batch_size)\n\neval_loader = torch.utils.data.DataLoader(\n    eval_dataset, batch_size=batch_size, shuffle=False,\n    num_workers=1, pin_memory=True, sampler=None)","metadata":{"execution":{"iopub.status.busy":"2024-11-01T16:18:28.925567Z","iopub.execute_input":"2024-11-01T16:18:28.925866Z","iopub.status.idle":"2024-11-01T16:18:28.995918Z","shell.execute_reply.started":"2024-11-01T16:18:28.925823Z","shell.execute_reply":"2024-11-01T16:18:28.995073Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Use the model**","metadata":{}},{"cell_type":"code","source":"def train(loader, model, crit, optimizer, epoch):\n    model.train()\n\n    loss_sum = 0\n    for clips, targets in tqdm(loader):\n        clips = clips.to(device)\n        targets = targets.to(device)\n\n        logits = model(clips)\n        loss = crit(logits, targets)\n        loss_sum += loss.data.cpu().numpy()\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n    return (loss_sum / len(loader))\n\ndef bce(probs, labels):\n    safelog =  lambda x: np.log(np.maximum(x, np.exp(-50.)))\n    return np.mean(-labels * safelog(probs) - (1 - labels) * safelog(1 - probs))\n\ndef validate(loader, model, crit):\n    model.eval()\n    sm = nn.Softmax(dim=1)\n    labels = np.zeros((len(loader.dataset)), dtype=np.float32)\n    probs = np.zeros((len(loader.dataset), 2), dtype=np.float32)\n    with torch.no_grad():\n        for i, (clips, targets) in enumerate(tqdm(loader)):\n            start = i*batch_size\n            end = start + clips.shape[0]\n            labels[start:end] = targets\n            clips = clips.to(device)\n\n            logits = model(clips)\n            probs[start:end] = sm(logits).cpu().numpy()\n\n    probs = probs.reshape(4, -1, 2).mean(axis=0)\n    labels = labels.reshape(4, -1).mean(axis=0)\n\n    preds = probs.argmax(axis=1)\n    correct = (preds == labels).sum()\n    acc = correct*100//preds.shape[0]\n    loss = bce(probs[:, 1], labels)\n    print('validation accuracy %d%%' % acc)\n    return loss\n\n\nmodel = ResNet50LSTMClassifier().to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\ntrain_losses = []\nval_losses = []\n\nfor epoch in range(epochs):\n    # Train for one epoch\n    train_loss = train(train_loader, model, criterion, optimizer, epoch)\n\n    # Evaluate on validation set\n    val_loss = validate(eval_loader, model, criterion)\n    print('epoch %d training loss %.2f validation loss %.2f\\n' % (\n          epoch, train_loss, val_loss))\n\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n\nprint('done')","metadata":{"execution":{"iopub.status.busy":"2024-11-01T15:12:09.343272Z","iopub.execute_input":"2024-11-01T15:12:09.343621Z","iopub.status.idle":"2024-11-01T16:13:41.416180Z","shell.execute_reply.started":"2024-11-01T15:12:09.343562Z","shell.execute_reply":"2024-11-01T16:13:41.415034Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nnum_epochs = range(1, epochs + 1)\n\nplt.figure(figsize=(14, 5))\n\n# Plot Training and Validation Loss\nplt.plot(num_epochs, train_losses, label=\"Training Loss\")\nplt.plot(num_epochs, val_losses, label=\"Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training and Validation Loss\")\nplt.legend()\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-01T16:18:39.977412Z","iopub.execute_input":"2024-11-01T16:18:39.977701Z","iopub.status.idle":"2024-11-01T16:18:40.377256Z","shell.execute_reply.started":"2024-11-01T16:18:39.977659Z","shell.execute_reply":"2024-11-01T16:18:40.376119Z"}},"outputs":[],"execution_count":null}]}