{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":924245,"sourceType":"datasetVersion","datasetId":464091}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>1. Import Library</b></div>","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport random\nimport time\nimport numpy as np\nimport pandas as pd\nimport copy\nfrom collections import defaultdict\n\nfrom glob import glob\nfrom tqdm import tqdm\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\n\nimport PIL\nfrom PIL import Image\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.fft\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda import amp\nimport torch.nn.functional as F\nfrom torchmetrics import MetricCollection, Accuracy, AUROC, Precision, Recall, F1Score\n\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.model_selection import StratifiedGroupKFold, train_test_split\n\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay, roc_curve, auc\n\nfrom colorama import Fore, Style\nc_ = Fore.BLUE\ns_ = Style.BRIGHT\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:35.089168Z","iopub.execute_input":"2025-10-04T07:55:35.089375Z","iopub.status.idle":"2025-10-04T07:55:52.052432Z","shell.execute_reply.started":"2025-10-04T07:55:35.089357Z","shell.execute_reply":"2025-10-04T07:55:52.051515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"{c_}{s_}Torch Version is {torch.__version__}\\n\")\nprint(f\"{c_}{s_}Timm Version is {timm.__version__}\\n\")\nprint(f\"{c_}{s_}Albumentation Version is {A.__version__}\\n\")\nprint(f\"{c_}{s_}Numpy Version is {np.__version__}\\n\")\nprint(f\"{c_}{s_}Pillow Version is {PIL.__version__}\\n\")\nprint(f\"{c_}{s_}OpenCV Version is {cv2.__version__}\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:52.053467Z","iopub.execute_input":"2025-10-04T07:55:52.053818Z","iopub.status.idle":"2025-10-04T07:55:52.059250Z","shell.execute_reply.started":"2025-10-04T07:55:52.053796Z","shell.execute_reply":"2025-10-04T07:55:52.058539Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Login W&B(Weighted and Bias)**","metadata":{}},{"cell_type":"code","source":"import wandb\n\n# wandb.init: 프로젝트 시작 \n# wandb.watch: 모델 gradient, parameter 실시간으로 저장\n# wandb.log: wandb dashboard 내 저장 \n# wandb.save: wandb 내 artifact 저장 \n# wandb.finish: 프로젝트 종료\n\ntry:\n    wandb.login(key=\"31788e8973e0fb7be0979e807096a77c6a31e90e\")\n    anonymous = None\nexcept:\n    anonymous = \"must\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:52.061004Z","iopub.execute_input":"2025-10-04T07:55:52.061296Z","iopub.status.idle":"2025-10-04T07:55:57.783916Z","shell.execute_reply.started":"2025-10-04T07:55:52.061268Z","shell.execute_reply":"2025-10-04T07:55:57.783096Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Global Settings**","metadata":{}},{"cell_type":"code","source":"class CFG:\n  seed = 2025\n  debug = False # if False, Full Training\n  project_name = 'DFDC-DeepFakeClassifier-Attention Expriment'\n  backbone = 'tf_efficientnet_b5.ns_jft_in1k' # Using Noisy Student Network\n  img_size = [380,380]\n  mean =[0.485, 0.456, 0.406] #[0.485, 0.456, 0.406]\n  std = [0.229, 0.224, 0.225] #[0.229, 0.224, 0.225]\n  dataset = 'DFDC' # Meta - DeepFake Detection Challenge\n  train_bs = 16\n  valid_bs = train_bs * 2\n  n_accumulate = 1\n  lr = 1e-3\n  min_lr = lr * 1e-3\n  epochs = 10\n  weight_decay = 1e-4\n  scheduler = 'CosineAnnealingLR'\n  loss = 'BCEWithLogitsLoss'\n\n  device = \"cuda\" if torch.cuda.is_available() else 'cpu'\n  device_cnt = torch.cuda.device_count()\n\nCFG.model_name = f\"dim-{CFG.img_size[0]}x{CFG.img_size[1]}-model-{CFG.backbone}-FCA\"\n\n\nprint(f\"{c_}{s_} => Device is {CFG.device}\")\nprint(f\"{c_}{s_} => Num of GPU is {CFG.device_cnt}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:57.784802Z","iopub.execute_input":"2025-10-04T07:55:57.785252Z","iopub.status.idle":"2025-10-04T07:55:57.907803Z","shell.execute_reply.started":"2025-10-04T07:55:57.785234Z","shell.execute_reply":"2025-10-04T07:55:57.907215Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Reproducibility**","metadata":{}},{"cell_type":"code","source":"def seed_everything(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    torch.manual_seed(SEED)\n    torch.cuda.manual_seed(SEED)\n\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n\nseed_everything(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:57.908463Z","iopub.execute_input":"2025-10-04T07:55:57.908724Z","iopub.status.idle":"2025-10-04T07:55:57.918510Z","shell.execute_reply.started":"2025-10-04T07:55:57.908694Z","shell.execute_reply":"2025-10-04T07:55:57.917857Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>2. Load Face Crop Dataset</b></div>\n\n| DFDC    | Train  | Valid | Test | Total  |\n| ------- | ------ | ----- | ---- | ------ |\n| Real   | 19154  | 2000  | 2500 | 23654  |\n| Fake | 100000 | 2000  | 2500 | 104500 |\n\nBefore tackling the detection of fake or real videos, we will first construct a robust pipeline for detecting fake or real images.\n\nFor this purpose, we utilize a face-cropped dataset available on Kaggle, which contains approximately 95,634 images. The dataset is split with 80% used for training, and each image is standardized to a resolution of 224×224 pixels to facilitate effective model training.\n\n[DeepFake Cropped Face Dataset](https://www.kaggle.com/datasets/dagnelies/deepfake-faces)\n","metadata":{}},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/deepfake-faces'\n\nIMAGE_PATH = os.path.join(ROOT_DIR, 'faces_224')\nMETA_PATH = os.path.join(ROOT_DIR, 'metadata.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:57.919307Z","iopub.execute_input":"2025-10-04T07:55:57.919580Z","iopub.status.idle":"2025-10-04T07:55:57.924063Z","shell.execute_reply.started":"2025-10-04T07:55:57.919557Z","shell.execute_reply":"2025-10-04T07:55:57.923199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"meta_df = pd.read_csv(META_PATH)\nimg_list = {os.path.basename(x).split('.')[0] + '.mp4': x for x in glob(f'{IMAGE_PATH}/*')}\nmeta_df['image_path'] = meta_df['videoname'].map(img_list)\nmeta_df['original'] = meta_df['original'].fillna(meta_df['videoname'])\nprint(f\"Shape of Meta DataFrame: {meta_df.shape}\")\nprint(display(meta_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:57.924929Z","iopub.execute_input":"2025-10-04T07:55:57.925197Z","iopub.status.idle":"2025-10-04T07:55:59.368379Z","shell.execute_reply.started":"2025-10-04T07:55:57.925175Z","shell.execute_reply":"2025-10-04T07:55:59.367733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trace = go.Pie(\n    labels = ['FAKE', 'REAL'],\n    values = meta_df['label'].value_counts().sort_index(),\n    hole = 0.5,\n    textinfo = 'label+percent',\n    marker = dict(colors=['rgba(255,0,0,0.5)', 'rgba(0,0,255,0.5)']),\n    showlegend=True,\n)\n\nlayout = go.Layout(\n    title='Distribution of FAKE or REAL Image',\n    legend = dict(title='label', bordercolor='black', borderwidth=1)\n)\n\nfig = go.Figure(data=[trace], layout=layout)\nfig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:59.369273Z","iopub.execute_input":"2025-10-04T07:55:59.369568Z","iopub.status.idle":"2025-10-04T07:55:59.607287Z","shell.execute_reply.started":"2025-10-04T07:55:59.369547Z","shell.execute_reply":"2025-10-04T07:55:59.606712Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Both the Face-Cropped Image Dataset and the original DFDC Video Dataset suffer from a severe class imbalance(Real 17% vs Fake 83%). So we need to design training strategies that mitigate this bias-leveraging balanced sampling, loss reweightning, and advanced augmentation.\n\n다만, 비율 차이가 클 경우 한쪽으로 치우쳐져 학습할 가능성이 높지만, 항상 그런거 아님. 만약 fake랑 real image의 특징이 차이가 많이 나며.. 분류 성능이 높음.","metadata":{}},{"cell_type":"code","source":"\ntrace1 = go.Histogram(\n    x = meta_df[meta_df['label']=='FAKE']['original_width'],\n    name = 'FAKE',\n    marker = dict(color='rgba(255,0,0,0.5)'),\n    xbins = dict(size=20),\n    showlegend=True,\n)\ntrace2 = go.Histogram(\n    x = meta_df[meta_df['label']=='REAL']['original_width'],\n    name = 'REAL',\n    marker = dict(color='rgba(0,0,255,0.5)'),\n    xbins = dict(size=20),\n    showlegend=True,\n)\n\ntrace3 = go.Histogram(\n    x = meta_df[meta_df['label']=='FAKE']['original_height'],\n    name = 'FAKE',\n    marker = dict(color='rgba(255,0,0,0.5)'),\n    xbins = dict(size=20),\n    showlegend=False,\n)\ntrace4 = go.Histogram(\n    x = meta_df[meta_df['label']=='REAL']['original_height'],\n    name = 'REAL',\n    marker = dict(color='rgba(0,0,255,0.5)'),\n    xbins = dict(size=20),\n    showlegend=False,\n)\n\nfig = make_subplots(rows=1, cols=2, subplot_titles=(\"Width\", \"Height\"))\nfig.add_traces([trace1, trace2], 1, 1)\nfig.add_traces([trace3, trace4], 1, 2)\n\nfig.update_layout(\n    title='Distribution of Original Width & Orignal',\n\n)\n\nfig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:59.610132Z","iopub.execute_input":"2025-10-04T07:55:59.610350Z","iopub.status.idle":"2025-10-04T07:55:59.972189Z","shell.execute_reply.started":"2025-10-04T07:55:59.610334Z","shell.execute_reply":"2025-10-04T07:55:59.971522Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"When training a classification model, it is necessary to standardize image dimensions. However, as shown in the graph below, the dataset exhibits a wide variation in image sizes. Simply resizing all images to a fixed size may lead to distortion and loss of critical visual information. Therefore, it is preferable to resize images while preserving their original aspect ratio.","metadata":{}},{"cell_type":"code","source":"# from scipy.stats import gaussian_kde\n\n# fake_tmp = meta_df[meta_df['label']=='FAKE']\n# true_tmp = meta_df[meta_df['label']=='REAL']\n\n# fake_xy = np.vstack([fake_tmp['original_width'],fake_tmp['original_height']])\n# fake_z = gaussian_kde(fake_xy)(fake_xy)\n# true_xy = np.vstack([true_tmp['original_width'],true_tmp['original_height']])\n# true_z = gaussian_kde(true_xy)(true_xy)\n\n\n# trace1 = go.Scatter(\n#     x = fake_tmp['original_width'],\n#     y = fake_tmp['original_height'],\n#     mode = 'markers',\n#     name = 'FAKE',\n#     marker=dict(\n#         color=fake_z,\n#         colorscale='coolwarm',\n#         showscale=False,\n#         size=8,\n#         opacity=0.7\n#     )\n# )\n\n# trace2 = go.Scatter(\n#     x = true_tmp['original_width'],\n#     y = true_tmp['original_height'],\n#     mode = 'markers',\n#     name = 'TRUE',\n#     marker=dict(\n#         color=true_z,\n#         colorscale='coolwarm',\n#         showscale=False,\n#         size=8,\n#         opacity=0.7\n#     )\n# )\n\n# fig = make_subplots(rows=1, cols=2, subplot_titles=[\"FAKE Image\", \"TRUE Image\"])\n# fig.add_trace(trace1, 1, 1)\n# fig.add_trace(trace2, 1, 2)\n\n# fig.update_layout(\n#     title = 'Distribution of Aspect Ratio',\n#     xaxis1 = dict(title='Width', ticklen = 5),\n#     yaxis1 = dict(title='Height', ticklen = 5),\n#     xaxis2 = dict(title='Width', ticklen = 5),\n#     yaxis2 = dict(title='Height', ticklen = 5),\n# )\n\n\n# fig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:59.972866Z","iopub.execute_input":"2025-10-04T07:55:59.973063Z","iopub.status.idle":"2025-10-04T07:55:59.977768Z","shell.execute_reply.started":"2025-10-04T07:55:59.973048Z","shell.execute_reply":"2025-10-04T07:55:59.976836Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>3. Build Dataset and DataLoader</b></div>","metadata":{}},{"cell_type":"code","source":"## Label Encoding\n## REAL: 0, FAKE: 1\n\nmeta_df['label'] = meta_df['label'].apply(lambda x: 1 if x == 'FAKE' else 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:55:59.978620Z","iopub.execute_input":"2025-10-04T07:55:59.978877Z","iopub.status.idle":"2025-10-04T07:56:00.026547Z","shell.execute_reply.started":"2025-10-04T07:55:59.978853Z","shell.execute_reply":"2025-10-04T07:56:00.026049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## 다양한 Augmentation이 필요\n## Albumentation에서 제공해주는 Weak Augmentation\n## MixUp, CutMix와 같은 Strong Augmentation\n## Test Time Augmentation(ex: HFlip)\n\n\ndef get_train_transform():\n    return A.Compose([\n        A.Resize(*CFG.img_size),\n        A.HorizontalFlip(p=0.5),\n    ])\n\ndef get_valid_transform():\n    return A.Compose([\n        A.Resize(*CFG.img_size),\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.027198Z","iopub.execute_input":"2025-10-04T07:56:00.027470Z","iopub.status.idle":"2025-10-04T07:56:00.031721Z","shell.execute_reply.started":"2025-10-04T07:56:00.027446Z","shell.execute_reply":"2025-10-04T07:56:00.030996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DFDCDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n\n        row = self.df.iloc[index]\n        img_path = row['image_path']\n        label = row['label']\n\n        # Load Image\n        img = Image.open(img_path).convert(\"RGB\")\n        img = np.array(img, dtype=np.float32)\n        \n        # Augmentation\n        if self.transforms:\n            img = self.transforms(image=img)['image']\n\n        # Normalization\n        img = (img - CFG.mean) / CFG.std\n\n        # Numpy -> Torch \n        img = torch.tensor(img, dtype=torch.float32).permute(2,0,1)\n        label = torch.tensor(label, dtype=torch.float32)\n\n        return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.032407Z","iopub.execute_input":"2025-10-04T07:56:00.032571Z","iopub.status.idle":"2025-10-04T07:56:00.042820Z","shell.execute_reply.started":"2025-10-04T07:56:00.032558Z","shell.execute_reply":"2025-10-04T07:56:00.042072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def split_data(meta_df, seed=CFG.seed, debug=CFG.debug, test_size=0.2, debug_ratio=0.1):\n\n    df_real = meta_df[meta_df['label'] == 0].reset_index(drop=True)\n    df_fake = meta_df[meta_df['label'] == 1].reset_index(drop=True)\n\n    real_train, real_val = train_test_split(\n        df_real,\n        test_size=test_size,\n        random_state=seed,\n        shuffle=True\n    )\n\n    fake_train = df_fake[df_fake['original'].isin(real_train['original'].unique())].reset_index(drop=True)\n    fake_val = df_fake[df_fake['original'].isin(real_val['original'].unique())].reset_index(drop=True)\n\n    fake_val = fake_val.sample(len(real_val), random_state=seed)\n\n    train_df = pd.concat([real_train, fake_train], axis=0)\n    valid_df = pd.concat([real_val, fake_val], axis=0)\n\n    if debug:\n        train_df = train_df.sample(len(train_df) // int(1/debug_ratio), random_state=seed)\n        valid_df = valid_df.sample(len(valid_df) // int(1/debug_ratio), random_state=seed)\n\n    train_df = train_df.sample(frac=1, random_state=seed).reset_index(drop=True)\n    valid_df = valid_df.sample(frac=1, random_state=seed).reset_index(drop=True)\n\n    return train_df, valid_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.043730Z","iopub.execute_input":"2025-10-04T07:56:00.044043Z","iopub.status.idle":"2025-10-04T07:56:00.055731Z","shell.execute_reply.started":"2025-10-04T07:56:00.044004Z","shell.execute_reply":"2025-10-04T07:56:00.055114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, valid_df = split_data(meta_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.056381Z","iopub.execute_input":"2025-10-04T07:56:00.056583Z","iopub.status.idle":"2025-10-04T07:56:00.137867Z","shell.execute_reply.started":"2025-10-04T07:56:00.056567Z","shell.execute_reply":"2025-10-04T07:56:00.137296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trace1 = go.Pie(\n        labels=['REAL', 'FAKE'],\n        values=train_df['label'].value_counts().sort_index(),\n        hole=0.5,\n        textinfo='label+percent',\n        marker=dict(colors=['rgba(0,0,255,0.5)', 'rgba(255,0,0,0.5)']),\n        showlegend=True\n    )\n\ntrace2 = go.Pie(\n        labels=['REAL', 'FAKE'],\n        values=valid_df['label'].value_counts().sort_index(),\n        hole=0.5,\n        textinfo='label+percent',\n        marker=dict(colors=['rgba(0,0,255,0.5)', 'rgba(255,0,0,0.5)']),\n        showlegend=False\n    )\n\n\nfig = make_subplots(\n    rows=1,\n    cols=2,\n    subplot_titles=[\"Train\", \"Valid\"],\n    specs=[[{\"type\": \"domain\"} for _ in range(2)]]\n)\n\n\nfig.add_trace(trace1, 1, 1); fig.add_trace(trace2, 1, 2)\n\nfig.update_layout(\n    title='Distribution of Fake & Real images',\n    legend = dict(title='label', bordercolor='black', borderwidth=1)\n)\nfig.show(renderer=\"iframe\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.138589Z","iopub.execute_input":"2025-10-04T07:56:00.138832Z","iopub.status.idle":"2025-10-04T07:56:00.217789Z","shell.execute_reply.started":"2025-10-04T07:56:00.138815Z","shell.execute_reply":"2025-10-04T07:56:00.217091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>4. Build CNN Model</b></div>\n\n| Column           | Value                      |\n| ---------------- | -------------------------- |\n| Classification   | Binary-Classification      |\n| Model            | EfficientNetB5_ns               |\n| Image Resolution | 380 X 380      |\n| Data             | Face Cropped Image with 80% training set |\n| Loss Function    | Binary Cross Entropy       |","metadata":{}},{"cell_type":"code","source":"## Model Detail\n\n# Dataset: ILSVRC/ImageNet-1k\n# Self-Training with Noisy Student\n\n# EfficientNetB0_ns: 5.3M params, (224x224)\n# EfficientNetB1_ns: 7.8M params, (240X240)\n# EfficientNetB2_ns: 9.1M params, (260X260)\n# EfficientNetB3_ns: 12.2M params, (300X300)\n# EfficientNetB4_ns: 19.3M params, (380X380)\n# EfficientNetB5_ns: 30.4M params, (456X456)\n# EfficientNetB6_ns: 43.0M params, (528X528)\n# EfficientNetB7_ns: 66.3M params, (600X600)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.218431Z","iopub.execute_input":"2025-10-04T07:56:00.218650Z","iopub.status.idle":"2025-10-04T07:56:00.222162Z","shell.execute_reply.started":"2025-10-04T07:56:00.218636Z","shell.execute_reply":"2025-10-04T07:56:00.221379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = timm.create_model(CFG.backbone, pretrained=True)\ntimm.data.resolve_model_data_config(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:00.223202Z","iopub.execute_input":"2025-10-04T07:56:00.223503Z","iopub.status.idle":"2025-10-04T07:56:02.790380Z","shell.execute_reply.started":"2025-10-04T07:56:00.223485Z","shell.execute_reply":"2025-10-04T07:56:02.789731Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git clone https://github.com/HanMoonSub/DeepGuard.git\nsys.path.append(\"DeepGuard\")\nfrom pooling import GeM\nfrom attention import *","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:02.791096Z","iopub.execute_input":"2025-10-04T07:56:02.791307Z","iopub.status.idle":"2025-10-04T07:56:03.941027Z","shell.execute_reply.started":"2025-10-04T07:56:02.791290Z","shell.execute_reply":"2025-10-04T07:56:03.940070Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Attention Module","metadata":{}},{"cell_type":"code","source":"class ECA(nn.Module):\n    def __init__(self, in_planes, gamma=2, b=1):\n        super(ECA, self).__init__()\n        t = int(abs((math.log2(in_planes) / gamma) + b))\n        k = t if t % 2 else t + 1  \n\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        \n        self.conv = nn.Conv1d(2, 1, kernel_size=k, padding=(k-1)//2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        B, C, _, _ = x.size()\n        \n        avg_out = self.avg_pool(x).squeeze(-1).transpose(-1, -2)  # (B, 1, C)\n        max_out = self.max_pool(x).squeeze(-1).transpose(-1, -2)  # (B, 1, C)\n        \n        out = torch.cat([avg_out, max_out], dim=1)\n\n        out = self.conv(out)   # (B, 1, C)\n        out = self.sigmoid(out).transpose(-1, -2).unsqueeze(-1)  # (B, C, 1, 1)\n\n        return x * out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:03.942234Z","iopub.execute_input":"2025-10-04T07:56:03.942505Z","iopub.status.idle":"2025-10-04T07:56:03.949593Z","shell.execute_reply.started":"2025-10-04T07:56:03.942482Z","shell.execute_reply":"2025-10-04T07:56:03.948651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AttentionGate(nn.Module):\n    def __init__(self):\n        super(AttentionGate, self).__init__()\n   \n        self.conv = nn.Sequential(\n            nn.Conv2d(2,1, kernel_size=3, stride=1, padding=\"same\", bias=False),\n            nn.BatchNorm2d(1, eps=1e-5, momentum=0.01, affine=True),\n            nn.ReLU(inplace=True),\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n\n        max_out, _ = torch.max(x, dim=1, keepdim=True) # (B,1,H,W)\n        avg_out = torch.mean(x, dim=1, keepdim=True) # (B,1,H,W)\n\n        cat = torch.cat([max_out, avg_out], dim=1)\n\n        x_out = self.conv(cat)\n\n        return x * self.sigmoid(x_out)\n\nclass TripletAttention(nn.Module):\n    def __init__(self):\n        super(TripletAttention, self).__init__()\n        \n        self.cw = AttentionGate() # Channel, Width\n        self.hc = AttentionGate() # Height, Channel\n        self.hw = AttentionGate() # Height, Width\n\n    def forward(self, x):\n        # x: (Batch, Channel, Height, Width)\n        \n        # Channel, Width\n        x_perm1 = x.permute(0,2,1,3).contiguous()\n        x_out1 = self.cw(x_perm1)\n        x_out1 = x_out1.permute(0,2,1,3).contiguous()\n\n        # Height, Channel\n        x_perm2 = x.permute(0,3,2,1).contiguous()\n        x_out2 = self.hc(x_perm2)\n        x_out2 = x_out2.permute(0,3,2,1).contiguous()\n\n        # Height, Width\n        # x_out = self.hw(x)\n        # x_out = (x_out + x_out1 + x_out2) / 3\n\n        \n\n        return (x_out1 + x_out2) / 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:03.950449Z","iopub.execute_input":"2025-10-04T07:56:03.950709Z","iopub.status.idle":"2025-10-04T07:56:03.965885Z","shell.execute_reply.started":"2025-10-04T07:56:03.950670Z","shell.execute_reply":"2025-10-04T07:56:03.965126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FCA(nn.Module):\n\n    def __init__(self, channel, reduction=16, top_k=16):\n        super(FCA, self).__init__()\n        self.channel = channel\n        self.reduction = reduction\n        self.top_k = top_k\n\n        self.fc = nn.Sequential(\n            nn.Linear(channel, channel // reduction, bias = False),\n            nn.ReLU(inplace=True),\n            nn.Linear(channel // reduction, channel, bias = False),\n            nn.Sigmoid()\n        )\n\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        \n        x_float = x.float()\n\n        freq_map = torch.fft.fft2(x_float, norm='ortho')\n        freq_map = torch.abs(freq_map)\n\n        mag, idx = torch.topk(freq_map.view(B, C, -1), self.top_k, dim=2)\n        descriptor = mag.mean(dim=2)\n\n        weights = self.fc(descriptor).view(B, C, 1, 1)\n\n        weights = weights.to(x.dtype)\n\n        return x * weights\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:03.966729Z","iopub.execute_input":"2025-10-04T07:56:03.966987Z","iopub.status.idle":"2025-10-04T07:56:03.980498Z","shell.execute_reply.started":"2025-10-04T07:56:03.966966Z","shell.execute_reply":"2025-10-04T07:56:03.979825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DeepfakeCNN(nn.Module):\n      def __init__(self, encoder_name, pretrained=True):\n          super(DeepfakeCNN, self).__init__()\n          self.feature_extractor = timm.create_model(encoder_name, pretrained=pretrained, features_only=True)\n          \n          dummy_input = torch.randn(1, 3, *CFG.img_size)\n          with torch.no_grad():\n              dummy_out = self.feature_extractor(dummy_input)\n\n          feature_dim = dummy_out[-1].size(1)\n\n          self.fca = FCA(feature_dim)\n          \n          self.classifier = nn.Linear(feature_dim, 1)\n\n      def forward(self, x):\n\n          ## Extract CNN Features\n          out = self.feature_extractor(x)[-1]\n\n          ## Squeeze and Excitation\n          out = self.fca(out)\n\n          ## Global Average Pooling\n          out = F.adaptive_avg_pool2d(out, 1)\n          \n          ## Reshape (B,C,1,1) -> (B,C)\n          out = out.view(out.size(0),-1)\n\n          ## Classifier\n          out = self.classifier(out)\n\n          return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.000869Z","iopub.status.idle":"2025-10-04T07:56:04.001120Z","shell.execute_reply.started":"2025-10-04T07:56:04.000996Z","shell.execute_reply":"2025-10-04T07:56:04.001007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model():\n    model = DeepfakeCNN(encoder_name=CFG.backbone, pretrained=True)\n\n    if torch.cuda.device_count() > 1:\n        print(f\"Using {torch.cuda.device_count()} Gpus\")\n        model = nn.DataParallel(model)\n\n    model.to(CFG.device)\n    \n    return model\n        \ndef load_model(path):\n\n    model = build_model()\n\n    if path is not None:\n        try:\n            state_dict = torch.load(path, map_location=CFG.device)\n            model.load_state_dict(state_dict)\n        except Exception as e:\n            print(f\"Error loading model weights form {path}: {e}\")\n    \n    model.eval()\n    \n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.001987Z","iopub.status.idle":"2025-10-04T07:56:04.002193Z","shell.execute_reply.started":"2025-10-04T07:56:04.002092Z","shell.execute_reply":"2025-10-04T07:56:04.002101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q torchinfo\nfrom torchinfo import summary\n\nsummary(\n    build_model(),\n    input_size=(1,3, *CFG.img_size),\n    col_names=[\"input_size\", \"output_size\", \"num_params\", \"mult_adds\"],\n    depth=2,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.003270Z","iopub.status.idle":"2025-10-04T07:56:04.003558Z","shell.execute_reply.started":"2025-10-04T07:56:04.003424Z","shell.execute_reply":"2025-10-04T07:56:04.003436Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T-SNE Visualization(Untrained)","metadata":{}},{"cell_type":"code","source":"import cudf, cuml, cupy\nfrom cuml.manifold import TSNE","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.005518Z","iopub.status.idle":"2025-10-04T07:56:04.005847Z","shell.execute_reply.started":"2025-10-04T07:56:04.005728Z","shell.execute_reply":"2025-10-04T07:56:04.005741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_, valid_df = split_data(meta_df, debug=False)\nprint(f\"Total Images are {len(valid_df)}\")\n\nds = DFDCDataset(valid_df, transforms=get_valid_transform())\ndata_loader = DataLoader(ds, shuffle=False, batch_size=CFG.valid_bs,\n                              num_workers=4, pin_memory=True, drop_last = False)\n\nmodel = build_model(); model.eval()\n\n\nprediction = []\n\nwith torch.no_grad():\n    for idx, batch in tqdm(enumerate(data_loader), total=len(data_loader)):\n        imgs, _ = batch\n        imgs = imgs.to(CFG.device)\n\n        backbone = model.module if isinstance(model, nn.DataParallel) else model\n\n        outs = backbone.feature_extractor(imgs)[-1]\n        outs = backbone.fca(outs)\n        outs = F.adaptive_avg_pool2d(outs, 1)\n        preds = outs.view(outs.size(0),-1)\n        \n        prediction.append(preds.cpu().numpy())      \n\nprediction = np.concatenate(prediction, axis=0)\ntmp2 = pd.DataFrame(prediction, columns=[f'embed_{i}' for i in range(prediction.shape[1])])\ntmp2['label'] = valid_df['label']\n\ntsne = TSNE(n_components=2, random_state=CFG.seed)\n\ncudf_tmp = cudf.from_pandas(tmp2[tmp2.columns[:-1]])\nembed_2d = tsne.fit_transform(cudf_tmp).to_numpy()\nembed_df = pd.DataFrame(embed_2d, columns=['x','y'])\n\nembed_df['label'] = tmp2['label']\n\nfig = go.Figure()\n\nfor label in embed_df['label'].unique():\n    temp = embed_df[embed_df['label'] == label]\n\n    fig.add_trace(\n        go.Scatter(\n            x = temp['x'],\n            y = temp['y'],\n            mode='markers',\n            marker=dict(size=6, opacity=0.7),\n            name=\"Fake\" if label == 1 else \"REAL\",\n        )\n    )\n\nfig.update_layout(\n    title=f'T-SNE Visualization with Untrained {CFG.backbone}',\n    legend=dict(title='Label', bordercolor='black', borderwidth=1),\n    xaxis1 = dict(title='t-SNE dimension 1', ticklen = 5),\n    yaxis1 = dict(title='t-SNE dimension 2', ticklen = 5),\n    \n)\n\nfig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.006701Z","iopub.status.idle":"2025-10-04T07:56:04.007009Z","shell.execute_reply.started":"2025-10-04T07:56:04.006842Z","shell.execute_reply":"2025-10-04T07:56:04.006858Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>5. Train CNN Model</b></div>","metadata":{}},{"cell_type":"code","source":"class AverageMeter(object):\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.sum = 0\n        self.avg = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.007875Z","iopub.status.idle":"2025-10-04T07:56:04.008230Z","shell.execute_reply.started":"2025-10-04T07:56:04.008077Z","shell.execute_reply":"2025-10-04T07:56:04.008092Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Class Imbalance Handling","metadata":{}},{"cell_type":"code","source":"def class_imbalance_handle(real_df, fake_df, k):\n    \n    n_real = len(real_df)\n    start = k * n_real\n    end = start + n_real\n\n    if end > len(fake_df):\n        remaining = fake_df[start:]\n        n_remaining = len(remaining)\n        extra = fake_df[:start].sample(n=n_real - n_remaining, replace=False, random_state=CFG.seed)\n        fake_slice = pd.concat([remaining, extra], axis=0).reset_index(drop=True)\n\n    else:\n        fake_slice = fake_df[start:end].reset_index(drop=True)\n\n    train_df = pd.concat([real_df, fake_slice], axis=0).sample(frac=1, random_state=CFG.seed).reset_index(drop=True)\n    return train_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.009464Z","iopub.status.idle":"2025-10-04T07:56:04.009823Z","shell.execute_reply.started":"2025-10-04T07:56:04.009587Z","shell.execute_reply":"2025-10-04T07:56:04.009603Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### DeepFake Trainer","metadata":{}},{"cell_type":"code","source":"class DeepfakeTrainer:\n    def __init__(self, model, optimizer, scheduler):\n\n        self.model = model\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.history = defaultdict(list)\n\n        self.train_epoch_loss = AverageMeter()\n        self.valid_epoch_loss = AverageMeter()\n\n        self.train_metrics, self.valid_metrics = self.build_metrics()\n\n        self.best_loss = 10**5\n        self.best_epoch = -1\n\n\n        self.criterion = nn.BCEWithLogitsLoss()\n\n        # To automcally log graidents\n        wandb.watch(self.model, log='all', log_freq = 100)\n\n        print(f\"{c_}{s_} ### Trainer Prepared, Device is {CFG.device}\")\n\n    def build_metrics(self):\n        \n        base_metrics = MetricCollection({\n            'acc': Accuracy(task=\"binary\"),\n            'auc': AUROC(task='binary'),\n            'precision': Precision(task='binary'),\n            'recall': Recall(task='binary'),\n            'f1': F1Score(task='binary'),\n        }).to(CFG.device)\n\n        return base_metrics.clone(prefix=\"train_\"), base_metrics.clone(prefix=\"valid_\")\n\n    def fit(self, train_df, valid_df):\n        \n        real_train_df = train_df[train_df['label']==0].reset_index(drop=True)\n        fake_train_df = train_df[train_df['label']==1].reset_index(drop=True)\n\n        n_splits = int(np.ceil(len(fake_train_df) / len(real_train_df)))\n\n        valid_dataset = DFDCDataset(valid_df, transforms=get_valid_transform())\n        valid_loader = DataLoader(\n            valid_dataset, shuffle=False, batch_size=CFG.valid_bs,\n            num_workers=4, pin_memory=True, drop_last=False\n        )\n        \n        for epoch in range(1, CFG.epochs + 1):\n            print(f'{c_}#'*25)\n            print(f'{c_}### Epoch {epoch}/{CFG.epochs}')\n            print(f'{c_}#'*25)\n\n            k = (epoch - 1) % n_splits\n            \n            handled_df = class_imbalance_handle(real_train_df, fake_train_df, k)\n            train_dataset = DFDCDataset(handled_df, transforms=get_train_transform())\n            train_loader = DataLoader(\n                train_dataset, shuffle=True, batch_size=CFG.train_bs,\n                num_workers=4, pin_memory=True, drop_last=True,\n            )\n            \n        \n            train_loss, train_metrics = self.train_one_epoch(train_loader)\n            valid_loss, valid_metrics = self.valid_one_epoch(valid_loader)\n\n            if self.scheduler is not None:\n               self.scheduler.step()\n\n\n            # Log the metrics\n            wandb.log({\n                \"Train Loss\": train_loss,\n                \"Valid Loss\": valid_loss,\n                **train_metrics,  \n                **valid_metrics,\n                \"LR\": self.scheduler.get_last_lr()[0],\n                # self.optimizer.param_groups[0]['lr']\n            })\n\n            print(f'Train Loss: {train_loss:.3f} | Valid Loss: {valid_loss:.3f}')\n\n            ## Save Best CheckPoint\n            if valid_loss <= self.best_loss:\n                print(f\"{c_}{s_}Valid Loss Decreased ({self.best_loss:0.4f} ---> {valid_loss:0.4f})\")\n\n                self.best_loss = valid_loss\n                self.best_epoch = epoch\n        \n                best_model_wts = copy.deepcopy(self.model.state_dict())\n                torch.save(best_model_wts, \"best-checkpoint.bin\")\n                best_artifact = wandb.Artifact(\n                    name = CFG.model_name,\n                    type = 'model',\n                    description = None,\n                    metadata = {\n                        \"epoch\": CFG.epochs, \n                        \"lr\": CFG.lr,\n                        \"min_lr\": CFG.min_lr,\n                        \"weight_decay\": CFG.weight_decay,\n                        \"scheduelr\": CFG.scheduler,\n                        \"loss\": CFG.loss\n                           }\n                )\n                best_artifact.add_file(\"best-checkpoint.bin\")\n                wandb.log_artifact(best_artifact, aliases=[\"best\"])\n                \n        ## Save Last CheckPoint\n        last_model_wts = copy.deepcopy(self.model.state_dict())\n        torch.save(last_model_wts, \"last-checkpoint.bin\")\n\n        ## Save Artifact in W&B\n        last_artifact = wandb.Artifact(\n            name = CFG.model_name,\n            type = 'model',\n            description = None,\n            metadata = {\n                \"epoch\": CFG.epochs, \n                \"lr\": CFG.lr,\n                \"min_lr\": CFG.min_lr,\n                \"weight_decay\": CFG.weight_decay,\n                \"scheduelr\": CFG.scheduler,\n                \"loss\": CFG.loss\n                       }\n        )\n        last_artifact.add_file(\"last-checkpoint.bin\")\n        wandb.log_artifact(last_artifact, aliases=[\"last\"])\n\n        self.model.load_state_dict(torch.load(\"last-checkpoint.bin\"))\n\n        return self.model\n\n\n    def train_one_epoch(self, train_loader):\n        self.model.train()\n\n        self.train_epoch_loss.reset()\n        self.train_metrics.reset()\n\n        scaler = amp.GradScaler()\n\n        pbar = tqdm(enumerate(train_loader), total=len(train_loader), desc='Training')\n\n        for step, (images, labels) in pbar:\n            images = images.to(CFG.device)\n            labels = labels.to(CFG.device)\n\n            batch_size = images.size(0)\n\n            with amp.autocast(enabled=True):\n                 y_pred = self.model(images)\n                 y_pred = y_pred.view(-1)\n                 loss = self.criterion(y_pred, labels)\n                 loss = loss / CFG.n_accumulate\n\n            scaler.scale(loss).backward()\n\n            if (step + 1) % CFG.n_accumulate == 0:\n                scaler.step(self.optimizer)\n                scaler.update()\n                self.optimizer.zero_grad()\n\n            preds = torch.sigmoid(y_pred)\n            self.train_metrics.update(preds, labels.int())\n            self.train_epoch_loss.update(loss.detach().item(), batch_size)\n\n            avg_loss = self.train_epoch_loss.avg\n            avg_metrics = self.train_metrics.compute()\n\n            mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n            pbar.set_postfix(\n                train_loss=f'{avg_loss:.4f}',\n                train_acc=f'{avg_metrics[\"train_acc\"]:.4f}',\n                train_auc=f'{avg_metrics[\"train_auc\"]:.4f}',\n                train_precision=f'{avg_metrics[\"train_precision\"]:.4f}',\n                train_recall=f'{avg_metrics[\"train_recall\"]:.4f}',\n                train_f1=f'{avg_metrics[\"train_f1\"]:.4f}',\n                lr=f'{self.optimizer.param_groups[0][\"lr\"]:.5f}',\n                gpu_mem=f'{mem} GB'\n            )\n        torch.cuda.empty_cache()\n\n        return avg_loss, avg_metrics\n\n    @torch.no_grad()\n    def valid_one_epoch(self, valid_loader):\n        self.model.eval()\n\n        self.valid_epoch_loss.reset()\n        self.valid_metrics.reset()\n\n        pbar = tqdm(enumerate(valid_loader), total=len(valid_loader), desc='Validation')\n\n        for step, (images, labels) in pbar:\n            images = images.to(CFG.device)\n            labels = labels.to(CFG.device)\n            batch_size = images.size(0)\n\n            y_pred = self.model(images)\n            y_pred = y_pred.view(-1)\n            loss = self.criterion(y_pred, labels)\n\n            preds = torch.sigmoid(y_pred)\n            self.valid_metrics.update(preds, labels.int())\n            self.valid_epoch_loss.update(loss.detach().item(), batch_size)\n\n            avg_loss = self.valid_epoch_loss.avg\n            avg_metrics = self.valid_metrics.compute()\n\n            mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n\n            pbar.set_postfix(\n                valid_loss=f'{avg_loss:.4f}',\n                valid_acc=f'{avg_metrics[\"valid_acc\"]:.4f}',\n                valid_auc=f'{avg_metrics[\"valid_auc\"]:.4f}',\n                valid_precision=f'{avg_metrics[\"valid_precision\"]:.4f}',\n                valid_recall=f'{avg_metrics[\"valid_recall\"]:.4f}',\n                valid_f1=f'{avg_metrics[\"valid_f1\"]:.4f}',\n                lr=f'{self.optimizer.param_groups[0][\"lr\"]:.5f}',\n                gpu_mem=f'{mem} GB'\n            )\n\n        torch.cuda.empty_cache()\n\n        return avg_loss, avg_metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.010763Z","iopub.status.idle":"2025-10-04T07:56:04.011066Z","shell.execute_reply.started":"2025-10-04T07:56:04.010908Z","shell.execute_reply":"2025-10-04T07:56:04.010923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fetch_scheduler(scheduler_type, optimizer):\n\n    if scheduler_type == 'CosineAnnealingLR':\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=CFG.epochs, eta_min=CFG.min_lr,\n            last_epoch=-1\n        )\n        return scheduler\n\n    elif scheduler_type == 'StepLR':\n        scheduler = torch.optim.lr_scheduler.StepLR(\n            optimizer, step_size=CFG.epochs//10, gamma=0.1, last_epoch=-1\n        )\n        return scheduler\n\n    elif scheduler_type == None:\n        return None\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.012280Z","iopub.status.idle":"2025-10-04T07:56:04.012593Z","shell.execute_reply.started":"2025-10-04T07:56:04.012443Z","shell.execute_reply":"2025-10-04T07:56:04.012456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython import display as ipd\n\n\nrun = wandb.init(project=CFG.project_name,\n                     config={\n                         \"epoch\": CFG.epochs, \n                         \"lr\": CFG.lr,\n                         \"min_lr\": CFG.min_lr,\n                         \"weight_decay\": CFG.weight_decay,\n                         \"scheduelr\": CFG.scheduler,\n                         \"loss\": CFG.loss,  \n                     },\n                     anonymous=anonymous,\n                     name=CFG.model_name,\n                     group = \"Full Training\" if not CFG.debug else \"Debug\",\n                     )\n\ntrain_df, valid_df = split_data(meta_df)\nmodel = build_model()\noptimizer = optim.Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = fetch_scheduler(CFG.scheduler, optimizer)\n\ntrainer = DeepfakeTrainer(model=model, optimizer=optimizer, scheduler=scheduler)\n\nmodel = trainer.fit(train_df, valid_df)\n\nrun.finish()\ndisplay(ipd.IFrame(run.url, width=1000, height=720))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.013895Z","iopub.status.idle":"2025-10-04T07:56:04.014116Z","shell.execute_reply.started":"2025-10-04T07:56:04.014009Z","shell.execute_reply":"2025-10-04T07:56:04.014018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"valid_ds = DFDCDataset(valid_df, transforms=get_valid_transform())\nvalid_loader = DataLoader(valid_ds, shuffle=False, batch_size=CFG.valid_bs,\n                              num_workers=4, pin_memory=True, drop_last = False)\n\n## Last Model\nmodel = load_model('/kaggle/working/last-checkpoint.bin')\n\nval_pred = []; val_true = []\n\nwith torch.no_grad():\n    for idx, batch in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n        imgs, labels = batch\n        imgs = imgs.to(CFG.device)\n        \n        preds = model(imgs)\n        preds = torch.sigmoid(preds)\n        val_pred.append(preds.cpu().numpy())\n        val_true.append(labels.cpu().numpy())\n\nval_pred = np.concatenate(val_pred, axis=0).squeeze()\nval_true = np.concatenate(val_true, axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.014989Z","iopub.status.idle":"2025-10-04T07:56:04.015348Z","shell.execute_reply.started":"2025-10-04T07:56:04.015165Z","shell.execute_reply":"2025-10-04T07:56:04.015181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fpr, tpr, _ = roc_curve(val_true, val_pred)\nroc_auc = auc(fpr, tpr)\n\ncm = confusion_matrix(val_true.astype(int), np.where(val_pred >= 0.5, 1, 0), labels=[0,1], normalize='true')\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=['Real', 'Fake'])\n\nfig,axes = plt.subplots(1,2, figsize=(12,6))\n\ndisp.plot(ax=axes[0], cmap='Blues')\n\naxes[1].plot(fpr, tpr, color=\"darkorange\", lw=2, label=f\"AUC = {roc_auc:.2f}\")\naxes[1].plot([0, 1], [0, 1], color=\"navy\", lw=2, linestyle=\"--\")  # baseline\naxes[1].set_xlim([0.0, 1.0])\naxes[1].set_ylim([0.0, 1.05])\naxes[1].set_xlabel(\"False Positive Rate\")\naxes[1].set_ylabel(\"True Positive Rate\")\naxes[1].set_title(\"ROC Curve\")\naxes[1].legend(loc=\"lower right\")\n\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.016671Z","iopub.status.idle":"2025-10-04T07:56:04.017028Z","shell.execute_reply.started":"2025-10-04T07:56:04.016849Z","shell.execute_reply":"2025-10-04T07:56:04.016867Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T-SNE Visualization with Trained Model","metadata":{}},{"cell_type":"code","source":"_, valid_df = split_data(meta_df, debug=False)\nprint(f\"Total Images are {len(valid_df)}\")\n\nds = DFDCDataset(valid_df, transforms=get_valid_transform())\ndata_loader = DataLoader(ds, shuffle=False, batch_size=CFG.valid_bs,\n                              num_workers=4, pin_memory=True, drop_last = False)\n\nmodel = load_model('/kaggle/working/last-checkpoint.bin')\n\nprediction = []\n\nwith torch.no_grad():\n    for idx, batch in tqdm(enumerate(data_loader), total=len(data_loader)):\n        imgs, _ = batch\n        imgs = imgs.to(CFG.device)\n\n        backbone = model.module if isinstance(model, nn.DataParallel) else model\n\n        outs = backbone.feature_extractor(imgs)[-1]\n        outs = backbone.fca(outs)\n        outs = F.adaptive_avg_pool2d(outs, 1)\n        \n        preds = outs.view(outs.size(0),-1)\n        \n        prediction.append(preds.cpu().numpy())      \n\nprediction = np.concatenate(prediction, axis=0)\ntmp2 = pd.DataFrame(prediction, columns=[f'embed_{i}' for i in range(prediction.shape[1])])\ntmp2['label'] = valid_df['label']\n\ntsne = TSNE(n_components=2, random_state=CFG.seed)\n\ncudf_tmp = cudf.from_pandas(tmp2[tmp2.columns[:-1]])\nembed_2d = tsne.fit_transform(cudf_tmp).to_numpy()\nembed_df = pd.DataFrame(embed_2d, columns=['x','y'])\n\nembed_df['label'] = tmp2['label']\n\nfig = go.Figure()\n\nfor label in embed_df['label'].unique():\n    temp = embed_df[embed_df['label'] == label]\n\n    fig.add_trace(\n        go.Scatter(\n            x = temp['x'],\n            y = temp['y'],\n            mode='markers',\n            marker=dict(size=6, opacity=0.7),\n            name=\"Fake\" if label == 1 else \"REAL\",\n        )\n    )\n\nfig.update_layout(\n    title=f'T-SNE Visualization with Untrained {CFG.backbone}',\n    legend=dict(title='Label', bordercolor='black', borderwidth=1),\n    xaxis1 = dict(title='t-SNE dimension 1', ticklen = 5),\n    yaxis1 = dict(title='t-SNE dimension 2', ticklen = 5),\n    \n)\n\nfig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.018512Z","iopub.status.idle":"2025-10-04T07:56:04.018791Z","shell.execute_reply.started":"2025-10-04T07:56:04.018625Z","shell.execute_reply":"2025-10-04T07:56:04.018634Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>6. Display Cam</b></div>\n\n### Cam: Class Activation Mapping","metadata":{}},{"cell_type":"code","source":"class CAM_FeatureExtractor:\n    def __init__(self, model, model_type='efficientnet', target_layer=None):\n        \"\"\"\n        model: nn.DataParallel\n        model_type: 'efficientnet', 'xception', ...\n        target_layers: list of layer names to extract\n        \"\"\"\n\n        self.model = model.module if isinstance(model, nn.DataParallel) else model\n        self.model_type = model_type\n        self.target_layer = target_layer or self._default_layer() \n        self.feature = None\n        self.hook = None\n        self._register_hook()\n\n    def _default_layer(self):\n        if self.model_type == 'efficientnet':\n            return 'blocks.6'\n        else:\n            raise ValueError(f\"Unsupported model type: {self.model_type}\")\n\n    def _register_hook(self):\n            \n        module = dict(self.model.feature_extractor.named_modules()).get(self.target_layer, None)\n        \n        if module is None:\n            raise ValueError(f\"Layer {self.target_layer} not found in model\")\n        \n        self.hook = module.register_forward_hook(self._hook_fn)\n\n    @torch.no_grad\n    def _hook_fn(self, module, input, output):\n        self.feature = self.model.fca(output).detach()\n\n    def remove_hook(self):\n        \n        if self.hook is not None:\n            self.hook.remove()\n            self.hook = None\n\n    @torch.no_grad\n    def __call__(self, inp):\n\n        pred = self.model(inp)\n        prob = torch.sigmoid(pred.squeeze(0)).item()\n\n        feature = self.feature.squeeze(0).cpu().numpy()\n        fc_weight = self.model.classifier.weight.detach().cpu().numpy()\n\n        return prob, feature, fc_weight\n\n\nclass CAM_FeatureVisualizer:\n    @staticmethod\n    def compute_cam(feature, fc_weight):\n        \"\"\"\n        feature: (C,H,W)\n        fc_weight: (1, C)\n        \n        \"\"\"\n\n        ## Weighted Sum\n        c,h,w = feature.shape\n        cam = fc_weight.squeeze(0).dot(feature.reshape(c, h*w))\n        cam = cam.reshape(h,w)\n        \n        ## Normalization(0~1)\n        cam -= cam.min()\n        cam = cam / (cam.max() + 1e-9)\n\n        return cam\n\n    @staticmethod\n    def overlay_heatmap(cam_resized, original_img, alpha=0.5):\n        heatmap = np.uint8(255 * cam_resized)\n        heatmap = cv2.applyColorMap(heatmap, cv2.COLORMAP_JET)\n        heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n        overlay = cv2.addWeighted(original_img, 1 - alpha, heatmap, alpha, 0)\n        \n        return overlay\n\n    @staticmethod\n    def resize_cam(cam, size):\n        return cv2.resize(cam, (size[1], size[0]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.019609Z","iopub.status.idle":"2025-10-04T07:56:04.019856Z","shell.execute_reply.started":"2025-10-04T07:56:04.019748Z","shell.execute_reply":"2025-10-04T07:56:04.019760Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = load_model(path='../working/best-checkpoint.bin')\n\ncam_extractor = CAM_FeatureExtractor(model, model_type=\"efficientnet\")\ncam_visualizer = CAM_FeatureVisualizer()\n\n_, valid_df = split_data(meta_df, debug=False)\nds = DFDCDataset(valid_df, transforms=get_valid_transform())\ndata_loader = DataLoader(ds, shuffle=False, batch_size=16,\n                         num_workers=4, pin_memory=True, drop_last=False)\n\nBatch_idx = 3\n\nfor idx, batch in enumerate(data_loader, start=1):\n    imgs, labels = batch\n    \n    plt.figure(figsize=(12, 12))\n    print(\"#\"*25); print('\\n')\n    \n    for i, (img, label) in enumerate(zip(imgs, labels)):\n        img_tensor = img.unsqueeze(0).to(CFG.device)\n\n        prob, feature, fc_weight = cam_extractor(img_tensor)\n\n        cam = cam_visualizer.compute_cam(feature, fc_weight)\n        cam_resized = cam_visualizer.resize_cam(cam, CFG.img_size)\n        \n        img_np = img.permute(1, 2, 0).cpu().numpy()\n        img_np -= img_np.min()\n        img_np = (img_np / (img_np.max() + 1e-9) * 255).astype(np.uint8)\n\n        overlay = cam_visualizer.overlay_heatmap(cam_resized, img_np, alpha=0.3)\n\n        plt.subplot(4, 4, i + 1)\n        plt.imshow(overlay)\n        plt.title(f\"GT: {label.item()}\\nPred: {prob:.2f}\", size=12)\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n\n    if Batch_idx == idx:\n        break\n\ncam_extractor.remove_hook()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.021500Z","iopub.status.idle":"2025-10-04T07:56:04.021771Z","shell.execute_reply.started":"2025-10-04T07:56:04.021624Z","shell.execute_reply":"2025-10-04T07:56:04.021633Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"padding: 25px;color:white;margin:10;font-size:60%;text-align:left;display:fill;border-radius:10px;background-color:#FFFFFF;overflow:hidden;background-color:#FFA500\"><b><span style='color:#FFA500'></span></b> <b>7. Display Grad-Cam</b></div>\n\n### Grad-Cam: Gradient Weighted Class Activation Mapping","metadata":{}},{"cell_type":"code","source":"class GradCAM_FeatureExtractor:\n    def __init__(self, model, model_type='efficientnet', target_layer=None):\n        self.model = model.module if isinstance(model, torch.nn.DataParallel) else model\n        self.model_type = model_type\n        self.target_layer = target_layer or self._default_layer()\n\n        self.gradients = None\n        self.activations = None\n\n        self.hook_forward = None\n        self.hook_backward = None\n        \n        self._register_hook() # Initial Regist Hook\n\n    def _default_layer(self):\n        if self.model_type == 'efficientnet':\n            return 'blocks.6'\n        else:\n            raise ValueError(f'Unspported model type: {self.model_type}')\n\n    def _register_hook(self):\n        module = dict(self.model.feature_extractor.named_modules()).get(self.target_layer, None)\n        if module is None:\n            raise ValueError(f\"Layer {self.target_layer} not found in model\")\n\n        self.hook_forward = module.register_forward_hook(self._hook_activation_fn)\n        self.hook_backward = module.register_full_backward_hook(self._hook_gradient_fn)\n\n    def _hook_activation_fn(self, module, input, output):\n        \n        self.activations = self.model.fca(output).detach()\n\n    def _hook_gradient_fn(self, module, grad_input, grad_output):\n        \n        self.gradients = grad_output[0].detach()\n        \n    def remove_hooks(self):\n        if self.hook_forward is not None:\n            self.hook_forward.remove()\n            self.hook_forward = None\n\n        if self.hook_backward is not None:\n            self.hook_backward.remove()\n            self.hook_backward = None\n\n    def __call__(self, inp):\n        \n        self.model.zero_grad()\n\n        pred = self.model(inp)\n        prob = torch.sigmoid(pred.squeeze(0)).item()\n\n        pred.backward()\n        \n        return prob, self.activations.squeeze(0), self.gradients.squeeze(0)\n\nclass GrdaCAM_FeatureVisualizer:\n    @staticmethod \n    def compute_gradcam(activations, gradients):\n        \"\"\"\n        activations: (C, H, W)\n        gradients:   (C, H, W)\n        \"\"\"\n\n        ## weighted sum\n        weights = gradients.mean(dim=[1,2], keepdim=True) # (C,1,1)\n        cam = (weights * activations).sum(dim=0, keepdim=True) # (1,H,W)\n        cam = F.relu(cam)\n        \n        ## Normalize [0,1]\n        cam = cam.squeeze().cpu().numpy()\n        cam -= cam.min()\n        cam /= (cam.max() + 1e-9)\n        return cam\n    \n    @staticmethod\n    def resize_gradcam(cam, size):\n        return cv2.resize(cam, (size[1], size[0]))\n\n    @staticmethod\n    def overlay_heatmap(gradcam_resized, original_img, alpha=0.4):\n        heatmap = np.uint8(255 * gradcam_resized)\n        heatmap = cv2.applyColorMap(heatmap, cv2.COLORMAP_JET)\n        heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n\n        overlay = cv2.addWeighted(original_img, 1 - alpha, heatmap, alpha, 0)\n\n        return overlay","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.023000Z","iopub.status.idle":"2025-10-04T07:56:04.023275Z","shell.execute_reply.started":"2025-10-04T07:56:04.023140Z","shell.execute_reply":"2025-10-04T07:56:04.023154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = load_model(path=\"../working/best-checkpoint.bin\")\ngradcam_extractor = GradCAM_FeatureExtractor(model, model_type=\"efficientnet\", target_layer=\"blocks.6\")\ngradcam_visualizer = GrdaCAM_FeatureVisualizer()\n\n_, valid_df = split_data(meta_df, debug=False)\nds = DFDCDataset(valid_df, transforms=get_valid_transform())\ndata_loader = DataLoader(ds, shuffle=False, batch_size=16,\n                         num_workers=4, pin_memory=True, drop_last=False)\n\nBatch_idx = 3\n\nfor idx, batch in enumerate(data_loader, start=1):\n    imgs, labels = batch\n    \n    plt.figure(figsize=(12, 12))\n    print(\"#\"*25); print('\\n')\n    \n    for i, (img, label) in enumerate(zip(imgs, labels)):\n        img_tensor = img.unsqueeze(0).to(CFG.device)\n\n        prob, activations, gradients = gradcam_extractor(img_tensor)\n\n        gradcam = gradcam_visualizer.compute_gradcam(activations, gradients)\n        gradcam_resized = gradcam_visualizer.resize_gradcam(gradcam, CFG.img_size)\n        \n        img_np = img.permute(1, 2, 0).cpu().numpy()\n        img_np -= img_np.min()\n        img_np = (img_np / (img_np.max() + 1e-9) * 255).astype(np.uint8)\n\n        overlay = gradcam_visualizer.overlay_heatmap(gradcam_resized, img_np, alpha=0.3)\n\n        plt.subplot(4, 4, i + 1)\n        plt.imshow(overlay)\n        plt.title(f\"GT: {label.item()}\\nPred: {prob:.2f}\", size=12)\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n\n    if Batch_idx == idx:\n        break\n\ngradcam_extractor.remove_hooks()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-04T07:56:04.024167Z","iopub.status.idle":"2025-10-04T07:56:04.024535Z","shell.execute_reply.started":"2025-10-04T07:56:04.024354Z","shell.execute_reply":"2025-10-04T07:56:04.024371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}