{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":16880,"databundleVersionId":858837,"isSourceIdPinned":false}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"f86389e3","cell_type":"markdown","source":"# Deepfake Detection Pipeline\n\n- starts from raw DFDC videos\n- uses grouped train/validation split to reduce leakage from the same source video\n- extracts face crops with a fallback crop when face detection fails\n- caches extracted frames so reruns are cheaper\n- handles class imbalance with `WeightedRandomSampler` and `pos_weight`\n- uses mixed precision, gradient clipping, learning-rate scheduling, and early stopping\n- evaluates at both frame level and video level\n- tunes a validation threshold after training\n","metadata":{}},{"id":"6238b9f5","cell_type":"code","source":"!pip install -q opencv-python timm facenet-pytorch scikit-learn albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:05.766493Z","iopub.execute_input":"2026-04-16T04:48:05.766958Z","iopub.status.idle":"2026-04-16T04:48:11.203053Z","shell.execute_reply.started":"2026-04-16T04:48:05.766927Z","shell.execute_reply":"2026-04-16T04:48:11.202305Z"}},"outputs":[],"execution_count":null},{"id":"02c41f1a","cell_type":"markdown","source":"## 1. Imports\n","metadata":{}},{"id":"88e7bca0","cell_type":"code","source":"import json\nimport math\nimport os\nimport random\nfrom collections import defaultdict\nfrom pathlib import Path\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom albumentations import (\n    Compose,\n    Normalize,\n    Resize,\n    HorizontalFlip,\n    ImageCompression,\n    GaussianBlur,\n    ColorJitter,\n    RandomBrightnessContrast,\n    ShiftScaleRotate,\n)\nfrom albumentations.pytorch import ToTensorV2\nfrom facenet_pytorch import MTCNN\nfrom sklearn.metrics import (\n    accuracy_score,\n    classification_report,\n    confusion_matrix,\n    f1_score,\n    log_loss,\n    precision_score,\n    recall_score,\n    roc_auc_score,\n)\nfrom sklearn.model_selection import GroupShuffleSplit\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom tqdm.auto import tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:11.204634Z","iopub.execute_input":"2026-04-16T04:48:11.205007Z","iopub.status.idle":"2026-04-16T04:48:24.724288Z","shell.execute_reply.started":"2026-04-16T04:48:11.204977Z","shell.execute_reply":"2026-04-16T04:48:24.723682Z"}},"outputs":[],"execution_count":null},{"id":"6c66fdb1","cell_type":"markdown","source":"## 2. Configuration","metadata":{}},{"id":"a8efc599","cell_type":"code","source":"class CFG:\n    # data paths\n    DATA_ROOT = \"/kaggle/input/competitions/deepfake-detection-challenge\"\n    OUTPUT_DIR = \"/kaggle/working/output_improved\"\n    FRAME_DIR = \"/kaggle/working/extracted_faces_improved\"\n    CACHE_DIR = \"/kaggle/working/cache_improved\"\n\n    # dataset scope\n    USE_ONLY_SAMPLE_VIDEOS = True\n    MAX_VIDEOS = None  # set an integer for faster experiments\n    BALANCE_CLASSES = False\n    VAL_RATIO = 0.2\n    SEED = 42\n\n    # image / video preprocessing\n    IMG_SIZE = 224\n    NUM_FRAMES = 12\n    FRAME_SAMPLING = \"uniform\"  # uniform or random\n    MIN_FACE_SIZE = 40\n    FACE_MARGIN = 20\n    MIN_EXTRACTED_FRAMES = 4\n    REUSE_EXTRACTED = True\n\n    # dataloader / train\n    BATCH_SIZE = 16\n    EPOCHS = 8\n    LR = 1e-4\n    WEIGHT_DECAY = 1e-4\n    NUM_WORKERS = 2\n    USE_AMP = True\n    GRAD_CLIP_NORM = 1.0\n    EARLY_STOPPING_PATIENCE = 3\n\n    # model\n    MODEL_NAME = \"tf_efficientnet_b0.ns_jft_in1k\"\n    PRETRAINED = True\n    DROPOUT = 0.2\n\n    # evaluation\n    DEFAULT_THRESHOLD = 0.5\n    AGGREGATION = \"topk_mean\"  # mean, max, topk_mean\n    TOPK = 4\n\n    DEVICE = (\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"mps\"\n        if hasattr(torch.backends, \"mps\") and torch.backends.mps.is_available()\n        else \"cpu\"\n    )\n\n\nfor path in [CFG.OUTPUT_DIR, CFG.FRAME_DIR, CFG.CACHE_DIR]:\n    os.makedirs(path, exist_ok=True)\n\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n\nseed_everything(CFG.SEED)\nprint(\"Using device:\", CFG.DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:24.725265Z","iopub.execute_input":"2026-04-16T04:48:24.725733Z","iopub.status.idle":"2026-04-16T04:48:24.993294Z","shell.execute_reply.started":"2026-04-16T04:48:24.725707Z","shell.execute_reply":"2026-04-16T04:48:24.992561Z"}},"outputs":[],"execution_count":null},{"id":"3a76c31f-478c-4a30-9906-45ed4a4dd341","cell_type":"markdown","source":"## 3. Load Metadata\n\nWe load DFDC metadata from training folders only, then build a grouped split.\nThe grouping key keeps each original video and its derived fakes in the same split to reduce leakage.\n","metadata":{}},{"id":"5ab35cad","cell_type":"code","source":"def discover_training_folders(data_root, use_only_sample_videos=True):\n    root = Path(data_root)\n    if use_only_sample_videos:\n        sample_dir = root / \"train_sample_videos\"\n        return [sample_dir] if sample_dir.exists() else []\n    return sorted(\n        folder\n        for folder in root.iterdir()\n        if folder.is_dir() and (folder / \"metadata.json\").exists()\n    )\n\n\ndef load_dfdc_metadata(data_root, use_only_sample_videos=True, max_videos=None):\n    rows = []\n    for folder in discover_training_folders(data_root, use_only_sample_videos):\n        metadata_path = folder / \"metadata.json\"\n        with open(metadata_path, \"r\") as f:\n            metadata = json.load(f)\n\n        items = sorted(metadata.items())\n        if max_videos is not None:\n            items = items[:max_videos]\n\n        for filename, info in items:\n            video_path = folder / filename\n            if not video_path.exists():\n                continue\n\n            label = 1 if info[\"label\"].upper() == \"FAKE\" else 0\n            original = info.get(\"original\")\n            video_id = video_path.stem\n            group_key = original if original is not None else video_id\n\n            rows.append({\n                \"video_path\": str(video_path),\n                \"video_id\": video_id,\n                \"label\": label,\n                \"original\": original,\n                \"group_key\": group_key,\n                \"folder\": folder.name,\n            })\n\n    df = pd.DataFrame(rows)\n    if df.empty:\n        raise ValueError(\"No training videos were found. Check CFG.DATA_ROOT.\")\n    return df\n\n\ndf = load_dfdc_metadata(\n    CFG.DATA_ROOT,\n    use_only_sample_videos=CFG.USE_ONLY_SAMPLE_VIDEOS,\n    max_videos=CFG.MAX_VIDEOS,\n)\n\nprint(df.shape)\ndf.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:24.995193Z","iopub.execute_input":"2026-04-16T04:48:24.995510Z","iopub.status.idle":"2026-04-16T04:48:31.285587Z","shell.execute_reply.started":"2026-04-16T04:48:24.995483Z","shell.execute_reply":"2026-04-16T04:48:31.284715Z"}},"outputs":[],"execution_count":null},{"id":"bad15d9c","cell_type":"code","source":"print(df[\"label\"].value_counts(dropna=False))\nprint(\"Unique groups:\", df[\"group_key\"].nunique())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:31.286716Z","iopub.execute_input":"2026-04-16T04:48:31.287092Z","iopub.status.idle":"2026-04-16T04:48:31.296446Z","shell.execute_reply.started":"2026-04-16T04:48:31.287053Z","shell.execute_reply":"2026-04-16T04:48:31.295684Z"}},"outputs":[],"execution_count":null},{"id":"6e9f395e","cell_type":"markdown","source":"## 4. Grouped Train / Validation Split\n","metadata":{}},{"id":"3824c231","cell_type":"code","source":"def make_balanced_subset(df, max_videos=None, seed=42):\n    real_df = df[df[\"label\"] == 0]\n    fake_df = df[df[\"label\"] == 1]\n\n    if max_videos is None:\n        n = min(len(real_df), len(fake_df))\n    else:\n        n = min(len(real_df), len(fake_df), max_videos // 2)\n\n    balanced_df = pd.concat([\n        real_df.sample(n=n, random_state=seed),\n        fake_df.sample(n=n, random_state=seed),\n    ])\n\n    return balanced_df.sample(frac=1, random_state=seed).reset_index(drop=True)\n\n\ndef make_grouped_split(df, val_ratio=0.2, seed=42):\n    splitter = GroupShuffleSplit(n_splits=1, test_size=val_ratio, random_state=seed)\n    train_idx, val_idx = next(splitter.split(df, y=df[\"label\"], groups=df[\"group_key\"]))\n    train_df = df.iloc[train_idx].reset_index(drop=True)\n    val_df = df.iloc[val_idx].reset_index(drop=True)\n    return train_df, val_df\n\n\nworking_df = df.copy()\n\nif CFG.BALANCE_CLASSES:\n    working_df = make_balanced_subset(\n        working_df,\n        max_videos=CFG.MAX_VIDEOS,\n        seed=CFG.SEED,\n    )\n\nprint(\"Working dataset shape:\", working_df.shape)\nprint(\"Working label counts:\")\nprint(working_df[\"label\"].value_counts())\n\ntrain_df, val_df = make_grouped_split(\n    working_df,\n    val_ratio=CFG.VAL_RATIO,\n    seed=CFG.SEED,\n)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:  \", val_df.shape)\nprint(\"Train label counts:\")\nprint(train_df[\"label\"].value_counts())\nprint(\"Val label counts:\")\nprint(val_df[\"label\"].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:31.297375Z","iopub.execute_input":"2026-04-16T04:48:31.297653Z","iopub.status.idle":"2026-04-16T04:48:31.315650Z","shell.execute_reply.started":"2026-04-16T04:48:31.297617Z","shell.execute_reply":"2026-04-16T04:48:31.315017Z"}},"outputs":[],"execution_count":null},{"id":"abd5b5e3","cell_type":"markdown","source":"## 5. Face Extractor With Fallback Crop\n\nIf MTCNN does not find a face on a sampled frame, we keep a centered square crop instead of dropping the frame.\nThis makes the pipeline more robust and reduces missing-video problems.\n","metadata":{}},{"id":"83cfcf5b","cell_type":"code","source":"mtcnn_device = \"cpu\" if CFG.DEVICE == \"mps\" else CFG.DEVICE\n\nmtcnn = MTCNN(\n    image_size=CFG.IMG_SIZE,\n    margin=CFG.FACE_MARGIN,\n    min_face_size=CFG.MIN_FACE_SIZE,\n    post_process=False,\n    device=mtcnn_device,\n)\n\n\ndef sample_frame_indices(num_total_frames, num_samples=8, mode=\"uniform\"):\n    if num_total_frames <= 0:\n        return []\n    if num_total_frames <= num_samples:\n        return list(range(num_total_frames))\n    if mode == \"random\":\n        return sorted(random.sample(range(num_total_frames), num_samples))\n    return np.linspace(0, num_total_frames - 1, num_samples).astype(int).tolist()\n\n\ndef center_crop_square(image):\n    h, w = image.shape[:2]\n    side = min(h, w)\n    y1 = (h - side) // 2\n    x1 = (w - side) // 2\n    return image[y1:y1 + side, x1:x1 + side]\n\n\ndef crop_largest_face(rgb_image, boxes):\n    if boxes is None or len(boxes) == 0:\n        return None\n\n    areas = []\n    for box in boxes:\n        x1, y1, x2, y2 = box\n        areas.append(max(0, x2 - x1) * max(0, y2 - y1))\n\n    best_idx = int(np.argmax(areas))\n    x1, y1, x2, y2 = boxes[best_idx].astype(int)\n    h, w = rgb_image.shape[:2]\n    x1 = max(0, x1)\n    y1 = max(0, y1)\n    x2 = min(w, x2)\n    y2 = min(h, y2)\n\n    if x2 <= x1 or y2 <= y1:\n        return None\n    return rgb_image[y1:y2, x1:x2]\n\n\ndef extract_faces_from_video(video_path, save_dir, label, num_frames=8):\n    save_dir = Path(save_dir)\n    save_dir.mkdir(parents=True, exist_ok=True)\n\n    existing = sorted(save_dir.glob(\"*.jpg\"))\n    if CFG.REUSE_EXTRACTED and len(existing) >= CFG.MIN_EXTRACTED_FRAMES:\n        return [str(path) for path in existing]\n\n    for old_file in save_dir.glob(\"*.jpg\"):\n        old_file.unlink()\n\n    cap = cv2.VideoCapture(str(video_path))\n    total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))\n    target_indices = set(sample_frame_indices(total_frames, num_frames, CFG.FRAME_SAMPLING))\n\n    saved_paths = []\n    frame_idx = 0\n\n    while cap.isOpened():\n        ok, frame = cap.read()\n        if not ok:\n            break\n\n        if frame_idx in target_indices:\n            rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n            boxes, _ = mtcnn.detect(rgb)\n            face = crop_largest_face(rgb, boxes)\n            used_fallback = face is None\n\n            if used_fallback:\n                face = center_crop_square(rgb)\n\n            if face is not None and face.size > 0:\n                face = cv2.resize(face, (CFG.IMG_SIZE, CFG.IMG_SIZE))\n                suffix = \"fallback\" if used_fallback else \"face\"\n                out_path = save_dir / f\"{Path(video_path).stem}_f{frame_idx}_{suffix}_label{label}.jpg\"\n                cv2.imwrite(str(out_path), cv2.cvtColor(face, cv2.COLOR_RGB2BGR))\n                saved_paths.append(str(out_path))\n\n        frame_idx += 1\n\n    cap.release()\n    return saved_paths\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:59.512096Z","iopub.execute_input":"2026-04-16T04:48:59.512897Z","iopub.status.idle":"2026-04-16T04:48:59.545178Z","shell.execute_reply.started":"2026-04-16T04:48:59.512865Z","shell.execute_reply":"2026-04-16T04:48:59.544577Z"}},"outputs":[],"execution_count":null},{"id":"76b2fb34","cell_type":"markdown","source":"## 6. Build Cached Frame Tables\n","metadata":{}},{"id":"10a724be","cell_type":"code","source":"def build_face_frame_table(df, split_name):\n    rows = []\n    split_dir = Path(CFG.FRAME_DIR) / split_name\n    split_dir.mkdir(parents=True, exist_ok=True)\n\n    for row in tqdm(df.itertuples(index=False), total=len(df), desc=f\"Extracting {split_name}\"):\n        video_dir = split_dir / row.video_id\n        saved_paths = extract_faces_from_video(\n            video_path=row.video_path,\n            save_dir=video_dir,\n            label=row.label,\n            num_frames=CFG.NUM_FRAMES,\n        )\n\n        for image_path in saved_paths:\n            rows.append({\n                \"image_path\": image_path,\n                \"label\": row.label,\n                \"video_id\": row.video_id,\n                \"video_path\": row.video_path,\n                \"group_key\": row.group_key,\n            })\n\n    frame_df = pd.DataFrame(rows)\n    cache_path = Path(CFG.CACHE_DIR) / f\"{split_name}_frames.csv\"\n    frame_df.to_csv(cache_path, index=False)\n    print(f\"Saved frame cache to: {cache_path}\")\n    return frame_df\n\n\ntrain_frames_df = build_face_frame_table(train_df, \"train\")\nval_frames_df = build_face_frame_table(val_df, \"val\")\n\nprint(train_frames_df.shape, val_frames_df.shape)\ntrain_frames_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T04:48:59.546540Z","iopub.execute_input":"2026-04-16T04:48:59.546811Z","iopub.status.idle":"2026-04-16T05:04:29.720807Z","shell.execute_reply.started":"2026-04-16T04:48:59.546787Z","shell.execute_reply":"2026-04-16T05:04:29.719965Z"}},"outputs":[],"execution_count":null},{"id":"3f6b99de","cell_type":"code","source":"def summarize_extraction(video_df, frame_df, split_name):\n    extracted_videos = frame_df[\"video_id\"].nunique()\n    total_videos = video_df[\"video_id\"].nunique()\n    frame_counts = frame_df.groupby(\"video_id\").size()\n    print(f\"{split_name} videos:           {total_videos}\")\n    print(f\"{split_name} extracted videos: {extracted_videos}\")\n    print(f\"{split_name} extracted frames: {len(frame_df)}\")\n    print(f\"{split_name} avg frames/video: {frame_counts.mean():.2f}\")\n    print(f\"{split_name} min frames/video: {frame_counts.min()}\")\n    print(f\"{split_name} max frames/video: {frame_counts.max()}\")\n\n\nsummarize_extraction(train_df, train_frames_df, \"Train\")\nsummarize_extraction(val_df, val_frames_df, \"Val\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:29.721890Z","iopub.execute_input":"2026-04-16T05:04:29.722196Z","iopub.status.idle":"2026-04-16T05:04:29.732185Z","shell.execute_reply.started":"2026-04-16T05:04:29.722170Z","shell.execute_reply":"2026-04-16T05:04:29.731331Z"}},"outputs":[],"execution_count":null},{"id":"25ca3cb3","cell_type":"markdown","source":"## 7. Augmentations and Dataset\n","metadata":{}},{"id":"4356b5af","cell_type":"code","source":"train_transform = Compose([\n    Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n    HorizontalFlip(p=0.5),\n    ImageCompression(quality_range=(50, 100), p=0.3),\n    GaussianBlur(blur_limit=(3, 5), p=0.2),\n    ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05, p=0.3),\n    RandomBrightnessContrast(p=0.2),\n    ShiftScaleRotate(\n        shift_limit=0.03,\n        scale_limit=0.05,\n        rotate_limit=5,\n        border_mode=cv2.BORDER_REFLECT_101,\n        p=0.3,\n    ),\n    Normalize(),\n    ToTensorV2(),\n])\n\nval_transform = Compose([\n    Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n    Normalize(),\n    ToTensorV2(),\n])\n\n\nclass DeepfakeFrameDataset(Dataset):\n    def __init__(self, frame_df, transforms=None):\n        self.frame_df = frame_df.reset_index(drop=True)\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.frame_df)\n\n    def __getitem__(self, idx):\n        row = self.frame_df.iloc[idx]\n        image = cv2.imread(row[\"image_path\"])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transforms is not None:\n            image = self.transforms(image=image)[\"image\"]\n        label = torch.tensor(row[\"label\"], dtype=torch.float32)\n        return image, label, row[\"video_id\"]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:29.733974Z","iopub.execute_input":"2026-04-16T05:04:29.734216Z","iopub.status.idle":"2026-04-16T05:04:29.749087Z","shell.execute_reply.started":"2026-04-16T05:04:29.734195Z","shell.execute_reply":"2026-04-16T05:04:29.748321Z"}},"outputs":[],"execution_count":null},{"id":"0a47c53c","cell_type":"code","source":"train_dataset = DeepfakeFrameDataset(train_frames_df, transforms=train_transform)\nval_dataset = DeepfakeFrameDataset(val_frames_df, transforms=val_transform)\n\n\ndef build_weighted_sampler(frame_df):\n    class_counts = frame_df[\"label\"].value_counts().sort_index()\n    class_weights = {label: 1.0 / count for label, count in class_counts.items()}\n    sample_weights = frame_df[\"label\"].map(class_weights).to_numpy(dtype=np.float64)\n    return WeightedRandomSampler(\n        weights=torch.as_tensor(sample_weights, dtype=torch.double),\n        num_samples=len(sample_weights),\n        replacement=True,\n    )\n\n\ntrain_sampler = build_weighted_sampler(train_frames_df)\npin_memory = CFG.DEVICE == \"cuda\"\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    sampler=train_sampler,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=pin_memory,\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=pin_memory,\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:  \", len(val_loader))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:29.749962Z","iopub.execute_input":"2026-04-16T05:04:29.750237Z","iopub.status.idle":"2026-04-16T05:04:29.763598Z","shell.execute_reply.started":"2026-04-16T05:04:29.750206Z","shell.execute_reply":"2026-04-16T05:04:29.762790Z"}},"outputs":[],"execution_count":null},{"id":"d060cc32","cell_type":"code","source":"images, labels, video_ids = next(iter(train_loader))\n\nfig, axes = plt.subplots(2, 4, figsize=(12, 6))\naxes = axes.flatten()\n\nfor i in range(min(8, len(images))):\n    image = images[i].permute(1, 2, 0).cpu().numpy()\n    image = (image - image.min()) / (image.max() - image.min() + 1e-8)\n    axes[i].imshow(image)\n    axes[i].set_title(f\"label={int(labels[i].item())}\")\n    axes[i].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:29.764420Z","iopub.execute_input":"2026-04-16T05:04:29.764755Z","iopub.status.idle":"2026-04-16T05:04:30.711247Z","shell.execute_reply.started":"2026-04-16T05:04:29.764720Z","shell.execute_reply":"2026-04-16T05:04:30.710259Z"}},"outputs":[],"execution_count":null},{"id":"98fe98e0","cell_type":"markdown","source":"## 8. Model, Loss, and Optimizer\n","metadata":{}},{"id":"c2650799","cell_type":"code","source":"class DeepfakeClassifier(nn.Module):\n    def __init__(self, model_name, pretrained=True, dropout=0.0):\n        super().__init__()\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            num_classes=1,\n            drop_rate=dropout,\n        )\n\n    def forward(self, x):\n        return self.backbone(x).squeeze(1)\n\n\nmodel = DeepfakeClassifier(\n    CFG.MODEL_NAME,\n    pretrained=CFG.PRETRAINED,\n    dropout=CFG.DROPOUT,\n).to(CFG.DEVICE)\n\npositive_count = max(1, int((train_frames_df[\"label\"] == 1).sum()))\nnegative_count = max(1, int((train_frames_df[\"label\"] == 0).sum()))\npos_weight = torch.tensor([negative_count / positive_count], device=CFG.DEVICE, dtype=torch.float32)\n\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\noptimizer = optim.AdamW(model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1,\n)\n\nscaler = torch.cuda.amp.GradScaler(enabled=(CFG.USE_AMP and CFG.DEVICE == \"cuda\"))\n\nprint(\"pos_weight:\", float(pos_weight.item()))\nmodel\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:30.712593Z","iopub.execute_input":"2026-04-16T05:04:30.712874Z","iopub.status.idle":"2026-04-16T05:04:32.726678Z","shell.execute_reply.started":"2026-04-16T05:04:30.712835Z","shell.execute_reply":"2026-04-16T05:04:32.726097Z"}},"outputs":[],"execution_count":null},{"id":"9324b956","cell_type":"markdown","source":"## 9. Training and Evaluation Helpers\n","metadata":{}},{"id":"16a503f7","cell_type":"code","source":"def aggregate_video_probabilities(frame_df, method=\"mean\", topk=4):\n    rows = []\n    for video_id, group in frame_df.groupby(\"video_id\"):\n        probs = group[\"prob_fake\"].to_numpy(dtype=float)\n        label = int(group[\"label\"].iloc[0])\n\n        if method == \"max\":\n            video_prob = float(np.max(probs))\n        elif method == \"topk_mean\":\n            k = min(topk, len(probs))\n            video_prob = float(np.mean(np.sort(probs)[-k:]))\n        else:\n            video_prob = float(np.mean(probs))\n\n        rows.append({\n            \"video_id\": video_id,\n            \"label\": label,\n            \"prob_fake\": video_prob,\n            \"num_frames\": len(group),\n        })\n\n    return pd.DataFrame(rows)\n\n\ndef compute_binary_metrics(y_true, y_prob, threshold=0.5, prefix=\"\"):\n    y_true = np.asarray(y_true, dtype=int)\n    y_prob = np.asarray(y_prob, dtype=float)\n    y_pred = (y_prob >= threshold).astype(int)\n\n    metrics = {\n        f\"{prefix}acc\": accuracy_score(y_true, y_pred),\n        f\"{prefix}precision\": precision_score(y_true, y_pred, zero_division=0),\n        f\"{prefix}recall\": recall_score(y_true, y_pred, zero_division=0),\n        f\"{prefix}f1\": f1_score(y_true, y_pred, zero_division=0),\n        f\"{prefix}auc\": roc_auc_score(y_true, y_prob) if len(np.unique(y_true)) > 1 else 0.0,\n        f\"{prefix}logloss\": log_loss(y_true, np.clip(y_prob, 1e-6, 1 - 1e-6)),\n    }\n    return metrics\n\n\ndef find_best_threshold(y_true, y_prob, thresholds=None):\n    if thresholds is None:\n        thresholds = np.linspace(0.1, 0.9, 81)\n\n    best = {\"threshold\": 0.5, \"f1\": -1.0}\n    for threshold in thresholds:\n        f1 = f1_score(y_true, (np.asarray(y_prob) >= threshold).astype(int), zero_division=0)\n        if f1 > best[\"f1\"]:\n            best = {\"threshold\": float(threshold), \"f1\": float(f1)}\n    return best\n\n\ndef train_one_epoch(model, loader, optimizer, criterion, device, scaler):\n    model.train()\n    running_loss = 0.0\n\n    for images, labels, _ in tqdm(loader, desc=\"Train\", leave=False):\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast(enabled=(CFG.USE_AMP and device == \"cuda\")):\n            logits = model(images)\n            loss = criterion(logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.GRAD_CLIP_NORM)\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * images.size(0)\n\n    return running_loss / len(loader.dataset)\n\n\n@torch.no_grad()\ndef evaluate_loader(model, loader, criterion, device, threshold=0.5):\n    model.eval()\n    running_loss = 0.0\n    frame_rows = []\n\n    for images, labels, video_ids in tqdm(loader, desc=\"Eval\", leave=False):\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        with torch.cuda.amp.autocast(enabled=(CFG.USE_AMP and device == \"cuda\")):\n            logits = model(images)\n            loss = criterion(logits, labels)\n\n        probs = torch.sigmoid(logits).detach().cpu().numpy()\n        labels_np = labels.detach().cpu().numpy()\n        running_loss += loss.item() * images.size(0)\n\n        for prob, label, video_id in zip(probs.tolist(), labels_np.tolist(), list(video_ids)):\n            frame_rows.append({\n                \"video_id\": video_id,\n                \"label\": int(label),\n                \"prob_fake\": float(prob),\n            })\n\n    frame_df = pd.DataFrame(frame_rows)\n    frame_metrics = compute_binary_metrics(\n        frame_df[\"label\"].values,\n        frame_df[\"prob_fake\"].values,\n        threshold=threshold,\n        prefix=\"frame_\",\n    )\n\n    video_df = aggregate_video_probabilities(\n        frame_df,\n        method=CFG.AGGREGATION,\n        topk=CFG.TOPK,\n    )\n    video_metrics = compute_binary_metrics(\n        video_df[\"label\"].values,\n        video_df[\"prob_fake\"].values,\n        threshold=threshold,\n        prefix=\"video_\",\n    )\n\n    return {\n        \"loss\": running_loss / len(loader.dataset),\n        \"frame_metrics\": frame_metrics,\n        \"video_metrics\": video_metrics,\n        \"frame_df\": frame_df,\n        \"video_df\": video_df,\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:32.727699Z","iopub.execute_input":"2026-04-16T05:04:32.728049Z","iopub.status.idle":"2026-04-16T05:04:32.743493Z","shell.execute_reply.started":"2026-04-16T05:04:32.728025Z","shell.execute_reply":"2026-04-16T05:04:32.742695Z"}},"outputs":[],"execution_count":null},{"id":"bf3c205b","cell_type":"markdown","source":"## 10. Train the Model\n\nThe checkpoint is chosen by validation `video_auc` because it is threshold-independent.\nThreshold tuning is done only after training finishes.\n","metadata":{}},{"id":"8c7d0e42","cell_type":"code","source":"best_video_auc = -1.0\nbest_model_path = Path(CFG.OUTPUT_DIR) / \"best_model_improved.pth\"\nhistory = []\nepochs_without_improvement = 0\n\nfor epoch in range(1, CFG.EPOCHS + 1):\n    print(f\"\\nEpoch [{epoch}/{CFG.EPOCHS}]\")\n\n    train_loss = train_one_epoch(model, train_loader, optimizer, criterion, CFG.DEVICE, scaler)\n    val_result = evaluate_loader(\n        model,\n        val_loader,\n        criterion,\n        CFG.DEVICE,\n        threshold=CFG.DEFAULT_THRESHOLD,\n    )\n\n    current_video_auc = val_result[\"video_metrics\"][\"video_auc\"]\n    scheduler.step(current_video_auc)\n\n    epoch_row = {\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"val_loss\": val_result[\"loss\"],\n        **val_result[\"frame_metrics\"],\n        **val_result[\"video_metrics\"],\n    }\n    history.append(epoch_row)\n\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_result['loss']:.4f}\")\n    print(\n        f\"Frame F1:   {val_result['frame_metrics']['frame_f1']:.4f} | \"\n        f\"Frame AUC: {val_result['frame_metrics']['frame_auc']:.4f}\"\n    )\n    print(\n        f\"Video F1:   {val_result['video_metrics']['video_f1']:.4f} | \"\n        f\"Video AUC: {val_result['video_metrics']['video_auc']:.4f} | \"\n        f\"Video LogLoss: {val_result['video_metrics']['video_logloss']:.4f}\"\n    )\n\n    if current_video_auc > best_video_auc:\n        best_video_auc = current_video_auc\n        epochs_without_improvement = 0\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"Saved best model to: {best_model_path}\")\n    else:\n        epochs_without_improvement += 1\n        print(f\"No improvement for {epochs_without_improvement} epoch(s).\")\n\n    if epochs_without_improvement >= CFG.EARLY_STOPPING_PATIENCE:\n        print(\"Early stopping triggered.\")\n        break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:04:32.744578Z","iopub.execute_input":"2026-04-16T05:04:32.745087Z","iopub.status.idle":"2026-04-16T05:08:29.026268Z","shell.execute_reply.started":"2026-04-16T05:04:32.745064Z","shell.execute_reply":"2026-04-16T05:08:29.024994Z"}},"outputs":[],"execution_count":null},{"id":"1c28a6f4","cell_type":"code","source":"history_df = pd.DataFrame(history)\nhistory_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:29.032995Z","iopub.execute_input":"2026-04-16T05:08:29.033605Z","iopub.status.idle":"2026-04-16T05:08:29.057974Z","shell.execute_reply.started":"2026-04-16T05:08:29.033525Z","shell.execute_reply":"2026-04-16T05:08:29.057099Z"}},"outputs":[],"execution_count":null},{"id":"1d16dd91","cell_type":"code","source":"plt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_loss\"], label=\"train_loss\")\nplt.plot(history_df[\"epoch\"], history_df[\"val_loss\"], label=\"val_loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend()\nplt.show()\n\nplt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"video_auc\"], label=\"video_auc\")\nplt.plot(history_df[\"epoch\"], history_df[\"video_f1\"], label=\"video_f1\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Score\")\nplt.title(\"Validation Video Metrics\")\nplt.legend()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:29.059367Z","iopub.execute_input":"2026-04-16T05:08:29.059827Z","iopub.status.idle":"2026-04-16T05:08:39.361131Z","shell.execute_reply.started":"2026-04-16T05:08:29.059802Z","shell.execute_reply":"2026-04-16T05:08:39.360462Z"}},"outputs":[],"execution_count":null},{"id":"76c1bead","cell_type":"markdown","source":"## 11. Load Best Model and Tune Threshold\n","metadata":{}},{"id":"7f1e62f4","cell_type":"code","source":"best_model = DeepfakeClassifier(\n    CFG.MODEL_NAME,\n    pretrained=False,\n    dropout=CFG.DROPOUT,\n).to(CFG.DEVICE)\nbest_model.load_state_dict(torch.load(best_model_path, map_location=CFG.DEVICE))\nbest_model.eval()\nprint(\"Loaded best checkpoint.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:39.362099Z","iopub.execute_input":"2026-04-16T05:08:39.362356Z","iopub.status.idle":"2026-04-16T05:08:39.550270Z","shell.execute_reply.started":"2026-04-16T05:08:39.362333Z","shell.execute_reply":"2026-04-16T05:08:39.549626Z"}},"outputs":[],"execution_count":null},{"id":"47755bcb","cell_type":"code","source":"val_result_default = evaluate_loader(\n    best_model,\n    val_loader,\n    criterion,\n    CFG.DEVICE,\n    threshold=CFG.DEFAULT_THRESHOLD,\n)\n\nbest_threshold_info = find_best_threshold(\n    val_result_default[\"video_df\"][\"label\"].values,\n    val_result_default[\"video_df\"][\"prob_fake\"].values,\n)\n\nval_result_tuned = evaluate_loader(\n    best_model,\n    val_loader,\n    criterion,\n    CFG.DEVICE,\n    threshold=best_threshold_info[\"threshold\"],\n)\n\nprint(\"Default threshold:\", CFG.DEFAULT_THRESHOLD)\nprint(\"Best validation threshold:\", best_threshold_info[\"threshold\"])\nprint(\"Best validation F1 at tuned threshold:\", best_threshold_info[\"f1\"])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:39.551228Z","iopub.execute_input":"2026-04-16T05:08:39.551542Z","iopub.status.idle":"2026-04-16T05:08:52.703359Z","shell.execute_reply.started":"2026-04-16T05:08:39.551518Z","shell.execute_reply":"2026-04-16T05:08:52.702577Z"}},"outputs":[],"execution_count":null},{"id":"4b5ea90d","cell_type":"markdown","source":"## 12. Final Validation Report\n","metadata":{}},{"id":"26be667b","cell_type":"code","source":"def print_metric_block(title, metrics):\n    print(title)\n    for key, value in metrics.items():\n        print(f\"  {key}: {value:.4f}\")\n\n\nprint_metric_block(\"Frame metrics @ default threshold\", val_result_default[\"frame_metrics\"])\nprint_metric_block(\"Video metrics @ default threshold\", val_result_default[\"video_metrics\"])\nprint()\nprint_metric_block(\"Frame metrics @ tuned threshold\", val_result_tuned[\"frame_metrics\"])\nprint_metric_block(\"Video metrics @ tuned threshold\", val_result_tuned[\"video_metrics\"])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:52.705105Z","iopub.execute_input":"2026-04-16T05:08:52.705380Z","iopub.status.idle":"2026-04-16T05:08:52.712208Z","shell.execute_reply.started":"2026-04-16T05:08:52.705344Z","shell.execute_reply":"2026-04-16T05:08:52.711561Z"}},"outputs":[],"execution_count":null},{"id":"29e1ada2","cell_type":"code","source":"val_video_df = val_result_tuned[\"video_df\"].copy()\nval_video_df[\"pred\"] = (val_video_df[\"prob_fake\"] >= best_threshold_info[\"threshold\"]).astype(int)\nval_video_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:52.713370Z","iopub.execute_input":"2026-04-16T05:08:52.714075Z","iopub.status.idle":"2026-04-16T05:08:52.734102Z","shell.execute_reply.started":"2026-04-16T05:08:52.714053Z","shell.execute_reply":"2026-04-16T05:08:52.733416Z"}},"outputs":[],"execution_count":null},{"id":"065ca30c","cell_type":"code","source":"y_true = val_video_df[\"label\"].values\ny_pred = val_video_df[\"pred\"].values\n\nprint(classification_report(y_true, y_pred, target_names=[\"Real\", \"Fake\"], zero_division=0))\nprint(\"Confusion Matrix:\")\nprint(confusion_matrix(y_true, y_pred))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:52.735018Z","iopub.execute_input":"2026-04-16T05:08:52.735284Z","iopub.status.idle":"2026-04-16T05:08:52.762506Z","shell.execute_reply.started":"2026-04-16T05:08:52.735253Z","shell.execute_reply":"2026-04-16T05:08:52.761720Z"}},"outputs":[],"execution_count":null},{"id":"ec045672","cell_type":"code","source":"val_predictions_path = Path(CFG.OUTPUT_DIR) / \"val_video_predictions_improved.csv\"\nval_video_df.to_csv(val_predictions_path, index=False)\nprint(f\"Saved validation predictions to: {val_predictions_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:52.763460Z","iopub.execute_input":"2026-04-16T05:08:52.764028Z","iopub.status.idle":"2026-04-16T05:08:52.770632Z","shell.execute_reply.started":"2026-04-16T05:08:52.763996Z","shell.execute_reply":"2026-04-16T05:08:52.770010Z"}},"outputs":[],"execution_count":null},{"id":"ba43531c","cell_type":"markdown","source":"## 13. Saliency Map\n\nThis section visualizes which image regions most influence the model's prediction for selected validation frames.\nThe implementation uses vanilla input-gradient saliency so it works directly with the current classifier.\n","metadata":{}},{"id":"a8568a4b","cell_type":"code","source":"def compute_saliency_map(model, image_tensor, device):\n    model.eval()\n    x = image_tensor.unsqueeze(0).to(device)\n    x.requires_grad_(True)\n\n    logits = model(x)\n    score = logits.squeeze()\n\n    model.zero_grad(set_to_none=True)\n    score.backward()\n\n    saliency = x.grad.detach().abs().max(dim=1)[0].squeeze(0).cpu().numpy()\n    saliency = saliency - saliency.min()\n    saliency = saliency / (saliency.max() + 1e-8)\n    return saliency\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:08:52.771617Z","iopub.execute_input":"2026-04-16T05:08:52.771952Z","iopub.status.idle":"2026-04-16T05:08:52.783360Z","shell.execute_reply.started":"2026-04-16T05:08:52.771891Z","shell.execute_reply":"2026-04-16T05:08:52.782494Z"}},"outputs":[],"execution_count":null},{"id":"21e412da","cell_type":"code","source":"import re\n\n\ndef get_frame_index_from_image_path(image_path):\n    match = re.search(r\"_f(\\d+)_\", Path(image_path).name)\n    if match is None:\n        return None\n    return int(match.group(1))\n\n\ndef read_video_frame(video_path, frame_idx):\n    cap = cv2.VideoCapture(str(video_path))\n    cap.set(cv2.CAP_PROP_POS_FRAMES, frame_idx)\n\n    ok, frame = cap.read()\n    cap.release()\n\n    if not ok:\n        return None\n\n    return cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n\n\ndef show_full_frame_context_with_heatmap(model, dataset, frame_df, indices=None, max_examples=4):\n    if indices is None:\n        indices = list(range(min(max_examples, len(dataset))))\n    else:\n        indices = indices[:max_examples]\n\n    fig, axes = plt.subplots(len(indices), 4, figsize=(18, 4 * len(indices)))\n\n    if len(indices) == 1:\n        axes = np.expand_dims(axes, axis=0)\n\n    for row_idx, sample_idx in enumerate(indices):\n        row = frame_df.iloc[sample_idx]\n\n        image_tensor, label, video_id = dataset[sample_idx]\n        image = unnormalize_image(image_tensor)\n\n        raw_heatmap = compute_saliency_map(model, image_tensor, CFG.DEVICE)\n        heatmap = enhance_heatmap(\n            raw_heatmap,\n            lower_percentile=70,\n            upper_percentile=99.5,\n            blur_sigma=2,\n        )\n\n        with torch.no_grad():\n            prob_fake = torch.sigmoid(\n                model(image_tensor.unsqueeze(0).to(CFG.DEVICE))\n            ).item()\n\n        frame_idx = get_frame_index_from_image_path(row[\"image_path\"])\n        full_frame = read_video_frame(row[\"video_path\"], frame_idx) if frame_idx is not None else None\n\n        if full_frame is not None:\n            axes[row_idx, 0].imshow(full_frame)\n            axes[row_idx, 0].set_title(f\"Full frame | f={frame_idx}\")\n        else:\n            axes[row_idx, 0].text(0.5, 0.5, \"Could not read full frame\", ha=\"center\", va=\"center\")\n            axes[row_idx, 0].set_title(\"Full frame unavailable\")\n        axes[row_idx, 0].axis(\"off\")\n\n        axes[row_idx, 1].imshow(image)\n        axes[row_idx, 1].set_title(f\"Model input crop | label={int(label.item())}\")\n        axes[row_idx, 1].axis(\"off\")\n\n        heat = axes[row_idx, 2].imshow(heatmap, cmap=\"magma\", vmin=0, vmax=1)\n        axes[row_idx, 2].set_title(\"Prediction Heatmap\")\n        axes[row_idx, 2].axis(\"off\")\n        fig.colorbar(heat, ax=axes[row_idx, 2], fraction=0.046, pad=0.04)\n\n        axes[row_idx, 3].imshow(image)\n        axes[row_idx, 3].imshow(heatmap, cmap=\"magma\", alpha=0.42, vmin=0, vmax=1)\n        axes[row_idx, 3].set_title(f\"Overlay | p(fake)={prob_fake:.3f}\")\n        axes[row_idx, 3].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:19:33.945838Z","iopub.execute_input":"2026-04-16T05:19:33.946364Z","iopub.status.idle":"2026-04-16T05:19:33.959938Z","shell.execute_reply.started":"2026-04-16T05:19:33.946332Z","shell.execute_reply":"2026-04-16T05:19:33.959071Z"}},"outputs":[],"execution_count":null},{"id":"fdb8097c","cell_type":"code","source":"positive_indices = val_frames_df.index[val_frames_df[\"label\"] == 1].tolist()\nnegative_indices = val_frames_df.index[val_frames_df[\"label\"] == 0].tolist()\n\nselected_indices = []\nif negative_indices:\n    selected_indices.append(negative_indices[0])\nif positive_indices:\n    selected_indices.append(positive_indices[0])\nif len(negative_indices) > 1:\n    selected_indices.append(negative_indices[1])\nif len(positive_indices) > 1:\n    selected_indices.append(positive_indices[1])\n\nshow_full_frame_context_with_heatmap(\n    best_model,\n    val_dataset,\n    val_frames_df,\n    indices=selected_indices,\n    max_examples=4,\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T05:19:33.961426Z","iopub.execute_input":"2026-04-16T05:19:33.961709Z","iopub.status.idle":"2026-04-16T05:20:47.623306Z","shell.execute_reply.started":"2026-04-16T05:19:33.961686Z","shell.execute_reply":"2026-04-16T05:20:47.622313Z"}},"outputs":[],"execution_count":null}]}