{"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":113558,"databundleVersionId":14456136,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install segmentation-models-pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:23:48.646705Z","iopub.execute_input":"2025-11-20T08:23:48.646925Z","iopub.status.idle":"2025-11-20T08:25:09.810043Z","shell.execute_reply.started":"2025-11-20T08:23:48.646902Z","shell.execute_reply":"2025-11-20T08:25:09.809285Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 1. Configuration & Imports","metadata":{}},{"cell_type":"code","source":"import gc\nimport torch\n\n# Force garbage collection\ngc.collect()\n\n# Clear CUDA cache\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    torch.cuda.ipc_collect()\n    \nprint(f\"Memory Allocated: {torch.cuda.memory_allocated() / 1024**3:.2f} GB\")\nprint(f\"Memory Reserved:  {torch.cuda.memory_reserved() / 1024**3:.2f} GB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:41:45.624471Z","iopub.execute_input":"2025-11-20T08:41:45.624780Z","iopub.status.idle":"2025-11-20T08:42:06.622638Z","shell.execute_reply.started":"2025-11-20T08:41:45.624760Z","shell.execute_reply":"2025-11-20T08:42:06.621817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport random\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import GroupKFold\nfrom tqdm.notebook import tqdm\n\n# --- Configuration ---\nclass CFG:\n    seed = 42\n    debug = False\n    \n    # Paths (Verify these match your Kaggle input)\n    base_path = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection\"\n    train_img_dir = f\"{base_path}/train_images\"\n    train_mask_dir = f\"{base_path}/train_masks\"\n    save_path = \"best_vit_dual_stream.pth\"\n    \n    # Model Architecture\n    # mit_b2 is lighter than b3, preventing OOM while maintaining high accuracy\n    encoder_name = 'mit_b2' \n    img_size = 320  # Reduced from 384 to save memory\n    \n    # Training Hyperparameters\n    epochs = 20\n    batch_size = 8      # Safe size for T4 GPU\n    accum_iter = 2      # Effective batch size = 16\n    lr = 1e-4\n    weight_decay = 1e-4\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    num_workers = 2     # Reduced to prevent shared memory errors\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed(CFG.seed)\nprint(\"Configuration Loaded.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:42:06.623859Z","iopub.execute_input":"2025-11-20T08:42:06.624100Z","iopub.status.idle":"2025-11-20T08:42:06.632821Z","shell.execute_reply.started":"2025-11-20T08:42:06.624084Z","shell.execute_reply":"2025-11-20T08:42:06.632123Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Utilities & Dual-Stream ViT Model","metadata":{}},{"cell_type":"code","source":"# --- 1. SRM Filter Layer ---\nclass SRMConv2d(nn.Module):\n    def __init__(self):\n        super().__init__()\n        filters = [\n            [[0, 0, 0, 0, 0], [0, -1, 2, -1, 0], [0, 2, -4, 2, 0], [0, -1, 2, -1, 0], [0, 0, 0, 0, 0]],\n            [[-1, 2, -2, 2, -1], [2, -6, 8, -6, 2], [-2, 8, -12, 8, -2], [2, -6, 8, -6, 2], [-1, 2, -2, 2, -1]],\n            [[0, 0, 0, 0, 0], [0, 0, 0, 0, 0], [0, 1, -2, 1, 0], [0, 0, 0, 0, 0], [0, 0, 0, 0, 0]]\n        ]\n        kernel = torch.tensor(filters, dtype=torch.float32).unsqueeze(1)\n        self.register_buffer('weight', kernel)\n        \n    def forward(self, x):\n        gray = x[:, 0:1, :, :] * 0.299 + x[:, 1:2, :, :] * 0.587 + x[:, 2:3, :, :] * 0.114\n        return F.conv2d(gray, self.weight, padding=2)\n\nclass DualStreamViT(nn.Module):\n    def __init__(self, encoder_name=CFG.encoder_name):\n        super().__init__()\n        \n        # Stream 1: RGB\n        self.base_model = smp.Unet(\n            encoder_name=encoder_name, \n            encoder_weights='imagenet', \n            in_channels=3, \n            classes=1,\n            activation=None,\n            decoder_attention_type='scse'\n        )\n        \n        # Stream 2: SRM\n        self.srm_layer = SRMConv2d()\n        \n        # Multi-scale SRM Encoder\n        self.srm_stage1 = nn.Sequential(nn.Conv2d(3, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(4))\n        self.srm_stage2 = nn.Sequential(nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2))\n        self.srm_stage3 = nn.Sequential(nn.Conv2d(64, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2))\n        self.srm_stage4 = nn.Sequential(nn.Conv2d(128, 256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(), nn.MaxPool2d(2))\n        \n        # Fusion layers\n        self.fuse1 = nn.Conv2d(32, 64, 1)\n        self.fuse2 = nn.Conv2d(64, 128, 1)\n        self.fuse3 = nn.Conv2d(128, 320, 1)\n        self.fuse4 = nn.Conv2d(256, 512, 1)\n\n    def forward(self, x):\n        # 1. Get RGB Features (List of tensors)\n        rgb_features = list(self.base_model.encoder(x))\n        \n        # 2. Get SRM Features\n        noise = self.srm_layer(x)\n        s1 = self.srm_stage1(noise)\n        s2 = self.srm_stage2(s1)\n        s3 = self.srm_stage3(s2)\n        s4 = self.srm_stage4(s3)\n        \n        # 3. Fuse (Check shapes to be safe against differing padding)\n        \n        if s1.shape[-2:] == rgb_features[-4].shape[-2:]:\n            rgb_features[-4] = rgb_features[-4] + self.fuse1(s1)\n            \n        if s2.shape[-2:] == rgb_features[-3].shape[-2:]:\n            rgb_features[-3] = rgb_features[-3] + self.fuse2(s2)\n            \n        if s3.shape[-2:] == rgb_features[-2].shape[-2:]:\n            rgb_features[-2] = rgb_features[-2] + self.fuse3(s3)\n            \n        if s4.shape[-2:] == rgb_features[-1].shape[-2:]:\n            rgb_features[-1] = rgb_features[-1] + self.fuse4(s4)\n            \n        # 4. Decode\n        decoder_output = self.base_model.decoder(rgb_features)\n        \n        masks = self.base_model.segmentation_head(decoder_output)\n        return masks\n\nprint(\"Model Class Defined Successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:45:39.582406Z","iopub.execute_input":"2025-11-20T08:45:39.582753Z","iopub.status.idle":"2025-11-20T08:45:39.597426Z","shell.execute_reply.started":"2025-11-20T08:45:39.582726Z","shell.execute_reply":"2025-11-20T08:45:39.596672Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Dataset & Transforms","metadata":{}},{"cell_type":"code","source":"def get_transforms(data):\n    if data == 'train':\n        return A.Compose([\n            # Ensure strictly padded to img_size with BLACK pixels (0)\n            A.PadIfNeeded(\n                min_height=CFG.img_size, \n                min_width=CFG.img_size, \n                border_mode=cv2.BORDER_CONSTANT, \n                value=0\n            ),\n            # Crop if larger\n            A.RandomCrop(height=CFG.img_size, width=CFG.img_size),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.Normalize(),\n            ToTensorV2(),\n        ])\n    elif data == 'valid':\n        return A.Compose([\n            A.PadIfNeeded(\n                min_height=CFG.img_size, \n                min_width=CFG.img_size, \n                border_mode=cv2.BORDER_CONSTANT, \n                value=0\n            ),\n            A.CenterCrop(height=CFG.img_size, width=CFG.img_size),\n            A.Normalize(),\n            ToTensorV2(),\n        ])\n\nclass ScientificForgeryDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        # Load Image\n        image = cv2.imread(row['image_path'])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        h, w = image.shape[:2]\n        \n        # Load Mask\n        if row['label'] == 0:\n            mask = np.zeros((h, w), dtype=np.float32)\n        else:\n            try:\n                mask = np.load(row['mask_path'])\n                # Handle Multi-channel masks if they exist\n                if mask.ndim == 3: \n                    mask = np.max(mask, axis=0)\n                mask = (mask > 0).astype(np.float32)\n            except:\n                mask = np.zeros((h, w), dtype=np.float32)\n\n        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image = augmented['image']\n            mask = augmented['mask']\n            \n        return image, mask.unsqueeze(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:46:01.707158Z","iopub.execute_input":"2025-11-20T08:46:01.707513Z","iopub.status.idle":"2025-11-20T08:46:01.716466Z","shell.execute_reply.started":"2025-11-20T08:46:01.707491Z","shell.execute_reply":"2025-11-20T08:46:01.715823Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Main Training Execution","metadata":{}},{"cell_type":"code","source":"# --- Early Stopping Helper ---\nclass EarlyStopping:\n    def __init__(self, patience=5, delta=0, path='checkpoint.pth'):\n        self.patience = patience\n        self.delta = delta\n        self.path = path\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.val_loss_min = np.Inf\n\n    def __call__(self, val_loss, model):\n        score = -val_loss\n        if self.best_score is None:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n        elif score < self.best_score + self.delta:\n            self.counter += 1\n            print(f'EarlyStopping counter: {self.counter} out of {self.patience}')\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n            self.counter = 0\n\n    def save_checkpoint(self, val_loss, model):\n        torch.save(model.state_dict(), self.path)\n        print(f'>>> Saved Best Model (Loss: {val_loss:.6f})')\n        self.val_loss_min = val_loss\n\n# --- Main Execution ---\nif __name__ == \"__main__\":\n    # 1. Prepare Data\n    data = []\n    auth_files = glob.glob(f\"{CFG.train_img_dir}/authentic/*.png\")\n    for path in auth_files:\n        data.append({\"image_path\": path, \"mask_path\": None, \"label\": 0})\n        \n    forged_files = glob.glob(f\"{CFG.train_img_dir}/forged/*.png\")\n    for path in forged_files:\n        file_id = os.path.basename(path).split('.')[0]\n        mask_path = f\"{CFG.train_mask_dir}/{file_id}.npy\"\n        if os.path.exists(mask_path):\n            data.append({\"image_path\": path, \"mask_path\": mask_path, \"label\": 1})\n            \n    df = pd.DataFrame(data)\n    \n    # GroupKFold to prevent data leakage from same source images\n    df['group_id'] = df['image_path'].apply(lambda x: os.path.basename(x).split('_')[0])\n    gkf = GroupKFold(n_splits=5)\n    train_idx, valid_idx = next(gkf.split(df, groups=df['group_id']))\n    \n    train_df = df.iloc[train_idx]\n    valid_df = df.iloc[valid_idx]\n    print(f\"Train Set: {len(train_df)} | Validation Set: {len(valid_df)}\")\n\n    # 2. DataLoaders\n    train_ds = ScientificForgeryDataset(train_df, transform=get_transforms('train'))\n    valid_ds = ScientificForgeryDataset(valid_df, transform=get_transforms('valid'))\n    \n    train_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_ds, batch_size=CFG.batch_size, shuffle=False, \n                              num_workers=CFG.num_workers, pin_memory=True)\n\n    # 3. Initialize Model\n    model = DualStreamViT(encoder_name=CFG.encoder_name).to(CFG.device)\n    \n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=5, T_mult=2, eta_min=1e-6)\n    \n    # Mixed Loss: Focal (for class imbalance) + Dice (for overlap)\n    loss_focal = smp.losses.FocalLoss(mode='binary')\n    loss_dice = smp.losses.DiceLoss(mode='binary')\n    \n    # Mixed Precision\n    scaler = torch.amp.GradScaler('cuda') # Updated API for torch 2.x\n    \n    early_stopping = EarlyStopping(patience=5, path=CFG.save_path)\n\n    # 4. Training Loop\n    print(f\"Starting Training on {CFG.device}...\")\n    \n    for epoch in range(CFG.epochs):\n        model.train()\n        train_loss = 0\n        loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{CFG.epochs}\")\n        \n        for i, (images, masks) in enumerate(loop):\n            images, masks = images.to(CFG.device), masks.to(CFG.device)\n            \n            with torch.amp.autocast('cuda'):\n                outputs = model(images)\n                # 60% Focal (Pixel classification), 40% Dice (Shape matching)\n                loss = (loss_focal(outputs, masks) * 0.6) + (loss_dice(outputs.sigmoid(), masks) * 0.4)\n                loss = loss / CFG.accum_iter\n            \n            scaler.scale(loss).backward()\n            \n            if ((i + 1) % CFG.accum_iter == 0) or (i + 1 == len(train_loader)):\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            train_loss += (loss.item() * CFG.accum_iter)\n            loop.set_postfix(loss=loss.item() * CFG.accum_iter)\n            \n        avg_train_loss = train_loss / len(train_loader)\n        \n        # Validation\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for images, masks in valid_loader:\n                images, masks = images.to(CFG.device), masks.to(CFG.device)\n                with torch.amp.autocast('cuda'):\n                    outputs = model(images)\n                    loss = (loss_focal(outputs, masks) * 0.6) + (loss_dice(outputs.sigmoid(), masks) * 0.4)\n                val_loss += loss.item()\n        \n        avg_val_loss = val_loss / len(valid_loader)\n        \n        # Update Scheduler\n        scheduler.step()\n        \n        print(f\"Epoch {epoch+1} Summary | Train Loss: {avg_train_loss:.4f} | Valid Loss: {avg_val_loss:.4f}\")\n        \n        # Early Stopping\n        early_stopping(avg_val_loss, model)\n        if early_stopping.early_stop:\n            print(\"Early stopping triggered.\")\n            break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:46:05.481286Z","iopub.execute_input":"2025-11-20T08:46:05.481843Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}