{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":16880,"datasetId":16880,"databundleId":16880,"isSourceIdPinned":false}],"isInternetEnabled":true,"language":"python","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q tensorboardX scipy scikit-learn opencv-python-headless tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:25.916297Z","iopub.execute_input":"2026-09-22T12:37:25.916689Z","iopub.status.idle":"2026-09-22T12:37:30.486674Z","shell.execute_reply.started":"2026-09-22T12:37:25.916618Z","shell.execute_reply":"2026-09-22T12:37:30.485618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y -q imageio-ffmpeg\n!pip install -q \"imageio-ffmpeg==0.4.5\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:30.48925Z","iopub.execute_input":"2026-09-22T12:37:30.489462Z","iopub.status.idle":"2026-09-22T12:37:37.906477Z","shell.execute_reply.started":"2026-09-22T12:37:30.489427Z","shell.execute_reply":"2026-09-22T12:37:37.905575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import imageio_ffmpeg, os\n\nffmpeg_path = imageio_ffmpeg.get_ffmpeg_exe()\nprint('ffmpeg binary tại:', ffmpeg_path)\n\nos.makedirs('/usr/local/bin', exist_ok=True)\nsymlink_path = '/usr/local/bin/ffmpeg'\nif os.path.exists(symlink_path) or os.path.islink(symlink_path):\n    os.remove(symlink_path)\nos.symlink(ffmpeg_path, symlink_path)\n\nimport subprocess\nresult = subprocess.run(['ffmpeg', '-version'], stdout=subprocess.PIPE, stderr=subprocess.PIPE)\nprint('ffmpeg return code:', result.returncode)\nprint(result.stdout.decode()[:200])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:37.908197Z","iopub.execute_input":"2026-09-22T12:37:37.908665Z","iopub.status.idle":"2026-09-22T12:37:38.229398Z","shell.execute_reply.started":"2026-09-22T12:37:37.908442Z","shell.execute_reply":"2026-09-22T12:37:38.228606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport copy\nimport glob\nimport time\nimport math\nimport random\nimport numbers\nimport pickle\nimport shutil\nimport argparse\nimport re\nimport csv\nimport subprocess\nimport collections.abc as collections\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nplt.switch_backend('agg')\nfrom collections import deque\nfrom datetime import datetime\n\nfrom PIL import ImageOps, Image\nfrom scipy.io import wavfile\nfrom scipy.interpolate import interp1d\nfrom scipy import signal\nfrom sklearn.metrics import roc_auc_score\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.autograd import Variable\nfrom torch.utils import data\nimport torch.utils.data\nimport torchvision\nfrom torchvision import transforms\nimport torchvision.transforms.functional as TF\nfrom tensorboardX import SummaryWriter\nfrom tqdm import tqdm\n\nprint('PyTorch version:', torch.__version__)\nprint('CUDA available:', torch.cuda.is_available())\nif torch.cuda.is_available():\n    print('GPU:', torch.cuda.get_device_name(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:38.23137Z","iopub.execute_input":"2026-09-22T12:37:38.231691Z","iopub.status.idle":"2026-09-22T12:37:39.552724Z","shell.execute_reply.started":"2026-09-22T12:37:38.231622Z","shell.execute_reply":"2026-09-22T12:37:39.551867Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 2. Configuration","metadata":{}},{"cell_type":"code","source":"class Config:\n    # Data paths - Kaggle mounts DFDC data here\n    DFDC_DATA_DIR = '/kaggle/input/competitions/deepfake-detection-challenge'   \n    \n    # Working directory for preprocessed data\n    WORK_DIR = '/kaggle/working'\n    TRAIN_DIR = os.path.join(WORK_DIR, 'train')\n    TEST_DIR = os.path.join(WORK_DIR, 'test')\n    \n    # Training hyperparameters\n    BATCH_SIZE = 16\n    NUM_WORKERS = 4\n    LEARNING_RATE = 1e-4\n    WEIGHT_DECAY = 1e-5\n    EPOCHS = 100\n    IMG_DIM = 224\n    SPATIAL_SIZE = 28\n    NET = 'resnet18'\n    \n    # Model options\n    WITH_ATTENTION = True\n    RESIDUAL_CONN = False\n    USING_PSEUDO_FAKE = True\n    \n    # Pseudo-fake augmentation parameters\n    AUD_MIN_FAKE_LEN = 2\n    AUD_MAX_FAKE_LEN = -1\n    VIS_MIN_FAKE_LEN = 2\n    VIS_MAX_FAKE_LEN = -1\n    \n    # Misc\n    PRINT_FREQ = 10\n    SEED = 0\n    SAVE_ALL = False\n\ncfg = Config()\nprint('Configuration loaded.')\nprint(f'  DFDC Data: {cfg.DFDC_DATA_DIR}')\nprint(f'  Working Dir: {cfg.WORK_DIR}')\nprint(f'  Epochs: {cfg.EPOCHS}, BS: {cfg.BATCH_SIZE}, LR: {cfg.LEARNING_RATE}')\nprint(f'  Attention: {cfg.WITH_ATTENTION}, Pseudo-Fake: {cfg.USING_PSEUDO_FAKE}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.556444Z","iopub.execute_input":"2026-09-22T12:37:39.556664Z","iopub.status.idle":"2026-09-22T12:37:39.564222Z","shell.execute_reply.started":"2026-09-22T12:37:39.556622Z","shell.execute_reply":"2026-09-22T12:37:39.563316Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 3. Utility Functions","metadata":{}},{"cell_type":"code","source":"def min_max_normalize(value, min_value, max_value):\n    return (value - min_value) / (max_value - min_value + 0.00000001)\n\n\ndef save_checkpoint(state, is_best=0, gap=1, filename='models/checkpoint.pth.tar', keep_all=False):\n    torch.save(state, filename)\n    last_epoch_path = os.path.join(os.path.dirname(filename),\n                                   'epoch%s.pth.tar' % str(state['epoch'] - gap))\n    if (state['epoch'] - gap) == 50:\n        os.makedirs(os.path.join(os.path.dirname(filename), 'halfepochs_results'), exist_ok=True)\n        if os.path.exists(last_epoch_path):\n            os.rename(last_epoch_path,\n                      os.path.join(os.path.dirname(filename), 'halfepochs_results',\n                                   'epoch%s.pth.tar' % str(state['epoch'] - gap)))\n        past_best = glob.glob(os.path.join(os.path.dirname(filename), 'model_best_*.pth.tar'))\n        for i in past_best:\n            os.rename(i, os.path.join(os.path.dirname(filename), 'halfepochs_results', os.path.basename(i)))\n    if not keep_all:\n        try:\n            os.remove(last_epoch_path)\n        except:\n            pass\n    if is_best:\n        past_best = glob.glob(os.path.join(os.path.dirname(filename), 'model_best_*.pth.tar'))\n        for i in past_best:\n            try:\n                os.remove(i)\n            except:\n                pass\n        torch.save(state, os.path.join(os.path.dirname(filename),\n                                       'model_best_epoch%s.pth.tar' % str(state['epoch'])))\n\n\ndef write_log(content, epoch, filename):\n    if not os.path.exists(filename):\n        log_file = open(filename, 'w')\n    else:\n        log_file = open(filename, 'a')\n    log_file.write('## Epoch %d:\\n' % epoch)\n    log_file.write('time: %s\\n' % str(datetime.now()))\n    log_file.write(content + '\\n\\n')\n    log_file.close()\n\n\ndef denorm(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]):\n    assert len(mean) == len(std) == 3\n    inv_mean = [-mean[i] / std[i] for i in range(3)]\n    inv_std = [1 / i for i in std]\n    return transforms.Normalize(mean=inv_mean, std=inv_std)\n\n\nclass AverageMeter(object):\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n        self.local_history = deque([])\n        self.local_avg = 0\n        self.history = []\n        self.dict = {}\n        self.save_dict = {}\n\n    def update(self, val, n=1, history=0, step=5):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n        if history:\n            self.history.append(val)\n        if step > 0:\n            self.local_history.append(val)\n            if len(self.local_history) > step:\n                self.local_history.popleft()\n            self.local_avg = np.average(self.local_history)\n\n    def dict_update(self, val, key):\n        if key in self.dict.keys():\n            self.dict[key].append(val)\n        else:\n            self.dict[key] = [val]\n\n    def __len__(self):\n        return self.count\n\nprint('Utility functions defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.566116Z","iopub.execute_input":"2026-09-22T12:37:39.566384Z","iopub.status.idle":"2026-09-22T12:37:39.587421Z","shell.execute_reply.started":"2026-09-22T12:37:39.566324Z","shell.execute_reply":"2026-09-22T12:37:39.586649Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 4. Augmentation Transforms","metadata":{}},{"cell_type":"code","source":"class Padding:\n    def __init__(self, pad):\n        self.pad = pad\n    def __call__(self, img):\n        return ImageOps.expand(img, border=self.pad, fill=0)\n\n\nclass Scale:\n    def __init__(self, size, interpolation=Image.NEAREST):\n        assert isinstance(size, int) or (isinstance(size, collections.Iterable) and len(size) == 2)\n        self.size = size\n        self.interpolation = interpolation\n\n    def __call__(self, imgmap):\n        img1 = imgmap[0]\n        if isinstance(self.size, int):\n            w, h = img1.size\n            if (w <= h and w == self.size) or (h <= w and h == self.size):\n                return imgmap\n            if w < h:\n                ow = self.size\n                oh = int(self.size * h / w)\n                return [i.resize((ow, oh), self.interpolation) for i in imgmap]\n            else:\n                oh = self.size\n                ow = int(self.size * w / h)\n                return [i.resize((ow, oh), self.interpolation) for i in imgmap]\n        else:\n            return [i.resize(self.size, self.interpolation) for i in imgmap]\n\n\nclass CenterCrop:\n    def __init__(self, size, consistent=True):\n        if isinstance(size, numbers.Number):\n            self.size = (int(size), int(size))\n        else:\n            self.size = size\n    def __call__(self, imgmap):\n        img1 = imgmap[0]\n        w, h = img1.size\n        th, tw = self.size\n        x1 = int(round((w - tw) / 2.))\n        y1 = int(round((h - th) / 2.))\n        return [i.crop((x1, y1, x1 + tw, y1 + th)) for i in imgmap]\n\n\nclass RandomHorizontalFlip:\n    def __init__(self, consistent=True, command=None):\n        self.consistent = consistent\n        if command == 'left':\n            self.threshold = 0\n        elif command == 'right':\n            self.threshold = 1\n        else:\n            self.threshold = 0.5\n    def __call__(self, imgmap):\n        if self.consistent:\n            if random.random() < self.threshold:\n                return [i.transpose(Image.FLIP_LEFT_RIGHT) for i in imgmap]\n            else:\n                return imgmap\n        else:\n            result = []\n            for i in imgmap:\n                if random.random() < self.threshold:\n                    result.append(i.transpose(Image.FLIP_LEFT_RIGHT))\n                else:\n                    result.append(i)\n            return result\n\n\nclass ToTensor:\n    def __call__(self, imgmap):\n        totensor = transforms.ToTensor()\n        return [totensor(i) for i in imgmap]\n\n\nclass Normalize:\n    def __init__(self, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]):\n        self.mean = mean\n        self.std = std\n    def __call__(self, imgmap):\n        normalize = transforms.Normalize(mean=self.mean, std=self.std)\n        return [normalize(i) for i in imgmap]\n\nprint('Augmentation transforms defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.5886Z","iopub.execute_input":"2026-09-22T12:37:39.588866Z","iopub.status.idle":"2026-09-22T12:37:39.609439Z","shell.execute_reply.started":"2026-09-22T12:37:39.588804Z","shell.execute_reply":"2026-09-22T12:37:39.608789Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 5. ResNet 2D3D Backbone","metadata":{}},{"cell_type":"code","source":"def conv3x3x3(in_planes, out_planes, stride=1, bias=False):\n    return nn.Conv3d(in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=bias)\n\ndef conv1x3x3(in_planes, out_planes, stride=1, bias=False):\n    return nn.Conv3d(in_planes, out_planes, kernel_size=(1, 3, 3),\n                     stride=(1, stride, stride), padding=(0, 1, 1), bias=bias)\n\n\ndef downsample_basic_block(x, planes, stride):\n    out = F.avg_pool3d(x, kernel_size=1, stride=stride)\n    zero_pads = torch.Tensor(\n        out.size(0), planes - out.size(1), out.size(2), out.size(3),\n        out.size(4)).zero_()\n    if isinstance(out.data, torch.cuda.FloatTensor):\n        zero_pads = zero_pads.cuda()\n    out = Variable(torch.cat([out.data, zero_pads], dim=1))\n    return out\n\n\nclass BasicBlock3d(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None,\n                 track_running_stats=True, use_final_relu=True):\n        super(BasicBlock3d, self).__init__()\n        bias = False\n        self.use_final_relu = use_final_relu\n        self.conv1 = conv3x3x3(inplanes, planes, stride, bias=bias)\n        self.bn1 = nn.BatchNorm3d(planes, track_running_stats=track_running_stats)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3x3(planes, planes, bias=bias)\n        self.bn2 = nn.BatchNorm3d(planes, track_running_stats=track_running_stats)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n        out = self.conv2(out)\n        out = self.bn2(out)\n        if self.downsample is not None:\n            residual = self.downsample(x)\n        out += residual\n        if self.use_final_relu:\n            out = self.relu(out)\n        return out\n\n\nclass BasicBlock2d(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None,\n                 track_running_stats=True, use_final_relu=True):\n        super(BasicBlock2d, self).__init__()\n        bias = False\n        self.use_final_relu = use_final_relu\n        self.conv1 = conv1x3x3(inplanes, planes, stride, bias=bias)\n        self.bn1 = nn.BatchNorm3d(planes, track_running_stats=track_running_stats)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv1x3x3(planes, planes, bias=bias)\n        self.bn2 = nn.BatchNorm3d(planes, track_running_stats=track_running_stats)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n        out = self.conv2(out)\n        out = self.bn2(out)\n        if self.downsample is not None:\n            residual = self.downsample(x)\n        out += residual\n        if self.use_final_relu:\n            out = self.relu(out)\n        return out\n\n\nclass ResNet2d3d_half(nn.Module):\n    def __init__(self, block, layers, track_running_stats=True):\n        super(ResNet2d3d_half, self).__init__()\n        self.inplanes = 64\n        self.track_running_stats = track_running_stats\n        bias = False\n        self.conv1 = nn.Conv3d(3, 64, kernel_size=(1, 7, 7), stride=(1, 2, 2),\n                               padding=(0, 3, 3), bias=bias)\n        self.bn1 = nn.BatchNorm3d(64, track_running_stats=track_running_stats)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool3d(kernel_size=(1, 3, 3), stride=(1, 2, 2), padding=(0, 1, 1))\n\n        if not isinstance(block, list):\n            block = [block] * 4\n\n        self.layer1 = self._make_layer(block[0], 64, layers[0])\n        self.layer2 = self._make_layer(block[1], 128, layers[1], stride=2, is_final=True)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                m.weight = nn.init.kaiming_normal_(m.weight, mode='fan_out')\n                if m.bias is not None:\n                    m.bias.data.zero_()\n            elif isinstance(m, nn.BatchNorm3d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n    def _make_layer(self, block, planes, blocks, stride=1, is_final=False):\n        downsample = None\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            if block == BasicBlock2d:\n                customized_stride = (1, stride, stride)\n            else:\n                customized_stride = stride\n            downsample = nn.Sequential(\n                nn.Conv3d(self.inplanes, planes * block.expansion,\n                          kernel_size=1, stride=customized_stride, bias=False),\n                nn.BatchNorm3d(planes * block.expansion,\n                               track_running_stats=self.track_running_stats)\n            )\n\n        layers = []\n        layers.append(block(self.inplanes, planes, stride, downsample,\n                            track_running_stats=self.track_running_stats))\n        self.inplanes = planes * block.expansion\n        if is_final:\n            for i in range(1, blocks - 1):\n                layers.append(block(self.inplanes, planes,\n                                    track_running_stats=self.track_running_stats))\n            layers.append(block(self.inplanes, planes,\n                                track_running_stats=self.track_running_stats,\n                                use_final_relu=False))\n        else:\n            for i in range(1, blocks):\n                layers.append(block(self.inplanes, planes,\n                                    track_running_stats=self.track_running_stats))\n        return nn.Sequential(*layers)\n\n    def forward(self, x, return_intermediate=False):\n        x0 = self.conv1(x)\n        x0 = self.bn1(x0)\n        x0 = self.relu(x0)\n        x0 = self.maxpool(x0)\n        x1 = self.layer1(x0)\n        x = self.layer2(x1)\n        if return_intermediate:\n            return x, x1\n        else:\n            return x\n\n\ndef resnet18_2d3d_half(**kwargs):\n    model = ResNet2d3d_half([BasicBlock3d, BasicBlock3d], [2, 2], **kwargs)\n    return model\n\n\ndef resnet34_2d3d_half(**kwargs):\n    model = ResNet2d3d_half([BasicBlock3d, BasicBlock3d], [6, 3], **kwargs)\n    return model\n\n\ndef select_resnet_half(network, track_running_stats=True):\n    param = {'feature_size': 512}\n    if network == 'resnet18':\n        model = resnet18_2d3d_half(track_running_stats=track_running_stats)\n        param['feature_size'] = 128\n    elif network == 'resnet34':\n        model = resnet34_2d3d_half(track_running_stats=track_running_stats)\n        param['feature_size'] = 128\n    else:\n        raise IOError('model type is wrong')\n    return model, param\n\n\ndef neq_load_customized(model, pretrained_dict):\n    model_dict = model.state_dict()\n    tmp = {}\n    print('\\n=======Check Weights Loading======')\n    print('Weights not used from pretrained file:')\n    for k, v in pretrained_dict.items():\n        if k in model_dict:\n            tmp[k] = v\n        else:\n            print(k)\n    print('---------------------------')\n    print('Weights not loaded into new model:')\n    for k, v in model_dict.items():\n        if k not in pretrained_dict:\n            print(k)\n    print('===================================\\n')\n    del pretrained_dict\n    model_dict.update(tmp)\n    del tmp\n    model.load_state_dict(model_dict)\n    return model\n\nprint('ResNet 2D3D backbone defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.611023Z","iopub.execute_input":"2026-09-22T12:37:39.611454Z","iopub.status.idle":"2026-09-22T12:37:39.647583Z","shell.execute_reply.started":"2026-09-22T12:37:39.611408Z","shell.execute_reply":"2026-09-22T12:37:39.646693Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 6. FGI Model (Audio-Visual Network)","metadata":{}},{"cell_type":"code","source":"class My_CNN_RawAud(nn.Module):\n    def __init__(self):\n        super(My_CNN_RawAud, self).__init__()\n        self.netcnnaud_layer1 = nn.Sequential(\n            nn.Conv1d(1, 128, kernel_size=80, stride=8),\n            nn.BatchNorm1d(128),\n            nn.ReLU(inplace=True),\n        )\n        self.netcnnaud_layer2 = nn.Sequential(\n            nn.MaxPool1d(kernel_size=4, stride=4),\n            nn.Conv1d(128, 192, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm1d(192),\n            nn.ReLU(inplace=True),\n            nn.Conv1d(192, 192, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm1d(192),\n        )\n\n    def forward(self, x):\n        x1 = self.netcnnaud_layer1(x)\n        x2 = self.netcnnaud_layer2(x1)\n        return x2\n\n\nclass My_Network(nn.Module):\n    def __init__(self, network='resnet18', with_attention=False,\n                 residual_conn=False, spatial_size=28):\n        super(My_Network, self).__init__()\n        self.__nFeatures__ = 24\n        self.__nChs__ = 32\n        self.__midChs__ = 32\n        self.aud2visspace = None\n        self.with_attention = with_attention\n        self.residual_conn = residual_conn\n        self.spatial_size = spatial_size\n\n        self.netcnnaud = My_CNN_RawAud()\n\n        a_size = [192, 1497]\n        v_size = [128, 15, spatial_size, spatial_size]\n        stride = int(a_size[1] / v_size[1])\n        kernel_size = int(a_size[1] - (v_size[1] - 1) * stride)\n        self.netcnnaud_to_vis = nn.Sequential(\n            nn.Conv1d(a_size[0], v_size[0] // 2, kernel_size=kernel_size, stride=stride),\n            nn.BatchNorm1d(v_size[0] // 2),\n            nn.ReLU(),\n            nn.Conv1d(v_size[0] // 2, v_size[0], kernel_size=3, stride=1, padding=1)\n        )\n\n        self.netcnnlip, self.param = select_resnet_half(network, track_running_stats=False)\n\n        map_size = v_size[2] * v_size[3]\n        self.final_fc = nn.Sequential(\n            nn.Linear(map_size, 1),\n            nn.Sigmoid()\n        )\n\n        out_emb = v_size[0] // 4\n        self.img_emb_layer = nn.Conv3d(v_size[0], out_emb, 1)\n        self.aud_emb_layer = nn.Conv1d(v_size[0], out_emb, 1)\n\n    def forward_aud(self, x):\n        (B, C) = x.shape\n        x = x.view(B, 1, C)\n        x = self.netcnnaud(x)\n        x = self.netcnnaud_to_vis(x)\n        return x\n\n    def forward_lip(self, x):\n        (B, N, C, NF, H, W) = x.shape\n        x = x.view(B * N, C, NF, H, W)\n        x = self.netcnnlip(x, return_intermediate=False)\n        if self.spatial_size != 28:\n            x = nn.functional.adaptive_avg_pool3d(x, (15, self.spatial_size, self.spatial_size))\n        return x\n\n    def forward(self, vid_seq, aud_seq):\n        vid_out = self.forward_lip(vid_seq)\n        aud_out = self.forward_aud(aud_seq)\n\n        aud_out = aud_out.view(aud_out.shape[0], aud_out.shape[1],\n                               aud_out.shape[2], 1, 1)\n\n        vid_aud_distance_ = torch.pow((vid_out - aud_out), 2)\n        vid_aud_distance_ = vid_aud_distance_.view(\n            vid_aud_distance_.shape[0],\n            vid_aud_distance_.shape[1] * vid_aud_distance_.shape[2],\n            vid_aud_distance_.shape[3] * vid_aud_distance_.shape[4]\n        )\n        vid_aud_distance_ = torch.sqrt(torch.sum(vid_aud_distance_, dim=1))\n\n        if self.with_attention:\n            img_emb = self.img_emb_layer(vid_out)\n            aud_emb = self.aud_emb_layer(\n                aud_out.view(aud_out.shape[0], aud_out.shape[1], aud_out.shape[2])\n            )\n            atts = []\n            for i_emb, a_emb in zip(img_emb, aud_emb):\n                atts.append(torch.tensordot(i_emb, a_emb, dims=([0, 1], [0, 1])))\n            att = torch.stack(atts)\n            att = att / (32 * 15)\n            att = att.view(att.shape[0], -1)\n            att = torch.nn.functional.softmax(att, dim=1)\n\n            if self.residual_conn:\n                vid_aud_distance_ = torch.mul(att, vid_aud_distance_) + vid_aud_distance_\n            else:\n                vid_aud_distance_ = torch.mul(att, vid_aud_distance_)\n        else:\n            att = None\n\n        final_out = self.final_fc(vid_aud_distance_)\n        return final_out, vid_aud_distance_, att\n\n\nprint('FGI Model defined.')\n_model = My_Network(with_attention=True)\ntotal_params = sum(p.numel() for p in _model.parameters())\nprint(f'Total parameters: {total_params:,}')\ndel _model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.648974Z","iopub.execute_input":"2026-09-22T12:37:39.649278Z","iopub.status.idle":"2026-09-22T12:37:39.725816Z","shell.execute_reply.started":"2026-09-22T12:37:39.649226Z","shell.execute_reply":"2026-09-22T12:37:39.725Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 7. Dataset","metadata":{}},{"cell_type":"code","source":"def pil_loader(path):\n    with open(path, 'rb') as f:\n        with Image.open(f) as img:\n            return img.convert('RGB')\n\n\ndef my_collate_rawaudio(batch):\n    batch = list(filter(lambda x: x is not None and x[1].size()[0] == 48000, batch))\n    if len(batch) == 0:\n        return [[], [], [], [], [], [], []]\n    return torch.utils.data.dataloader.default_collate(batch)\n\n\nclass deepfake_3d_rawaudio(data.Dataset):\n    def __init__(self, out_dir, mode='train', transform=None,\n                 vis_min_fake_len=2, vis_max_fake_len=-1,\n                 aud_min_fake_len=2, aud_max_fake_len=-1,\n                 using_pseudo_fake=False, dataset_name='dfdc'):\n        assert dataset_name in ['dfdc', 'fakeavceleb']\n        self.mode = mode\n        self.transform = transform\n        self.out_dir = out_dir\n        self.dataset_name = dataset_name\n\n        if mode == 'train':\n            split = os.path.join(self.out_dir, 'train_split.csv')\n            video_info = pd.read_csv(split, header=None)\n        elif mode == 'test':\n            split = os.path.join(self.out_dir, 'test_imbalance_split.csv')\n            video_info = pd.read_csv(split, header=None)\n        else:\n            raise ValueError('wrong mode')\n\n        self.using_pseudo_fake = using_pseudo_fake if mode == 'train' else False\n\n        self.label_dict_encode = {'fake': 0, 'real': 1}\n        self.label_dict_decode = {'0': 'fake', '1': 'real'}\n\n        self.video_info = video_info\n        self.vis_min_fake_len = vis_min_fake_len\n        self.vis_max_fake_len = vis_max_fake_len\n        self.aud_min_fake_len = aud_min_fake_len\n        self.aud_max_fake_len = aud_max_fake_len\n\n    def _generate_pseudo_fake(self, t_seq, audio, other_t_seq, other_audio):\n        chosen_type = random.choice([0, 1, 2])\n        if chosen_type == 0:\n            audio = self._augment_pseudo_fake(audio, audio.shape[0],\n                                              minimum_fake_length=self.aud_min_fake_len,\n                                              maximum_fake_length=self.aud_max_fake_len,\n                                              other_data=other_audio)\n        elif chosen_type == 1:\n            t_seq = self._augment_pseudo_fake(t_seq, t_seq.shape[0],\n                                              minimum_fake_length=self.vis_min_fake_len,\n                                              maximum_fake_length=self.vis_max_fake_len,\n                                              other_data=other_t_seq)\n        elif chosen_type == 2:\n            audio = self._augment_pseudo_fake(audio, audio.shape[0],\n                                              minimum_fake_length=self.aud_min_fake_len,\n                                              maximum_fake_length=self.aud_max_fake_len,\n                                              other_data=other_audio)\n            t_seq = self._augment_pseudo_fake(t_seq, t_seq.shape[0],\n                                              minimum_fake_length=self.vis_min_fake_len,\n                                              maximum_fake_length=self.vis_max_fake_len,\n                                              other_data=other_t_seq)\n        return t_seq, audio, chosen_type\n\n    def _select_pseudo_fake_window(self, data_length, minimum=2, maximum=-1):\n        if 0 < minimum <= 1:\n            minimum = max(2, int(minimum * data_length))\n        if minimum == -1:\n            minimum = data_length\n        if maximum == -1:\n            maximum = data_length\n        elif 0 < maximum <= 1:\n            maximum = min(minimum, int(maximum * data_length))\n        assert minimum >= 2\n        assert maximum >= minimum\n        fake_len = random.randint(minimum, maximum)\n        start_pos = random.randint(0, data_length - fake_len)\n        end_pos = min(data_length, start_pos + fake_len)\n        return fake_len, start_pos, end_pos\n\n    def _replace_with_other(self, data_tensor, fake_len, start_pos, end_pos, other_data):\n        data_tensor[start_pos:end_pos] = other_data[start_pos:end_pos]\n        return data_tensor\n\n    def _augment_pseudo_fake(self, data_tensor, time_len,\n                             minimum_fake_length=2, maximum_fake_length=-1,\n                             other_data=None):\n        fake_len, start_pos, end_pos = self._select_pseudo_fake_window(\n            data_length=time_len, minimum=minimum_fake_length, maximum=maximum_fake_length)\n        data_tensor = self._replace_with_other(data_tensor, fake_len, start_pos, end_pos, other_data)\n        return data_tensor\n\n    def _get_other_item(self, index):\n        if self.using_pseudo_fake and self.mode == 'train':\n            success = 0\n            while success == 0:\n                try:\n                    other_index = random.randint(0, self.__len__() - 1)\n                    while index == other_index:\n                        other_index = random.randint(0, self.__len__() - 1)\n                    other_vpath, other_audiopath, other_label = self.video_info.iloc[other_index]\n                    other_vpath = os.path.join(self.out_dir, other_vpath)\n                    other_audiopath = os.path.join(self.out_dir, other_audiopath)\n                    other_seq = [pil_loader(os.path.join(other_vpath, img)) for img in\n                                 sorted(os.listdir(other_vpath))]\n                    other_sample_rate, other_audio = wavfile.read(other_audiopath)\n                    other_t_seq = self.transform(other_seq)\n                    other_t_seq = torch.stack(other_t_seq, 0)\n                    other_normalized_raw_audio = min_max_normalize(\n                        other_audio, int(other_audio.min()), int(other_audio.max()))\n                    other_normalized_raw_audio = torch.from_numpy(\n                        other_normalized_raw_audio.astype(float)).float()\n                    other_normalized_raw_audio = (other_normalized_raw_audio - 0.5) / 0.5\n                    if len(other_normalized_raw_audio) < 48000 or len(other_t_seq) < 30:\n                        continue\n                    success = 1\n                except:\n                    continue\n        else:\n            other_t_seq = None\n            other_normalized_raw_audio = None\n        return other_t_seq, other_normalized_raw_audio\n\n    def __getitem__(self, index):\n        success = 0\n        while success == 0:\n            try:\n                vpath, audiopath, label = self.video_info.iloc[index]\n                vpath = os.path.join(self.out_dir, vpath)\n                audiopath = os.path.join(self.out_dir, audiopath)\n                seq = [pil_loader(os.path.join(vpath, img)) for img in sorted(os.listdir(vpath))]\n                sample_rate, audio = wavfile.read(audiopath)\n                t_seq = self.transform(seq)\n                (C, H, W) = t_seq[0].size()\n                t_seq = torch.stack(t_seq, 0)\n                normalized_raw_audio = min_max_normalize(audio, int(audio.min()), int(audio.max()))\n                normalized_raw_audio = torch.from_numpy(normalized_raw_audio.astype(float)).float()\n                normalized_raw_audio = (normalized_raw_audio - 0.5) / 0.5\n                success = 1\n            except:\n                index = random.randint(0, self.__len__() - 1)\n                continue\n\n        if self.using_pseudo_fake:\n            use_pseudo_fake = random.randint(0, 1)\n            if use_pseudo_fake:\n                other_t_seq, other_normalized_raw_audio = self._get_other_item(index)\n                t_seq, normalized_raw_audio, chosen_type = self._generate_pseudo_fake(\n                    t_seq, normalized_raw_audio, other_t_seq, other_normalized_raw_audio)\n                label = 'fake'\n\n        t_seq = t_seq.view(1, 30, C, H, W).transpose(1, 2)\n        vid = self.label_dict_encode[label]\n        return t_seq, normalized_raw_audio, torch.LongTensor([vid]), audiopath\n\n    def __len__(self):\n        return len(self.video_info)\n\n    def encode_label(self, label_name):\n        return self.label_dict_encode[label_name]\n\nprint('Dataset class defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.72728Z","iopub.execute_input":"2026-09-22T12:37:39.727617Z","iopub.status.idle":"2026-09-22T12:37:39.759032Z","shell.execute_reply.started":"2026-09-22T12:37:39.727555Z","shell.execute_reply":"2026-09-22T12:37:39.758362Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 8. Preprocessing: Extract Frames and Audio from DFDC Videos\n\nThis section processes raw DFDC videos:\n1. Converts videos to 30fps\n2. Extracts frames as JPG images (resized to 224x224)\n3. Extracts audio as WAV (48kHz, mono, PCM)\n4. Splits into 1-second chunks (30 frames + corresponding audio)\n\n> **Note**: Face detection (S3FD) is skipped. Uses full frame (similar to `--dont_crop_face`).","metadata":{}},{"cell_type":"code","source":"def preprocess_video_simple(video_path, output_base_dir, label, max_chunks=None):\n    \"\"\"\n    Simplified preprocessing for a single video.\n    \"\"\"\n    video_name = os.path.splitext(os.path.basename(video_path))[0]\n    pytmp_dir = os.path.join(output_base_dir, 'pytmp', label, video_name)\n    \n    if os.path.exists(pytmp_dir) and len(os.listdir(pytmp_dir)) > 0:\n        return\n    \n    temp_frames = os.path.join(output_base_dir, 'temp_frames', video_name)\n    temp_audio = os.path.join(output_base_dir, 'temp_audio')\n    os.makedirs(temp_frames, exist_ok=True)\n    os.makedirs(temp_audio, exist_ok=True)\n    os.makedirs(pytmp_dir, exist_ok=True)\n    \n    try:\n        cmd_frames = (\n            f'ffmpeg -y -i \"{video_path}\" -r 30 -qscale:v 2 '\n            f'-vf scale=224:224 -threads 1 -f image2 '\n            f'\"{os.path.join(temp_frames, \"%06d.jpg\")}\"'\n        )\n        subprocess.call(cmd_frames, shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)\n        \n        audio_path = os.path.join(temp_audio, f'{video_name}.wav')\n        cmd_audio = (\n            f'ffmpeg -y -i \"{video_path}\" -ac 1 -vn '\n            f'-acodec pcm_s16le -ar 48000 \"{audio_path}\"'\n        )\n        subprocess.call(cmd_audio, shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)\n        \n        if not os.path.exists(audio_path):\n            shutil.rmtree(pytmp_dir, ignore_errors=True)\n            return\n        \n        total_frames = len([f for f in os.listdir(temp_frames) if f.endswith('.jpg')])\n        chunk_count = 0\n        \n        for frame_start in range(0, total_frames, 30):\n            if frame_start + 30 > total_frames:\n                continue\n            if max_chunks and chunk_count >= max_chunks:\n                break\n            \n            chunk_id = '%05d' % (frame_start // 30)\n            chunk_dir = os.path.join(pytmp_dir, chunk_id)\n            os.makedirs(chunk_dir, exist_ok=True)\n            \n            for i in range(frame_start + 1, frame_start + 31):\n                src = os.path.join(temp_frames, '%06d.jpg' % i)\n                if os.path.exists(src):\n                    shutil.copy(src, chunk_dir)\n            \n            chunk_audio = os.path.join(pytmp_dir, f'{chunk_id}.wav')\n            audio_start = frame_start / 30.0\n            audio_end = (frame_start + 30) / 30.0\n            cmd_chunk_audio = (\n                f'ffmpeg -y -i \"{audio_path}\" -ss {audio_start:.3f} '\n                f'-to {audio_end:.3f} \"{chunk_audio}\"'\n            )\n            subprocess.call(cmd_chunk_audio, shell=True,\n                          stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)\n            chunk_count += 1\n        \n        shutil.rmtree(temp_frames, ignore_errors=True)\n        if os.path.exists(audio_path):\n            os.remove(audio_path)\n            \n    except Exception as e:\n        print(f'Error processing {video_name}: {e}')\n        shutil.rmtree(pytmp_dir, ignore_errors=True)\n        shutil.rmtree(temp_frames, ignore_errors=True)\n\nprint('Preprocessing functions defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.760423Z","iopub.execute_input":"2026-09-22T12:37:39.760707Z","iopub.status.idle":"2026-09-22T12:37:39.776143Z","shell.execute_reply.started":"2026-09-22T12:37:39.760653Z","shell.execute_reply":"2026-09-22T12:37:39.775096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\n\ndef organize_dfdc_data(dfdc_dir, work_dir, max_videos_per_class=None):\n    \"\"\"\n    Organize DFDC data into real/fake folders and preprocess.\n    Supports both train_sample_videos and full dfdc_train_part_XX format.\n    \"\"\"\n    train_dir = os.path.join(work_dir, 'train')\n    \n    metadata_path = os.path.join(dfdc_dir, 'metadata.json')\n    if not os.path.exists(metadata_path):\n        metadata_path = os.path.join(dfdc_dir, 'train_sample_videos', 'metadata.json')\n    \n    if not os.path.exists(metadata_path):\n        part_dirs = sorted(glob.glob(os.path.join(dfdc_dir, 'dfdc_train_part_*')))\n        if not part_dirs:\n            print(f'ERROR: Cannot find DFDC data in {dfdc_dir}')\n            print(f'Contents of {dfdc_dir}:')\n            if os.path.exists(dfdc_dir):\n                for item in os.listdir(dfdc_dir)[:20]:\n                    print(f'  {item}')\n            return\n        \n        real_count = 0\n        fake_count = 0\n        for part_dir in part_dirs:\n            part_metadata = os.path.join(part_dir, 'metadata.json')\n            if not os.path.exists(part_metadata):\n                continue\n            with open(part_metadata, 'r') as f:\n                metadata = json.load(f)\n            \n            for video_name, info in metadata.items():\n                video_path = os.path.join(part_dir, video_name)\n                if not os.path.exists(video_path):\n                    continue\n                label = 'real' if info['label'] == 'REAL' else 'fake'\n                \n                if max_videos_per_class:\n                    if label == 'real' and real_count >= max_videos_per_class:\n                        continue\n                    if label == 'fake' and fake_count >= max_videos_per_class:\n                        continue\n                \n                preprocess_video_simple(video_path, train_dir, label, max_chunks=5)\n                \n                if label == 'real':\n                    real_count += 1\n                else:\n                    fake_count += 1\n                \n                if (real_count + fake_count) % 50 == 0:\n                    print(f'Processed {real_count} real, {fake_count} fake videos...')\n        \n        print(f'Total processed: {real_count} real, {fake_count} fake videos')\n        return\n    \n    with open(metadata_path, 'r') as f:\n        metadata = json.load(f)\n    \n    video_dir = os.path.dirname(metadata_path)\n    real_count = 0\n    fake_count = 0\n    \n    for video_name, info in tqdm(metadata.items(), desc='Preprocessing videos'):\n        video_path = os.path.join(video_dir, video_name)\n        if not os.path.exists(video_path):\n            continue\n        label = 'real' if info['label'] == 'REAL' else 'fake'\n        \n        if max_videos_per_class:\n            if label == 'real' and real_count >= max_videos_per_class:\n                continue\n            if label == 'fake' and fake_count >= max_videos_per_class:\n                continue\n        \n        preprocess_video_simple(video_path, train_dir, label, max_chunks=5)\n        \n        if label == 'real':\n            real_count += 1\n        else:\n            fake_count += 1\n    \n    print(f'\\nTotal processed: {real_count} real, {fake_count} fake videos')\n\nprint('DFDC organizer defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.777409Z","iopub.execute_input":"2026-09-22T12:37:39.777717Z","iopub.status.idle":"2026-09-22T12:37:39.793675Z","shell.execute_reply.started":"2026-09-22T12:37:39.777679Z","shell.execute_reply":"2026-09-22T12:37:39.792989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================== RUN PREPROCESSING ======================\n# Set MAX_VIDEOS = None for full dataset, or a small number for testing\n\nprint('Starting DFDC data preprocessing...')\nprint(f'Input directory: {cfg.DFDC_DATA_DIR}')\nprint(f'Output directory: {cfg.WORK_DIR}')\n\n# Adjust this based on your needs:\n#   None  -> Process ALL videos (long time)\n#   50    -> Quick test\n#   500   -> Medium experiment\nMAX_VIDEOS = 300\n\norganize_dfdc_data(cfg.DFDC_DATA_DIR, cfg.WORK_DIR, max_videos_per_class=MAX_VIDEOS)\n\nfor temp_dir in ['temp_frames', 'temp_audio']:\n    temp_path = os.path.join(cfg.TRAIN_DIR, temp_dir)\n    if os.path.exists(temp_path):\n        shutil.rmtree(temp_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:39.795036Z","iopub.execute_input":"2026-09-22T12:37:39.795513Z","iopub.status.idle":"2026-09-22T12:37:41.920027Z","shell.execute_reply.started":"2026-09-22T12:37:39.795275Z","shell.execute_reply":"2026-09-22T12:37:41.919287Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 9. Generate CSV Split Files","metadata":{}},{"cell_type":"code","source":"def write_csv_list(data_list, path):\n    with open(path, 'w', newline='') as f:\n        writer = csv.writer(f, delimiter=',')\n        for row in data_list:\n            if row:\n                writer.writerow(row)\n    print(f'Split saved to {path} ({len(data_list)} entries)')\n\n\ndef generate_splits(work_dir, train_test_ratio=0.8):\n    pytmp_real = os.path.join(work_dir, 'train', 'pytmp', 'real')\n    pytmp_fake = os.path.join(work_dir, 'train', 'pytmp', 'fake')\n    \n    real_videos = []\n    fake_videos = []\n    \n    if os.path.exists(pytmp_real):\n        real_videos = [d for d in os.listdir(pytmp_real)\n                       if os.path.isdir(os.path.join(pytmp_real, d))]\n    if os.path.exists(pytmp_fake):\n        fake_videos = [d for d in os.listdir(pytmp_fake)\n                       if os.path.isdir(os.path.join(pytmp_fake, d))]\n    \n    print(f'Found {len(real_videos)} real videos, {len(fake_videos)} fake videos')\n    \n    random.shuffle(real_videos)\n    random.shuffle(fake_videos)\n    \n    real_train_count = int(len(real_videos) * train_test_ratio)\n    fake_train_count = int(len(fake_videos) * train_test_ratio)\n    \n    real_train = real_videos[:real_train_count]\n    real_test = real_videos[real_train_count:]\n    fake_train = fake_videos[:fake_train_count]\n    fake_test = fake_videos[fake_train_count:]\n    \n    print(f'Train: {len(real_train)} real, {len(fake_train)} fake')\n    print(f'Test: {len(real_test)} real, {len(fake_test)} fake')\n    \n    train_set = []\n    base_train = os.path.join('train', 'pytmp')\n    \n    for video_name in fake_train:\n        video_dir = os.path.join(pytmp_fake, video_name)\n        for chunk in os.listdir(video_dir):\n            chunk_path = os.path.join(video_dir, chunk)\n            if os.path.isdir(chunk_path):\n                audio_file = chunk + '.wav'\n                audio_path = os.path.join(video_dir, audio_file)\n                if os.path.exists(audio_path):\n                    train_set.append([\n                        os.path.join(base_train, 'fake', video_name, chunk),\n                        os.path.join(base_train, 'fake', video_name, audio_file),\n                        'fake'\n                    ])\n    \n    for video_name in real_train:\n        video_dir = os.path.join(pytmp_real, video_name)\n        for chunk in os.listdir(video_dir):\n            chunk_path = os.path.join(video_dir, chunk)\n            if os.path.isdir(chunk_path):\n                audio_file = chunk + '.wav'\n                audio_path = os.path.join(video_dir, audio_file)\n                if os.path.exists(audio_path):\n                    train_set.append([\n                        os.path.join(base_train, 'real', video_name, chunk),\n                        os.path.join(base_train, 'real', video_name, audio_file),\n                        'real'\n                    ])\n    \n    test_set = []\n    for video_name in fake_test:\n        video_dir = os.path.join(pytmp_fake, video_name)\n        for chunk in os.listdir(video_dir):\n            chunk_path = os.path.join(video_dir, chunk)\n            if os.path.isdir(chunk_path):\n                audio_file = chunk + '.wav'\n                audio_path = os.path.join(video_dir, audio_file)\n                if os.path.exists(audio_path):\n                    test_set.append([\n                        os.path.join(base_train, 'fake', video_name, chunk),\n                        os.path.join(base_train, 'fake', video_name, audio_file),\n                        'fake'\n                    ])\n    \n    for video_name in real_test:\n        video_dir = os.path.join(pytmp_real, video_name)\n        for chunk in os.listdir(video_dir):\n            chunk_path = os.path.join(video_dir, chunk)\n            if os.path.isdir(chunk_path):\n                audio_file = chunk + '.wav'\n                audio_path = os.path.join(video_dir, audio_file)\n                if os.path.exists(audio_path):\n                    test_set.append([\n                        os.path.join(base_train, 'real', video_name, chunk),\n                        os.path.join(base_train, 'real', video_name, audio_file),\n                        'real'\n                    ])\n    \n    random.shuffle(train_set)\n    random.shuffle(test_set)\n    \n    write_csv_list(train_set, os.path.join(work_dir, 'train_split.csv'))\n    write_csv_list(test_set, os.path.join(work_dir, 'test_imbalance_split.csv'))\n    \n    return len(train_set), len(test_set)\n\n\nn_train, n_test = generate_splits(cfg.WORK_DIR)\nprint(f'\\nTrain samples: {n_train}, Test samples: {n_test}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:41.921519Z","iopub.execute_input":"2026-09-22T12:37:41.921815Z","iopub.status.idle":"2026-09-22T12:37:41.997191Z","shell.execute_reply.started":"2026-09-22T12:37:41.921727Z","shell.execute_reply":"2026-09-22T12:37:41.99643Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 10. Data Loaders","metadata":{}},{"cell_type":"code","source":"def get_rawaudio_data(transform, cfg, mode='train'):\n    print(f'Loading data for \"{mode}\"...')\n    dataset = deepfake_3d_rawaudio(\n        out_dir=cfg.WORK_DIR, mode=mode, transform=transform,\n        vis_min_fake_len=cfg.VIS_MIN_FAKE_LEN, vis_max_fake_len=cfg.VIS_MAX_FAKE_LEN,\n        aud_min_fake_len=cfg.AUD_MIN_FAKE_LEN, aud_max_fake_len=cfg.AUD_MAX_FAKE_LEN,\n        using_pseudo_fake=cfg.USING_PSEUDO_FAKE if mode == 'train' else False,\n        dataset_name='dfdc'\n    )\n    sampler = data.RandomSampler(dataset)\n\n    if mode == 'train':\n        data_loader = data.DataLoader(\n            dataset, batch_size=cfg.BATCH_SIZE, sampler=sampler, shuffle=False,\n            num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True,\n            collate_fn=my_collate_rawaudio\n        )\n    elif mode == 'test':\n        data_loader = data.DataLoader(\n            dataset, batch_size=1, sampler=sampler, shuffle=False,\n            num_workers=cfg.NUM_WORKERS, pin_memory=True,\n            collate_fn=my_collate_rawaudio\n        )\n    else:\n        raise ValueError(f'No mode {mode}')\n\n    print(f'\"{mode}\" dataset size: {len(dataset)}')\n    return data_loader\n\n\ntrain_transform = transforms.Compose([\n    Scale(size=(cfg.IMG_DIM, cfg.IMG_DIM)),\n    ToTensor(),\n    Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n])\n\ntest_transform = transforms.Compose([\n    Scale(size=(cfg.IMG_DIM, cfg.IMG_DIM)),\n    ToTensor(),\n    Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n])\n\ntrain_loader = get_rawaudio_data(train_transform, cfg, 'train')\nprint(f'Train loader: {len(train_loader)} batches')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:41.998561Z","iopub.execute_input":"2026-09-22T12:37:41.998817Z","iopub.status.idle":"2026-09-22T12:37:42.015812Z","shell.execute_reply.started":"2026-09-22T12:37:41.99876Z","shell.execute_reply":"2026-09-22T12:37:42.014733Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 11. Training Loop","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(data_loader, model, optimizer, criterion, epoch, cfg, writer=None):\n    losses = AverageMeter()\n    real_distances = AverageMeter()\n    fake_distances = AverageMeter()\n    model.train()\n    cuda_dev = torch.device('cuda')\n\n    for idx, batch_data in enumerate(data_loader):\n        if len(batch_data[0]) == 0:\n            continue\n        video_seq, audio_seq, target, audiopath = batch_data\n        target = 1 - target\n\n        tic = time.time()\n        video_seq = video_seq.to(cuda_dev)\n        audio_seq = audio_seq.to(cuda_dev)\n        target = target.to(cuda_dev)\n        B = video_seq.size(0)\n\n        out, vid_out_dist, att = model(video_seq, audio_seq)\n        del video_seq, audio_seq\n\n        loss = criterion(out, target.float())\n        losses.update(loss.item(), B)\n\n        n_real = torch.sum(target.view(-1) == 0).detach().cpu().item()\n        if n_real > 0:\n            real_distances.update(torch.mean(vid_out_dist[target.view(-1) == 0]).item())\n        n_fake = torch.sum(target.view(-1) == 1).detach().cpu().item()\n        if n_fake > 0:\n            fake_distances.update(torch.mean(vid_out_dist[target.view(-1) == 1]).item())\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        if idx % cfg.PRINT_FREQ == 0:\n            elapsed = time.time() - tic\n            print(f'Epoch [{epoch}][{idx}/{len(data_loader)}] '\n                  f'Time: {elapsed:.2f}s  '\n                  f'Loss: {losses.val:.4f} ({losses.local_avg:.4f})  '\n                  f'Real: {real_distances.val:.4f}  Fake: {fake_distances.val:.4f}')\n\n        if writer:\n            writer.add_scalar('local/loss', losses.val, epoch * len(data_loader) + idx)\n\n    return losses.avg, real_distances.avg, fake_distances.avg\n\n\ndef test_model(data_loader, model, cuda_device):\n    real_distances = AverageMeter()\n    fake_distances = AverageMeter()\n    model.eval()\n    test_pred = {}\n    test_target = {}\n    test_num_chunks = {}\n\n    with torch.no_grad():\n        for idx, batch_data in tqdm(enumerate(data_loader), total=len(data_loader), desc='Testing'):\n            if len(batch_data[0]) == 0:\n                continue\n            video_seq, audio_seq, target, audiopath = batch_data\n            target = 1 - target\n\n            video_seq = video_seq[0].unsqueeze(0).to(cuda_device)\n            audio_seq = audio_seq[0].unsqueeze(0).to(cuda_device)\n            target = target[0].unsqueeze(0).to(cuda_device)\n\n            pred, vid_aud_dist, att = model(video_seq, audio_seq)\n            del video_seq, audio_seq\n\n            tar = target[0, :].view(-1).item()\n            vid_name = audiopath[0].split('/')[-2]\n\n            if test_pred.get(vid_name):\n                test_pred[vid_name] += pred[0].view(-1).item()\n                test_num_chunks[vid_name] += 1\n            else:\n                test_pred[vid_name] = pred[0].view(-1).item()\n                test_num_chunks[vid_name] = 1\n\n            if not test_target.get(vid_name):\n                test_target[vid_name] = tar\n\n            if tar == 1:\n                fake_distances.update(torch.mean(vid_aud_dist[0]).item())\n            else:\n                real_distances.update(torch.mean(vid_aud_dist[0]).item())\n\n    pred_fake = []\n    pred_real = []\n    for video, score in test_pred.items():\n        tar = test_target[video]\n        num_chunks = test_num_chunks[video]\n        mean_pred = score / num_chunks\n        if tar == 1:\n            pred_fake.append(mean_pred)\n        else:\n            pred_real.append(mean_pred)\n\n    if len(pred_fake) > 0 and len(pred_real) > 0:\n        pred_all = np.concatenate([np.array(pred_fake), np.array(pred_real)])\n        gt_all = np.concatenate([np.ones(len(pred_fake)), np.zeros(len(pred_real))])\n        auc = roc_auc_score(gt_all, pred_all)\n    else:\n        auc = 0.0\n\n    print(f'Fake dist: {fake_distances.avg:.4f}  Real dist: {real_distances.avg:.4f}  AUC: {auc:.4f}')\n    return auc, real_distances.avg, fake_distances.avg\n\nprint('Training and testing functions defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:42.01697Z","iopub.execute_input":"2026-09-22T12:37:42.017236Z","iopub.status.idle":"2026-09-22T12:37:42.03967Z","shell.execute_reply.started":"2026-09-22T12:37:42.017198Z","shell.execute_reply":"2026-09-22T12:37:42.038858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================== INITIALIZE MODEL ======================\n\ntorch.manual_seed(cfg.SEED)\nnp.random.seed(cfg.SEED)\ncuda = torch.device('cuda')\n\nmodel = My_Network(\n    with_attention=cfg.WITH_ATTENTION,\n    residual_conn=cfg.RESIDUAL_CONN,\n    spatial_size=cfg.SPATIAL_SIZE\n)\nmodel = model.cuda()\nmodel = nn.DataParallel(model)\n\ncriterion = nn.BCELoss()\noptimizer = optim.Adam(model.parameters(), lr=cfg.LEARNING_RATE, weight_decay=cfg.WEIGHT_DECAY)\n\nlog_dir = os.path.join(cfg.WORK_DIR, 'logs')\nmodel_dir = os.path.join(cfg.WORK_DIR, 'models')\nos.makedirs(log_dir, exist_ok=True)\nos.makedirs(model_dir, exist_ok=True)\nwriter = SummaryWriter(logdir=log_dir)\n\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f'Trainable: {trainable_params:,} / Total: {total_params:,} parameters')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:42.040814Z","iopub.execute_input":"2026-09-22T12:37:42.041088Z","iopub.status.idle":"2026-09-22T12:37:45.509172Z","shell.execute_reply.started":"2026-09-22T12:37:42.041015Z","shell.execute_reply":"2026-09-22T12:37:45.508414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================== MAIN TRAINING LOOP ======================\n\nprint('=' * 60)\nprint('Starting Training')\nprint(f'  Epochs: {cfg.EPOCHS}, BS: {cfg.BATCH_SIZE}, LR: {cfg.LEARNING_RATE}')\nprint(f'  Attention: {cfg.WITH_ATTENTION}, Pseudo-Fake: {cfg.USING_PSEUDO_FAKE}')\nprint('=' * 60)\n\nleast_loss = float('inf')\ntrain_losses = []\ntrain_real_dists = []\ntrain_fake_dists = []\n\ntorch.backends.cudnn.benchmark = True\n\nfor epoch in range(cfg.EPOCHS):\n    print(f'\\n--- Epoch {epoch+1}/{cfg.EPOCHS} ---')\n    \n    train_loss, real_distance, fake_distance = train_one_epoch(\n        train_loader, model, optimizer, criterion, epoch, cfg, writer\n    )\n    \n    train_losses.append(train_loss)\n    train_real_dists.append(real_distance)\n    train_fake_dists.append(fake_distance)\n    \n    writer.add_scalar('global/loss', train_loss, epoch)\n    writer.add_scalar('global/real_distance', real_distance, epoch)\n    writer.add_scalar('global/fake_distance', fake_distance, epoch)\n    \n    print(f'Epoch {epoch+1}: Loss={train_loss:.4f}, '\n          f'Real={real_distance:.4f}, Fake={fake_distance:.4f}')\n    \n    is_best = train_loss <= least_loss\n    least_loss = min(least_loss, train_loss)\n    \n    save_checkpoint({\n        'epoch': epoch + 1,\n        'net': cfg.NET,\n        'state_dict': model.state_dict(),\n        'least_loss': least_loss,\n        'optimizer': optimizer.state_dict(),\n        'iteration': epoch\n    }, is_best, filename=os.path.join(model_dir, f'epoch{epoch+1}.pth.tar'),\n       keep_all=cfg.SAVE_ALL)\n    \n    if is_best:\n        print(f'  * New best model! (loss={train_loss:.4f})')\n\nprint(f'\\nTraining finished! Best loss: {least_loss:.4f}')\nwriter.close()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-22T12:37:45.510526Z","iopub.execute_input":"2026-09-22T12:37:45.510733Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 12. Training Visualization","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(train_losses, 'b-', linewidth=2)\naxes[0].set_title('Training Loss', fontsize=14)\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('BCE Loss')\naxes[0].grid(True, alpha=0.3)\n\naxes[1].plot(train_real_dists, 'g-', linewidth=2, label='Real')\naxes[1].plot(train_fake_dists, 'r-', linewidth=2, label='Fake')\naxes[1].set_title('Audio-Visual Distance', fontsize=14)\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Distance')\naxes[1].legend()\naxes[1].grid(True, alpha=0.3)\n\ngap = [f - r for f, r in zip(train_fake_dists, train_real_dists)]\naxes[2].plot(gap, 'm-', linewidth=2)\naxes[2].set_title('Distance Gap (Fake - Real)', fontsize=14)\naxes[2].set_xlabel('Epoch')\naxes[2].set_ylabel('Gap')\naxes[2].axhline(y=0, color='k', linestyle='--', alpha=0.3)\naxes[2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(os.path.join(cfg.WORK_DIR, 'training_curves.png'), dpi=150, bbox_inches='tight')\nplt.show()\nprint('Training curves saved.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 13. Testing and Evaluation","metadata":{}},{"cell_type":"code","source":"# Load best model\nbest_model_path = glob.glob(os.path.join(model_dir, 'model_best_*.pth.tar'))\nif best_model_path:\n    best_model_path = best_model_path[0]\n    print(f'Loading best model: {best_model_path}')\n    checkpoint = torch.load(best_model_path)\n    model.load_state_dict(checkpoint['state_dict'])\n    print(f'Loaded model from epoch {checkpoint[\"epoch\"]}')\nelse:\n    print('No best model found, using latest model state.')\n\n# Create test loader\ntest_cfg = copy.deepcopy(cfg)\ntest_cfg.USING_PSEUDO_FAKE = False\n\ntest_loader = get_rawaudio_data(test_transform, test_cfg, 'test')\nprint(f'Test loader: {len(test_loader)} samples')\n\nprint('\\n--- Running Evaluation ---')\nauc, real_dist, fake_dist = test_model(test_loader, model, cuda)\n\nprint('\\n' + '=' * 60)\nprint(f'  FINAL RESULTS')\nprint(f'  AUC Score: {auc:.4f}')\nprint(f'  Real Distance: {real_dist:.4f}')\nprint(f'  Fake Distance: {fake_dist:.4f}')\nprint('=' * 60)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 14. Save Final Model","metadata":{}},{"cell_type":"code","source":"final_model_path = os.path.join(cfg.WORK_DIR, 'fgi_deepfake_detector_final.pth')\ntorch.save({\n    'state_dict': model.state_dict(),\n    'config': {\n        'with_attention': cfg.WITH_ATTENTION,\n        'residual_conn': cfg.RESIDUAL_CONN,\n        'spatial_size': cfg.SPATIAL_SIZE,\n        'net': cfg.NET,\n    },\n    'auc': auc,\n}, final_model_path)\n\nfile_size = os.path.getsize(final_model_path) / (1024 * 1024)\nprint(f'Final model saved: {final_model_path} ({file_size:.2f} MB)')\nprint(f'AUC: {auc:.4f}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## 15. Inference on New Videos","metadata":{}},{"cell_type":"code","source":"def predict_video(video_path, model, transform, cuda_device, max_chunks=10):\n    \"\"\"\n    Run deepfake prediction on a single video.\n    Returns probability of being fake (closer to 1 = FAKE, closer to 0 = REAL).\n    \"\"\"\n    model.eval()\n    temp_dir = '/kaggle/working/temp_inference'\n    os.makedirs(temp_dir, exist_ok=True)\n    \n    frames_dir = os.path.join(temp_dir, 'frames')\n    os.makedirs(frames_dir, exist_ok=True)\n    cmd = f'ffmpeg -y -i \"{video_path}\" -r 30 -qscale:v 2 -vf scale=224:224 -f image2 \"{os.path.join(frames_dir, \"%06d.jpg\")}\"'\n    subprocess.call(cmd, shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)\n    \n    audio_path = os.path.join(temp_dir, 'audio.wav')\n    cmd = f'ffmpeg -y -i \"{video_path}\" -ac 1 -vn -acodec pcm_s16le -ar 48000 \"{audio_path}\"'\n    subprocess.call(cmd, shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)\n    \n    if not os.path.exists(audio_path):\n        shutil.rmtree(temp_dir, ignore_errors=True)\n        return None\n    \n    frames = sorted([f for f in os.listdir(frames_dir) if f.endswith('.jpg')])\n    total_frames = len(frames)\n    sample_rate, full_audio = wavfile.read(audio_path)\n    \n    predictions = []\n    with torch.no_grad():\n        for chunk_start in range(0, total_frames, 30):\n            if chunk_start + 30 > total_frames:\n                break\n            if max_chunks and len(predictions) >= max_chunks:\n                break\n            \n            chunk_frames = [pil_loader(os.path.join(frames_dir, frames[i]))\n                           for i in range(chunk_start, chunk_start + 30)]\n            t_seq = transform(chunk_frames)\n            (C, H, W) = t_seq[0].size()\n            t_seq = torch.stack(t_seq, 0)\n            t_seq = t_seq.view(1, 1, 30, C, H, W).transpose(2, 3)\n            \n            audio_start = int(chunk_start / 30.0 * sample_rate)\n            audio_end = int((chunk_start + 30) / 30.0 * sample_rate)\n            chunk_audio = full_audio[audio_start:audio_end]\n            \n            if len(chunk_audio) < 48000:\n                chunk_audio = np.pad(chunk_audio, (0, 48000 - len(chunk_audio)))\n            elif len(chunk_audio) > 48000:\n                chunk_audio = chunk_audio[:48000]\n            \n            norm_audio = min_max_normalize(chunk_audio, int(chunk_audio.min()), int(chunk_audio.max()))\n            norm_audio = torch.from_numpy(norm_audio.astype(float)).float()\n            norm_audio = (norm_audio - 0.5) / 0.5\n            norm_audio = norm_audio.unsqueeze(0)\n            \n            pred, _, _ = model(t_seq.to(cuda_device), norm_audio.to(cuda_device))\n            predictions.append(pred[0].item())\n    \n    shutil.rmtree(temp_dir, ignore_errors=True)\n    return np.mean(predictions) if predictions else None\n\n\nprint('Inference function defined.')\nprint('Usage: score = predict_video(\"/path/to/video.mp4\", model, test_transform, cuda)')\nprint('  score close to 1 = FAKE, close to 0 = REAL')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example inference (uncomment to test):\n\n# sample_video = '/kaggle/input/deepfake-detection-challenge/train_sample_videos/aaqaifqrwn.mp4'\n# if os.path.exists(sample_video):\n#     score = predict_video(sample_video, model, test_transform, cuda)\n#     if score is not None:\n#         label = 'FAKE' if score > 0.5 else 'REAL'\n#         print(f'Video: {os.path.basename(sample_video)}')\n#         print(f'Score: {score:.4f} -> Prediction: {label}')\n\nprint('\\nNotebook complete!')\nprint('To train on full DFDC dataset, set MAX_VIDEOS = None in the preprocessing cell.')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}