{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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":5380830,"sourceType":"datasetVersion","datasetId":3120670},{"sourceId":6892693,"sourceType":"datasetVersion","datasetId":3959649},{"sourceId":7615428,"sourceType":"datasetVersion","datasetId":4434986}],"dockerImageVersionId":29845,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"_uuid":"b5ef8bdb-dcd5-4864-94f7-6e1b1ae4ca84","_cell_guid":"1537f2ec-a91e-4f20-a951-a729e9cdcbd3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms, models\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_score, recall_score, confusion_matrix, roc_auc_score, classification_report\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nfrom tqdm import tqdm\nimport random\nimport copy\nfrom torch.cuda.amp import autocast, GradScaler\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Set seeds for reproducibility\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(42)\n\n# Constants and Configuration\nCONFIG = {\n    'data_root': '/kaggle/input/faceforensics/FF++',\n    'batch_size': 4,\n    'num_frames': 20,\n    'frame_size': 224, # EfficientNet-B0 default input size\n    'num_epochs': 25,\n    'learning_rate': 1e-4, # May need adjustment for EfficientNet (try 5e-5 or 1e-5 if unstable)\n    'weight_decay': 1e-5, # May need adjustment\n    'lstm_hidden_dim': 512,\n    'dropout_rate': 0.5,\n    'use_attention': True,\n    'use_spatial_dropout': True, # Renamed, applied after feature extraction before LSTM\n    'use_temporal_dropout': True, # Dropout within LSTM layers\n    'use_focal_loss': True,\n    'focal_alpha': 0.25,\n    'focal_gamma': 2.0,\n    'scheduler_patience': 3,\n    'scheduler_factor': 0.5,\n    'early_stopping_patience': 5,\n    'model_save_path': './model_checkpoints_efficientnet', # Separate directory\n    'num_workers': 2, # Reduced further just in case of loader issues\n    'layers_to_freeze': 3, # Number of initial feature blocks in EfficientNet to freeze (0=stem, 1=stage1, etc.)\n}\n\n# Make sure model checkpoint directory exists\nos.makedirs(CONFIG['model_save_path'], exist_ok=True)\n\n# Helper Functions (Identical to previous versions)\ndef extract_frames(video_path, num_frames):\n    \"\"\"Extract evenly spaced frames from a video.\"\"\"\n    cap = cv2.VideoCapture(video_path)\n    if not cap.isOpened():\n        print(f\"Error: Could not open video {video_path}\")\n        # Return blank frames if video cannot be opened\n        return np.zeros((num_frames, CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n\n    total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))\n    if total_frames == 0:\n        print(f\"Warning: Video {video_path} has zero frames.\")\n        cap.release()\n        return np.zeros((num_frames, CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n\n    if total_frames <= num_frames:\n        indices = np.linspace(0, total_frames - 1, num_frames, endpoint=True, dtype=int)\n    else:\n        indices = np.linspace(0, total_frames - 1, num_frames, endpoint=True, dtype=int)\n\n    frames = []\n    processed_indices = set() # To handle cases where linspace gives duplicate indices for low total_frames\n\n    for idx in indices:\n        if idx in processed_indices: # If index already processed (due to duplication), try next available\n             original_idx = idx\n             while idx in processed_indices and idx < total_frames - 1:\n                 idx += 1\n             # If we exhausted options, just reuse the original (might happen if num_frames > total_frames)\n             if idx in processed_indices:\n                 idx = original_idx # Fallback to original index if no other unique frame is available\n\n        if idx >= total_frames: # Boundary check\n              idx = total_frames - 1\n\n        if idx not in processed_indices:\n             cap.set(cv2.CAP_PROP_POS_FRAMES, idx)\n             ret, frame = cap.read()\n             if ret:\n                 frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n                 frame = cv2.resize(frame, (CONFIG['frame_size'], CONFIG['frame_size']))\n                 frames.append(frame)\n                 processed_indices.add(idx)\n             else:\n                 # Attempt to read the previous frame if reading fails\n                 if idx > 0:\n                     cap.set(cv2.CAP_PROP_POS_FRAMES, idx - 1)\n                     ret, frame = cap.read()\n                     if ret:\n                         frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n                         frame = cv2.resize(frame, (CONFIG['frame_size'], CONFIG['frame_size']))\n                         frames.append(frame)\n                         processed_indices.add(idx) # Still mark idx as processed conceptually\n                     else:\n                         frames.append(np.zeros((CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8))\n                         processed_indices.add(idx)\n                 else: # If first frame fails, add zeros\n                    frames.append(np.zeros((CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8))\n                    processed_indices.add(idx)\n\n\n    cap.release()\n\n    # Pad if fewer frames were extracted than needed\n    while len(frames) < num_frames:\n        if frames: # Pad with the last valid frame if available\n            frames.append(frames[-1].copy())\n        else: # Pad with zeros if no frames were read at all\n            frames.append(np.zeros((CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8))\n\n    # Ensure all frames have the correct shape just in case\n    final_frames = []\n    for i in range(len(frames)):\n        frame = frames[i]\n        if frame.shape != (CONFIG['frame_size'], CONFIG['frame_size'], 3):\n             try:\n                 frame = cv2.resize(frame, (CONFIG['frame_size'], CONFIG['frame_size']))\n             except Exception as resize_e:\n                 print(f\"Error resizing frame {i} from video {video_path}. Shape was {frame.shape}. Error: {resize_e}\")\n                 frame = np.zeros((CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8) # Use blank frame on error\n        final_frames.append(frame)\n\n\n    return np.array(final_frames[:num_frames]) # Ensure exactly num_frames are returned\n\n\ndef get_class_weights(labels):\n    \"\"\"Calculate class weights inversely proportional to class frequencies.\"\"\"\n    class_counts = np.bincount(labels)\n    # Handle potential zero counts if a class is missing in the training split (unlikely with stratify)\n    if len(class_counts) < 2: # Assuming binary classification\n        return torch.FloatTensor([1.0, 1.0])\n    if class_counts[0] == 0 or class_counts[1] == 0:\n        print(\"Warning: One class has zero samples in the training set.\")\n        # Assign equal weight or handle as appropriate\n        return torch.FloatTensor([1.0, 1.0])\n\n    total_samples = len(labels)\n    class_weights = total_samples / (len(class_counts) * class_counts)\n    return torch.FloatTensor(class_weights)\n\ndef plot_confusion_matrix(cm, epoch=None, save_path=CONFIG['model_save_path']):\n    \"\"\"Plot confusion matrix.\"\"\"\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', cbar=False, xticklabels=['Predicted Fake', 'Predicted Real'], yticklabels=['True Fake', 'True Real'])\n    plt.xlabel('Predicted Labels')\n    plt.ylabel('True Labels')\n    title = f'Confusion Matrix'\n    if epoch is not None:\n        title += f' - Epoch {epoch+1}'\n    plt.title(title)\n    if epoch is not None:\n        plt.savefig(os.path.join(save_path, f'confusion_matrix_epoch_{epoch+1}.png'))\n    plt.close()\n\ndef plot_metrics(metrics, save_path=CONFIG['model_save_path']):\n    \"\"\"Plot training metrics.\"\"\"\n    epochs = range(1, len(metrics['train_loss']) + 1)\n\n    plt.figure(figsize=(15, 12))\n\n    # Plot loss\n    plt.subplot(2, 2, 1)\n    plt.plot(epochs, metrics['train_loss'], 'b-', label='Training Loss')\n    plt.plot(epochs, metrics['val_loss'], 'r-', label='Validation Loss')\n    plt.title('Training and Validation Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.ylim(0, 0.8)  # Loss between 0 and 0.8\n    plt.legend()\n    plt.grid(True)\n\n    # Plot accuracy\n    plt.subplot(2, 2, 2)\n    plt.plot(epochs, metrics['train_acc'], 'b-', label='Training Accuracy')\n    plt.plot(epochs, metrics['val_acc'], 'r-', label='Validation Accuracy')\n    plt.title('Training and Validation Accuracy')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n    plt.ylim(0, 1)  # Accuracy between 0 and 1\n    plt.legend()\n    plt.grid(True)\n\n    # Plot AUC-ROC\n    plt.subplot(2, 2, 3)\n    plt.plot(epochs, metrics['val_auc'], 'g-')\n    plt.title('Validation AUC-ROC')\n    plt.xlabel('Epochs')\n    plt.ylabel('AUC-ROC')\n    plt.grid(True)\n\n    # Plot Precision and Recall\n    plt.subplot(2, 2, 4)\n    plt.plot(epochs, metrics['val_precision'], 'c-', label='Precision (Val)')\n    plt.plot(epochs, metrics['val_recall'], 'm-', label='Recall (Val)')\n    plt.title('Validation Precision and Recall')\n    plt.xlabel('Epochs')\n    plt.ylabel('Score')\n    plt.legend()\n    plt.grid(True)\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(save_path, 'training_metrics.png'))\n    plt.close()\n\n\n# Custom Attention Module (Identical)\nclass TemporalAttention(nn.Module):\n    def __init__(self, hidden_dim):\n        super(TemporalAttention, self).__init__()\n        self.attention = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.ReLU(),\n            nn.Linear(hidden_dim // 2, 1)\n        )\n\n    def forward(self, lstm_output):\n        # lstm_output shape: (batch, seq_len, hidden_dim)\n        attn_weights = F.softmax(self.attention(lstm_output), dim=1)\n        # Apply attention weights\n        context = torch.sum(attn_weights * lstm_output, dim=1)\n        return context, attn_weights\n\n# Video Dataset (Added error handling for file existence)\nclass DeepfakeDataset(Dataset):\n    def __init__(self, video_paths, labels, transform=None):\n        self.video_paths = video_paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.video_paths)\n\n    def __getitem__(self, idx):\n        video_path = self.video_paths[idx]\n        label = self.labels[idx]\n\n        if not os.path.exists(video_path):\n            print(f\"Warning: Video file not found: {video_path}. Returning zeros.\")\n            frames = np.zeros((CONFIG['num_frames'], CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n        else:\n            try:\n                frames = extract_frames(video_path, CONFIG['num_frames'])\n                if frames.shape[0] != CONFIG['num_frames']:\n                     print(f\"Warning: extract_frames returned {frames.shape[0]} frames for {video_path}, expected {CONFIG['num_frames']}. Check padding.\")\n                     # Re-ensure padding (should be handled in extract_frames, but as safeguard)\n                     if frames.shape[0] < CONFIG['num_frames']:\n                         padding = np.zeros((CONFIG['num_frames'] - frames.shape[0], CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n                         if frames.shape[0] > 0: # Pad with last frame if possible\n                              padding = np.repeat(frames[-1:], CONFIG['num_frames'] - frames.shape[0], axis=0)\n                         frames = np.concatenate((frames, padding), axis=0)\n                     else: # Truncate if too many (shouldn't happen with linspace)\n                         frames = frames[:CONFIG['num_frames']]\n\n\n            except Exception as e:\n                print(f\"Error processing video {idx}: {video_path} during frame extraction.\")\n                print(f\"Error details: {str(e)}\")\n                frames = np.zeros((CONFIG['num_frames'], CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n\n        # Apply transformations to each frame\n        transformed_frames = []\n        if self.transform:\n            try:\n                for frame_idx, frame in enumerate(frames):\n                     # Ensure frame is uint8 for ToPILImage\n                     if not isinstance(frame, np.ndarray):\n                          print(f\"Warning: Frame {frame_idx} for video {video_path} is not a numpy array (type: {type(frame)}). Using zero frame.\")\n                          frame = np.zeros((CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n                     elif frame.dtype != np.uint8:\n                         # print(f\"Warning: Frame {frame_idx} for video {video_path} has dtype {frame.dtype}. Converting to uint8.\")\n                         frame = frame.astype(np.uint8)\n\n                     # Basic check for valid image dimensions\n                     if frame.shape != (CONFIG['frame_size'], CONFIG['frame_size'], 3):\n                         print(f\"Warning: Frame {frame_idx} for video {video_path} has unexpected shape {frame.shape}. Attempting resize or using zero frame.\")\n                         try:\n                             frame = cv2.resize(frame, (CONFIG['frame_size'], CONFIG['frame_size']))\n                             if frame.shape != (CONFIG['frame_size'], CONFIG['frame_size'], 3): # Check resize result\n                                 raise ValueError(\"Resize did not produce expected shape\")\n                         except Exception as resize_e:\n                             print(f\"  Resize failed: {resize_e}. Using zero frame.\")\n                             frame = np.zeros((CONFIG['frame_size'], CONFIG['frame_size'], 3), dtype=np.uint8)\n\n\n                     transformed_frames.append(self.transform(frame))\n\n                frames_tensor = torch.stack(transformed_frames)\n            except Exception as e_transform:\n                 print(f\"Error applying transform to frames from video {idx}: {video_path}\")\n                 print(f\"Error details: {str(e_transform)}\")\n                 # Return a dummy tensor in case of transformation errors\n                 frames_tensor = torch.zeros((CONFIG['num_frames'], 3, CONFIG['frame_size'], CONFIG['frame_size']))\n\n        else: # If no transform provided\n            try:\n                # Convert numpy array to tensor: (N, H, W, C) -> (N, C, H, W)\n                frames_tensor = torch.from_numpy(frames).permute(0, 3, 1, 2).float() / 255.0\n            except Exception as e_notransform:\n                print(f\"Error converting numpy frames to tensor for video {idx}: {video_path}\")\n                print(f\"Error details: {str(e_notransform)}\")\n                frames_tensor = torch.zeros((CONFIG['num_frames'], 3, CONFIG['frame_size'], CONFIG['frame_size']))\n\n\n        # Final check for tensor shape consistency\n        if frames_tensor.shape != (CONFIG['num_frames'], 3, CONFIG['frame_size'], CONFIG['frame_size']):\n            print(f\"Warning: Final tensor shape for video {idx} ({video_path}) is {frames_tensor.shape}, expected {(CONFIG['num_frames'], 3, CONFIG['frame_size'], CONFIG['frame_size'])}. Returning zeros.\")\n            frames_tensor = torch.zeros((CONFIG['num_frames'], 3, CONFIG['frame_size'], CONFIG['frame_size']))\n\n        return frames_tensor, torch.tensor(label, dtype=torch.float) # Ensure label is float for BCE/Focal loss\n\n\n# Focal Loss (Identical)\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(self, inputs, targets):\n        # Ensure targets are same type and shape as inputs\n        targets = targets.type_as(inputs).view_as(inputs)\n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-BCE_loss) # calculates p_t = sigmoid(input) for positive targets and 1 - sigmoid(input) for negative targets\n        alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets) # alpha_t\n        focal_loss = alpha_t * (1 - pt)**self.gamma * BCE_loss\n        return focal_loss.mean()\n\n# DeepFake Detection Model (Modified for EfficientNet Freezing)\nclass DeepfakeDetector(nn.Module):\n    def __init__(self, config):\n        super(DeepfakeDetector, self).__init__()\n        self.config = config\n\n        # Load EfficientNet-B0 pre-trained on ImageNet\n        efficientnet = models.efficientnet_b0(weights='DEFAULT')\n\n        # --- MODIFIED FREEZING STRATEGY for EfficientNet ---\n        # EfficientNet structure: features (Sequential), avgpool, classifier (Linear)\n        # We will freeze initial blocks within 'features'\n        feature_blocks = efficientnet.features\n        layers_to_freeze = config.get('layers_to_freeze', 3) # Default to 3 if not in config\n        print(f\"Attempting to freeze the first {layers_to_freeze} feature blocks of EfficientNet...\")\n\n        ct = 0\n        # Iterate through the direct children of the 'features' Sequential module\n        for idx, child in enumerate(feature_blocks.children()):\n            if ct < layers_to_freeze:\n                print(f\"  Freezing EfficientNet feature block {idx}: {child.__class__.__name__}\")\n                for param in child.parameters():\n                    param.requires_grad = False\n            else:\n                print(f\"  NOT Freezing EfficientNet feature block {idx}: {child.__class__.__name__}\")\n                for param in child.parameters():\n                    param.requires_grad = True # Ensure later layers are trainable\n            ct += 1\n        print(f\"Finished setting requires_grad for EfficientNet feature blocks.\")\n        # --- END MODIFIED FREEZING ---\n\n        # Use the modified features and the original avgpool, discard original classifier\n        self.feature_extractor = nn.Sequential(\n            efficientnet.features,\n            efficientnet.avgpool\n        )\n\n        # Feature dimension from EfficientNet-B0 after avgpool\n        self.feature_dim = 1280\n\n        # Check feature dimension (optional sanity check)\n        # with torch.no_grad():\n        #     dummy_input = torch.randn(1, 3, config['frame_size'], config['frame_size'])\n        #     dummy_output = self.feature_extractor(dummy_input)\n        #     print(f\"Sanity Check: Output shape after feature extractor: {dummy_output.shape}\") # Should be [1, 1280, 1, 1]\n        #     actual_dim = dummy_output.view(dummy_output.size(0), -1).shape[1]\n        #     if actual_dim != self.feature_dim:\n        #          print(f\"Warning: Expected feature_dim {self.feature_dim}, but got {actual_dim}. Adjusting.\")\n        #          self.feature_dim = actual_dim\n\n\n        # Spatial dropout (applied after feature extraction, before LSTM)\n        # Note: Dropout2d usually applied on feature maps (CxHxW). Here applying on flattened features (effectively 1D dropout)\n        self.spatial_dropout = nn.Dropout(p=config['dropout_rate']) if config['use_spatial_dropout'] else nn.Identity()\n\n        # Bidirectional LSTM for temporal analysis\n        self.lstm = nn.LSTM(\n            input_size=self.feature_dim,\n            hidden_size=config['lstm_hidden_dim'],\n            num_layers=2, # Using 2 LSTM layers\n            batch_first=True,\n            bidirectional=True,\n            dropout=config['dropout_rate'] if config['use_temporal_dropout'] and 2 > 1 else 0 # LSTM dropout only applies between layers if num_layers > 1\n        )\n\n        # Attention mechanism\n        self.use_attention = config['use_attention']\n        lstm_output_dim = config['lstm_hidden_dim'] * 2 # *2 for bidirectional\n\n        if self.use_attention:\n            self.attention = TemporalAttention(lstm_output_dim)\n            # Final classifier using attended features\n            self.classifier = nn.Sequential(\n                nn.Linear(lstm_output_dim, config['lstm_hidden_dim']), # Input from attention context vector\n                nn.ReLU(),\n                nn.Dropout(config['dropout_rate']),\n                nn.Linear(config['lstm_hidden_dim'], 1)\n            )\n        else:\n            # Final classifier using the last LSTM hidden state (or average/max pooling)\n            # Using last hidden state here for simplicity if no attention\n             self.classifier = nn.Sequential(\n                nn.Linear(lstm_output_dim, config['lstm_hidden_dim']), # Input is last hidden state\n                nn.ReLU(),\n                nn.Dropout(config['dropout_rate']),\n                nn.Linear(config['lstm_hidden_dim'], 1)\n            )\n\n    def forward(self, x):\n        # x shape: (batch_size, seq_len, C, H, W)\n        batch_size, seq_len, c, h, w = x.size()\n\n        # Process each frame through CNN\n        # Reshape input for CNN: (batch_size * seq_len, C, H, W)\n        x_cnn = x.view(batch_size * seq_len, c, h, w)\n        cnn_out = self.feature_extractor(x_cnn) # Output shape: (batch*seq_len, feature_dim, 1, 1)\n\n        # Flatten the features\n        cnn_features = cnn_out.view(batch_size, seq_len, -1) # (batch, seq_len, feature_dim)\n\n        # Apply spatial dropout to the sequence of features\n        if self.config['use_spatial_dropout']:\n             cnn_features = self.spatial_dropout(cnn_features) # Apply dropout across the feature dimension for each time step\n\n        # Process sequence with LSTM\n        lstm_out, _ = self.lstm(cnn_features)  # lstm_out shape: (batch, seq_len, hidden_dim*2)\n\n        # Apply attention or use final state\n        if self.use_attention:\n            context, attn_weights = self.attention(lstm_out) # context shape: (batch, hidden_dim*2)\n            # We could potentially return attn_weights for visualization/analysis\n        else:\n            # Use the output of the last time step from the BiLSTM\n            context = lstm_out[:, -1, :] # context shape: (batch, hidden_dim*2)\n\n        # Final classification\n        output = self.classifier(context) # output shape: (batch, 1)\n        return output.squeeze(-1) # Squeeze the last dimension -> (batch)\n\n\n# Training function (Identical, but added save paths to plots)\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, device, num_epochs, config):\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_val_loss = float('inf')\n    best_val_acc = 0.0\n    early_stopping_counter = 0\n    save_path = config['model_save_path'] # Get save path from config\n\n    # Create scaler for mixed precision training\n    scaler = GradScaler()\n\n    # Store metrics\n    metrics = {\n        'train_loss': [], 'train_acc': [],\n        'val_loss': [], 'val_acc': [],\n        'val_precision': [], 'val_recall': [],\n        'val_auc': []\n    }\n\n    # Create tqdm progress bar for epochs\n    epoch_loop = tqdm(range(num_epochs), desc=\"Training Progress\", unit=\"epoch\")\n\n    for epoch in epoch_loop:\n        epoch_loop.set_description(f\"Epoch {epoch+1}/{num_epochs}\")\n\n        # Training phase\n        model.train()\n        running_loss = 0.0\n        correct_train = 0\n        total_train = 0\n\n        train_loop = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\", leave=False, unit=\"batch\")\n        for inputs, labels in train_loop:\n            inputs = inputs.to(device, non_blocking=True) # Use non_blocking for potential speedup\n            labels = labels.to(device, non_blocking=True)\n\n            # Zero the parameter gradients\n            optimizer.zero_grad(set_to_none=True) # More efficient potentially\n\n            # Forward pass with mixed precision\n            with autocast():\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n            # Backward and optimize with gradient scaling\n            scaler.scale(loss).backward()\n            # Optional: Gradient clipping (can help stabilize training)\n            # scaler.unscale_(optimizer) # Unscale gradients before clipping\n            # torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n\n            # Statistics\n            running_loss += loss.item() * inputs.size(0)\n            preds = torch.sigmoid(outputs.detach()) >= 0.5 # Use detach() for metrics calculation\n            correct_train += (preds == labels).sum().item()\n            total_train += labels.size(0)\n\n            # Update progress bar postfix\n            train_loop.set_postfix(loss=loss.item(), acc=correct_train/total_train if total_train > 0 else 0)\n\n        epoch_train_loss = running_loss / total_train if total_train > 0 else 0\n        epoch_train_acc = correct_train / total_train if total_train > 0 else 0\n        metrics['train_loss'].append(epoch_train_loss)\n        metrics['train_acc'].append(epoch_train_acc)\n\n        # Validation phase\n        model.eval()\n        running_loss_val = 0.0\n        correct_val = 0\n        total_val = 0\n        all_preds = []\n        all_labels = []\n        all_outputs = [] # For AUC\n\n        val_loop = tqdm(val_loader, desc=f\"Epoch {epoch+1} Validation\", leave=False, unit=\"batch\")\n        with torch.no_grad():\n            for inputs, labels in val_loop:\n                inputs = inputs.to(device, non_blocking=True)\n                labels = labels.to(device, non_blocking=True)\n\n                # Forward pass (no autocast needed for eval unless specific layers require it)\n                outputs = model(inputs)\n                loss = criterion(outputs, labels) # Calculate loss for monitoring\n\n                # Statistics\n                running_loss_val += loss.item() * inputs.size(0)\n                probs = torch.sigmoid(outputs)\n                preds = probs >= 0.5\n                correct_val += (preds == labels).sum().item()\n                total_val += labels.size(0)\n\n                # Store predictions and labels for metrics\n                all_preds.extend(preds.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n                all_outputs.extend(probs.cpu().numpy()) # Store probabilities for AUC\n\n                # Update progress bar postfix\n                val_loop.set_postfix(loss=loss.item(), acc=correct_val/total_val if total_val > 0 else 0)\n\n        epoch_val_loss = running_loss_val / total_val if total_val > 0 else 0\n        epoch_val_acc = correct_val / total_val if total_val > 0 else 0\n\n        # Calculate additional metrics (handle potential division by zero)\n        try:\n            val_precision = precision_score(all_labels, all_preds, zero_division=0)\n        except ValueError: # Handle cases where labels/preds might be all one class temporarily\n             val_precision = 0.0\n        try:\n            val_recall = recall_score(all_labels, all_preds, zero_division=0)\n        except ValueError:\n            val_recall = 0.0\n        try:\n            # Ensure there are samples of both classes for AUC calculation\n            if len(np.unique(all_labels)) > 1:\n                 val_auc = roc_auc_score(all_labels, all_outputs)\n            else:\n                 print(f\"Warning: Only one class present in validation labels for epoch {epoch+1}. AUC set to 0.5.\")\n                 val_auc = 0.5 # Or handle as undefined/skip\n        except ValueError as e_auc:\n             print(f\"Could not calculate AUC for epoch {epoch+1}: {e_auc}. Setting to 0.5.\")\n             val_auc = 0.5\n\n\n        # Generate confusion matrix for this epoch\n        cm = confusion_matrix(all_labels, all_preds)\n        plot_confusion_matrix(cm, epoch, save_path=save_path) # Pass save_path\n\n        # Print epoch results\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(f\"  Train Loss: {epoch_train_loss:.4f} | Train Acc: {epoch_train_acc:.4f}\")\n        print(f\"  Val Loss:   {epoch_val_loss:.4f} | Val Acc:   {epoch_val_acc:.4f}\")\n        print(f\"  Precision:  {val_precision:.4f} | Recall:    {val_recall:.4f} | AUC: {val_auc:.4f}\")\n        print(f\"  Confusion Matrix:\\n{cm}\")\n        print(\"-\" * 60)\n\n        # Update learning rate based on validation loss\n        scheduler.step(epoch_val_loss)\n\n        # Save metrics\n        metrics['val_loss'].append(epoch_val_loss)\n        metrics['val_acc'].append(epoch_val_acc)\n        metrics['val_precision'].append(val_precision)\n        metrics['val_recall'].append(val_recall)\n        metrics['val_auc'].append(val_auc)\n\n        # Save model if it's the best so far based on validation accuracy\n        if epoch_val_acc > best_val_acc:\n            print(f\"Validation accuracy improved from {best_val_acc:.4f} to {epoch_val_acc:.4f}. Saving model...\")\n            best_val_acc = epoch_val_acc\n            best_val_loss = epoch_val_loss # Also save best loss associated with best accuracy\n            best_model_wts = copy.deepcopy(model.state_dict())\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': best_model_wts,\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'val_acc': best_val_acc,\n                'val_loss': best_val_loss,\n                'val_precision': val_precision,\n                'val_recall': val_recall,\n                'val_auc': val_auc,\n                'config': config # Save config with the model\n            }, os.path.join(save_path, \"best_model.pth\"))\n            early_stopping_counter = 0\n        else:\n            early_stopping_counter += 1\n            print(f\"Validation accuracy did not improve. Counter: {early_stopping_counter}/{config['early_stopping_patience']}\")\n\n\n        # Save checkpoint periodically (e.g., every 5 epochs)\n        if (epoch + 1) % 5 == 0:\n             print(f\"Saving checkpoint at epoch {epoch+1}...\")\n             torch.save({\n                 'epoch': epoch,\n                 'model_state_dict': model.state_dict(),\n                 'optimizer_state_dict': optimizer.state_dict(),\n                 'scheduler_state_dict': scheduler.state_dict(),\n                 'val_acc': epoch_val_acc,\n                 'val_loss': epoch_val_loss,\n                 'config': config\n             }, os.path.join(save_path, f\"checkpoint_epoch_{epoch+1}.pth\"))\n\n        # Early stopping\n        if early_stopping_counter >= config['early_stopping_patience']:\n            print(f\"Early stopping triggered after {epoch+1} epochs.\")\n            break\n\n    # Plot overall training metrics at the end\n    plot_metrics(metrics, save_path=save_path) # Pass save_path\n\n    print(f\"Best Validation Accuracy: {best_val_acc:.4f}\")\n    # Load best model weights back into the model\n    model.load_state_dict(best_model_wts)\n    return model, metrics\n\n\n# Load and process data (Identical, but added check for empty lists)\ndef prepare_dataset(config):\n    \"\"\"Process the dataset and create data loaders.\"\"\"\n    print(\"Preparing dataset...\")\n    data_root = config['data_root']\n\n    # Parse the list file (assuming Celeb-DF v2 structure)\n    # Training data comes from Celeb-real, YouTube-real, Celeb-synthesis\n    # Test data is defined in List_of_testing_videos.txt (we exclude these from training/validation)\n    test_list_path = os.path.join(data_root, 'List_of_testing_videos.txt')\n    video_paths = []\n    labels = []\n\n    # Read the test list to exclude these videos\n    test_videos_relative_paths = set()\n    if os.path.exists(test_list_path):\n        try:\n            with open(test_list_path, 'r') as f:\n                # Skip header line if present (check first line format)\n                first_line = f.readline().strip()\n                # Simple check if it looks like a header vs data\n                if not (len(first_line.split()) >= 2 and first_line.split()[0].isdigit()):\n                     print(\"Skipping potential header line in test list.\")\n                else:\n                     # Process the first line if it's data\n                     parts = first_line.split()\n                     if len(parts) >= 2:\n                         # Label seems to be 1 for fake, 0 for real in test list? Double check dataset spec.\n                         # Assuming format: label relative_path (e.g., 1 Celeb-synthesis/id0_id1_0000.mp4)\n                         test_videos_relative_paths.add(parts[1])\n\n\n                # Process rest of the file\n                for line in f:\n                    parts = line.strip().split()\n                    if len(parts) >= 2:\n                        test_videos_relative_paths.add(parts[1])\n            print(f\"Found {len(test_videos_relative_paths)} test videos to exclude from training/validation.\")\n        except Exception as e:\n             print(f\"Error reading test list file {test_list_path}: {e}. Proceeding without exclusions.\")\n             test_videos_relative_paths = set()\n\n    else:\n        print(f\"Warning: Test list file not found at {test_list_path}. Cannot exclude test videos.\")\n\n\n    # Collect all video paths from training-related directories\n    real_dirs = ['real']\n    fake_dirs = ['fake']\n    all_video_files = [] # Store tuples of (full_path, label)\n\n    # Process real videos\n    for dir_name in real_dirs:\n        dir_path = os.path.join(data_root, dir_name)\n        if not os.path.isdir(dir_path):\n            print(f\"Warning: Directory {dir_path} not found. Skipping...\")\n            continue\n        print(f\"Processing directory: {dir_path}\")\n        count = 0\n        for file_name in os.listdir(dir_path):\n            if file_name.lower().endswith('.mp4'):\n                full_path = os.path.join(dir_path, file_name)\n                relative_path = os.path.join(dir_name, file_name) # Path relative to data_root used in test list\n\n                # Skip if this video is in the test set\n                if relative_path in test_videos_relative_paths:\n                    # print(f\"Skipping test video: {relative_path}\")\n                    continue\n\n                all_video_files.append((full_path, 1)) # Real = 1\n                count += 1\n        print(f\"  Found {count} real videos (excluding test set).\")\n\n\n    # Process fake videos\n    for dir_name in fake_dirs:\n        dir_path = os.path.join(data_root, dir_name)\n        if not os.path.isdir(dir_path):\n            print(f\"Warning: Directory {dir_path} not found. Skipping...\")\n            continue\n        print(f\"Processing directory: {dir_path}\")\n        count = 0\n        for file_name in os.listdir(dir_path):\n            if file_name.lower().endswith('.mp4'):\n                full_path = os.path.join(dir_path, file_name)\n                relative_path = os.path.join(dir_name, file_name)\n\n                # Skip if this video is in the test set\n                if relative_path in test_videos_relative_paths:\n                    # print(f\"Skipping test video: {relative_path}\")\n                    continue\n\n                all_video_files.append((full_path, 0)) # Fake = 0\n                count += 1\n        print(f\"  Found {count} fake videos (excluding test set).\")\n\n    if not all_video_files:\n         raise ValueError(f\"No video files found in specified directories ({real_dirs}, {fake_dirs}) under {data_root} or all were excluded as test videos.\")\n\n    # Separate paths and labels\n    video_paths = [item[0] for item in all_video_files]\n    labels = np.array([item[1] for item in all_video_files])\n\n    # Print class distribution\n    unique, counts = np.unique(labels, return_counts=True)\n    if len(unique) > 0:\n        class_distribution = dict(zip(unique, counts))\n        print(f\"Class distribution before split (0=Fake, 1=Real): {class_distribution}\")\n        # Check for imbalance\n        if 0 not in class_distribution or 1 not in class_distribution:\n             print(\"Warning: Only one class found in the collected data.\")\n        elif counts[0] == 0 or counts[1] == 0:\n             print(\"Warning: One class has zero samples.\")\n\n    else:\n         raise ValueError(\"No labels collected. Cannot proceed.\")\n\n\n    # Split data into train and validation sets\n    try:\n         train_paths, val_paths, train_labels, val_labels = train_test_split(\n             video_paths, labels,\n             test_size=0.2, # 20% for validation\n             random_state=42,\n             stratify=labels # Important for imbalanced datasets\n         )\n    except ValueError as e_split:\n         if \"n_splits=2 cannot be greater than the number of members in each class\" in str(e_split):\n              print(\"Error during train/test split: Not enough samples in one class for stratification. Check data collection and class distribution.\")\n              raise\n         else:\n              print(f\"An unexpected error occurred during train/test split: {e_split}\")\n              raise\n\n\n    print(f\"Train set size: {len(train_paths)}\")\n    print(f\"Validation set size: {len(val_paths)}\")\n    unique_train, counts_train = np.unique(train_labels, return_counts=True)\n    print(f\"Train class distribution (0=Fake, 1=Real): {dict(zip(unique_train, counts_train))}\")\n    unique_val, counts_val = np.unique(val_labels, return_counts=True)\n    print(f\"Validation class distribution (0=Fake, 1=Real): {dict(zip(unique_val, counts_val))}\")\n\n\n    # Create transformations (Standard for ImageNet pre-trained models)\n    # Normalization values for models pre-trained on ImageNet\n    normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n\n    train_transform = transforms.Compose([\n        transforms.ToPILImage(), # Convert numpy array (H, W, C) to PIL Image\n        transforms.Resize((config['frame_size'], config['frame_size'])),\n        transforms.RandomHorizontalFlip(p=0.5), # Data augmentation\n        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05), # Data augmentation\n        # transforms.RandomAffine(degrees=10, translate=(0.05, 0.05), scale=(0.95, 1.05)), # More augmentation (optional)\n        transforms.ToTensor(), # Convert PIL Image to tensor (C, H, W) and scales pixels to [0, 1]\n        normalize # Normalize tensor values\n    ])\n\n    val_transform = transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.Resize((config['frame_size'], config['frame_size'])),\n        transforms.ToTensor(),\n        normalize\n    ])\n\n    # Create datasets\n    train_dataset = DeepfakeDataset(train_paths, train_labels, transform=train_transform)\n    val_dataset = DeepfakeDataset(val_paths, val_labels, transform=val_transform)\n\n    # Calculate class weights for imbalanced data using ONLY training labels\n    class_weights = get_class_weights(train_labels)\n    print(f\"Calculated class weights for WeightedRandomSampler (Fake=0, Real=1): {class_weights.tolist()}\")\n\n    # Create weighted sampler for training data to handle imbalance\n    # Weight for each sample is the inverse frequency of its class\n    sample_weights = torch.tensor([class_weights[label] for label in train_labels])\n    sampler = WeightedRandomSampler(\n        weights=sample_weights,\n        num_samples=len(sample_weights), # Draw as many samples as the original dataset size\n        replacement=True # Allow drawing same sample multiple times within an epoch\n    )\n\n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config['batch_size'],\n        sampler=sampler, # Use the weighted sampler for training\n        num_workers=config['num_workers'],\n        pin_memory=True, # Speeds up data transfer to GPU\n        prefetch_factor=2 if config['num_workers'] > 0 else None, # How many batches to prefetch\n        persistent_workers=True if config['num_workers'] > 0 else False, # Keep workers alive between epochs\n        # collate_fn=None, # Default collate_fn is usually fine\n    )\n\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config['batch_size'], # Can often use a larger batch size for validation if memory allows\n        shuffle=False, # No need to shuffle validation data\n        num_workers=config['num_workers'],\n        pin_memory=True,\n        prefetch_factor=2 if config['num_workers'] > 0 else None,\n        persistent_workers=True if config['num_workers'] > 0 else False,\n    )\n\n    print(f\"Data loaders created.\")\n    print(f\"Training samples: {len(train_dataset)}, Validation samples: {len(val_dataset)}\")\n\n    return train_loader, val_loader, class_weights # Returning class_weights maybe useful later if needed for loss\n\n\ndef main():\n    # Check for GPU\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    if torch.cuda.is_available():\n        print(f\"GPU Name: {torch.cuda.get_device_name(0)}\")\n\n    try:\n        # Create model save directory (using path from config)\n        os.makedirs(CONFIG['model_save_path'], exist_ok=True)\n        print(f\"Model checkpoints will be saved to: {CONFIG['model_save_path']}\")\n\n        # Prepare data\n        try:\n            train_loader, val_loader, class_weights = prepare_dataset(CONFIG)\n        except Exception as e_data:\n            print(f\"FATAL: Error preparing dataset: {str(e_data)}\")\n            import traceback\n            traceback.print_exc()\n            # Optionally print directory structure for debugging data path issues\n            print(\"\\n--- Directory Structure ---\")\n            try:\n                for root, dirs, files in os.walk(CONFIG['data_root'], topdown=True):\n                     level = root.replace(CONFIG['data_root'], '').count(os.sep)\n                     indent = ' ' * 4 * (level)\n                     print(f'{indent}{os.path.basename(root)}/')\n                     subindent = ' ' * 4 * (level + 1)\n                     # Limit files shown per directory\n                     files_shown = 0\n                     max_files_to_show = 3\n                     for f in files:\n                          if files_shown < max_files_to_show:\n                              print(f'{subindent}{f}')\n                              files_shown += 1\n                          else:\n                              print(f'{subindent}[... and {len(files) - max_files_to_show} more files]')\n                              break\n                     # Prune deeper exploration if needed for brevity\n                     # if level >= 1:\n                     #      dirs[:] = [] # Don't go deeper than level 1\n            except Exception as e_walk:\n                print(f\"Could not print directory structure: {e_walk}\")\n            print(\"--- End Directory Structure ---\\n\")\n\n            return # Stop execution if data prep fails\n\n\n        # Initialize model\n        print(\"Initializing model...\")\n        model = DeepfakeDetector(CONFIG).to(device)\n\n        # Count trainable parameters (after freezing)\n        trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        total_params = sum(p.numel() for p in model.parameters())\n        print(f\"Total parameters: {total_params:,}\")\n        print(f\"Trainable parameters: {trainable_params:,}\")\n\n\n        # Define loss function\n        if CONFIG['use_focal_loss']:\n            # Note: Class weights from sampler handle imbalance at data level.\n            # Focal loss handles imbalance at loss level (focusing on hard examples).\n            # Using both can be effective. We don't pass class weights directly to FocalLoss here.\n            criterion = FocalLoss(alpha=CONFIG['focal_alpha'], gamma=CONFIG['focal_gamma'])\n            print(f\"Using Focal Loss (alpha={CONFIG['focal_alpha']}, gamma={CONFIG['focal_gamma']})\")\n        else:\n            # If not using Focal Loss, consider using pos_weight in BCEWithLogitsLoss for imbalance\n            # pos_weight = weight for the positive class (Real=1)\n            # If class_weights = [weight_fake, weight_real], pos_weight = weight_real / weight_fake\n            # pos_weight_val = class_weights[1] / class_weights[0] if class_weights[0] > 0 else 1.0\n            # criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(pos_weight_val, device=device))\n            # print(f\"Using BCE Loss with pos_weight={pos_weight_val:.2f}\")\n             criterion = nn.BCEWithLogitsLoss() # Simpler BCE without explicit weighting if sampler is used\n             print(\"Using standard BCE Loss (relying on WeightedRandomSampler for imbalance)\")\n\n\n        # Define optimizer\n        optimizer = optim.AdamW(\n            model.parameters(), # Pass all parameters; optimizer respects requires_grad=False\n            lr=CONFIG['learning_rate'],\n            weight_decay=CONFIG['weight_decay']\n        )\n        print(f\"Using AdamW optimizer (lr={CONFIG['learning_rate']}, weight_decay={CONFIG['weight_decay']})\")\n        # --- Consider adjusting LR if performance is poor ---\n        print(\"--> NOTE: If initial training is unstable or accuracy is very low, consider lowering the learning rate (e.g., 1e-5 or 5e-5) for EfficientNet.\")\n\n\n        # Learning rate scheduler\n        scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n            optimizer,\n            mode='min', # Reduce LR when validation loss stops decreasing\n            factor=CONFIG['scheduler_factor'],\n            patience=CONFIG['scheduler_patience'],\n            verbose=True\n        )\n        print(f\"Using ReduceLROnPlateau scheduler (factor={CONFIG['scheduler_factor']}, patience={CONFIG['scheduler_patience']})\")\n\n        # Train the model\n        print(\"\\n--- Starting Training ---\")\n        model, metrics = train_model(\n            model=model,\n            train_loader=train_loader,\n            val_loader=val_loader,\n            criterion=criterion,\n            optimizer=optimizer,\n            scheduler=scheduler,\n            device=device,\n            num_epochs=CONFIG['num_epochs'],\n            config=CONFIG # Pass config for saving paths etc.\n        )\n\n        print(\"\\n--- Training Completed ---\")\n\n        # Save the final model (which is the best model based on validation accuracy)\n        # The best model state is already loaded back into 'model' by train_model\n        final_model_path = os.path.join(CONFIG['model_save_path'], \"final_best_model.pth\")\n        torch.save({\n            'model_state_dict': model.state_dict(), # Contains the best weights\n            'optimizer_state_dict': optimizer.state_dict(), # State at the end of training\n            'scheduler_state_dict': scheduler.state_dict(), # State at the end of training\n            'config': CONFIG,\n            'metrics': metrics # Optionally save training history\n        }, final_model_path)\n\n        print(f\"Best model (based on validation accuracy) saved to {final_model_path}\")\n\n    except Exception as e:\n        print(f\"\\n--- An error occurred during the main execution ---\")\n        import traceback\n        traceback.print_exc()\n\nif __name__ == \"__main__\":\n    main()","metadata":{"_uuid":"cf09951f-65e0-48c2-939c-24aa997d97cf","_cell_guid":"7a3d67a0-725b-474c-aef3-1f21f4db9bd3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-08-07T14:29:26.716969Z","iopub.execute_input":"2025-08-07T14:29:26.717248Z"}},"outputs":[],"execution_count":null}]}