{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":113558,"databundleVersionId":14878066,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom torchvision import transforms # <--- NEW\nfrom tqdm import tqdm\nimport random\n\n# --- CONFIG ---\nBASE_PATH = '/kaggle/input/recodai-luc-scientific-image-forgery-detection' \nTRAIN_IMG_PATH = os.path.join(BASE_PATH, 'train_images')\nTRAIN_MASK_PATH = os.path.join(BASE_PATH, 'train_masks')\nSUPP_IMG_PATH = os.path.join(BASE_PATH, 'supplemental_images')\nSUPP_MASK_PATH = os.path.join(BASE_PATH, 'supplemental_masks')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# --- DATASET WITH NORMALIZATION ---\nclass ScientificForgeryDatasetV2(Dataset):\n    def __init__(self, img_dirs, mask_dirs, img_size=512, transform=True):\n        self.img_size = img_size\n        self.transform = transform\n        self.images = []\n        for img_d, mask_d in zip(img_dirs, mask_dirs):\n            if os.path.exists(img_d):\n                for root, _, files in os.walk(img_d):\n                    for f in files:\n                        if f.lower().endswith(('.png', '.jpg', '.jpeg', '.tif', '.tiff')):\n                            self.images.append((os.path.join(root, f), f, mask_d))\n\n        # THIS IS THE MAGIC SAUCE\n        self.normalize = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((img_size, img_size)),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n\n    def __len__(self):\n        return len(self.images)\n    \n    def augment_data(self, img, mask):\n        if random.random() > 0.5: # Flip H\n            img = cv2.flip(img, 1); mask = cv2.flip(mask, 1)\n        if random.random() > 0.5: # Flip V\n            img = cv2.flip(img, 0); mask = cv2.flip(mask, 0)\n        if random.random() > 0.5: # Rotate\n            img = cv2.rotate(img, cv2.ROTATE_90_CLOCKWISE)\n            mask = cv2.rotate(mask, cv2.ROTATE_90_CLOCKWISE)\n        return img, mask\n\n    def __getitem__(self, idx):\n        img_path, filename, mask_dir_path = self.images[idx]\n        image = cv2.imread(img_path)\n        if image is None: return self._get_empty_batch()\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        file_id = filename.rsplit('.', 1)[0] \n        mask_path = os.path.join(mask_dir_path, file_id + \".npy\")\n        mask = np.zeros(image.shape[:2], dtype=np.float32)\n        \n        if os.path.exists(mask_path):\n            try:\n                m = np.load(mask_path)\n                if m.ndim == 3: m = np.max(m, axis=0)\n                if m.ndim == 3: m = np.max(m, axis=-1)\n                mask = (m > 0).astype(np.float32)\n            except: pass\n        \n        # Resize mask manually (Image is handled by transforms)\n        mask = cv2.resize(mask, (self.img_size, self.img_size), interpolation=cv2.INTER_NEAREST)\n        \n        if self.transform:\n            image, mask = self.augment_data(image, mask)\n\n        # APPLY NORMALIZATION\n        image_tensor = self.normalize(image) # Handles ToTensor and Normalize\n        mask_tensor = torch.from_numpy(mask).unsqueeze(0).float()\n        \n        return image_tensor, mask_tensor\n    \n    def _get_empty_batch(self):\n        return torch.zeros(3, self.img_size, self.img_size), torch.zeros(1, self.img_size, self.img_size)\n\n# --- REST OF TRAINING SETUP ---\n# Load Data\nfull_dataset = ScientificForgeryDatasetV2(\n    img_dirs=[TRAIN_IMG_PATH, SUPP_IMG_PATH], \n    mask_dirs=[TRAIN_MASK_PATH, SUPP_MASK_PATH],\n    img_size=512\n)\ntrain_dataset, val_dataset = random_split(full_dataset, [int(0.9*len(full_dataset)), len(full_dataset)-int(0.9*len(full_dataset))])\nval_dataset.dataset.transform = False\ntrain_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=2)\n\n# Model\nclass NativeResNetUNet(nn.Module):\n    def __init__(self, n_class=1):\n        super().__init__()\n        try:\n            self.base_model = models.resnet34(weights='DEFAULT')\n            print(\"✅ Loaded Pre-Trained ResNet34 Weights!\")\n        except:\n            self.base_model = models.resnet34(weights=None)\n            \n        self.base_layers = list(self.base_model.children())\n        self.layer0 = nn.Sequential(*self.base_layers[:3])\n        self.layer1 = nn.Sequential(*self.base_layers[3:5])\n        self.layer2 = self.base_layers[5]\n        self.layer3 = self.base_layers[6]\n        self.layer4 = self.base_layers[7]\n        self.up4 = self.decoder_block(512, 256)\n        self.up3 = self.decoder_block(256+256, 128)\n        self.up2 = self.decoder_block(128+128, 64)\n        self.up1 = self.decoder_block(64+64, 64)\n        self.final_conv = nn.Conv2d(64, n_class, kernel_size=1)\n        \n    def decoder_block(self, in_channels, out_channels):\n        return nn.Sequential(\n            nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        x = self.layer0(x)\n        layer1 = self.layer1(x)\n        layer2 = self.layer2(layer1)\n        layer3 = self.layer3(layer2)\n        layer4 = self.layer4(layer3)\n        x = self.up4(layer4)\n        x = torch.cat([x, layer3], dim=1)\n        x = self.up3(x)\n        x = torch.cat([x, layer2], dim=1)\n        x = self.up2(x)\n        x = torch.cat([x, layer1], dim=1)\n        x = self.up1(x)\n        return self.final_conv(x)\n\n# Training Loop\nmodel = NativeResNetUNet().to(device)\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\ncriterion = torch.nn.BCEWithLogitsLoss()\nEPOCHS = 15\n\nprint(\"--- STARTING FINAL TRAINING (Normalized) ---\")\nbest_loss = float('inf')\n\nfor epoch in range(EPOCHS):\n    model.train()\n    loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\")\n    for images, masks in loop:\n        images, masks = images.to(device), masks.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        if outputs.shape != masks.shape:\n             outputs = torch.nn.functional.interpolate(outputs, size=masks.shape[2:], mode='bilinear')\n        loss = criterion(outputs, masks)\n        loss.backward()\n        optimizer.step()\n        loop.set_postfix(loss=loss.item())\n\n    # Validation\n    model.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for images, masks in val_loader:\n            images, masks = images.to(device), masks.to(device)\n            outputs = model(images)\n            if outputs.shape != masks.shape:\n                outputs = torch.nn.functional.interpolate(outputs, size=masks.shape[2:], mode='bilinear')\n            val_loss += criterion(outputs, masks).item()\n    \n    avg_val_loss = val_loss / len(val_loader)\n    print(f\"Epoch {epoch+1} Val Loss: {avg_val_loss:.4f}\")\n\n    if avg_val_loss < best_loss:\n        best_loss = avg_val_loss\n        torch.save(model.state_dict(), \"resnet34_normalized.pth\")\n        print(\"New Best Model Saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-17T05:22:50.365507Z","iopub.execute_input":"2025-12-17T05:22:50.366321Z","iopub.status.idle":"2025-12-17T06:09:16.542766Z","shell.execute_reply.started":"2025-12-17T05:22:50.366288Z","shell.execute_reply":"2025-12-17T06:09:16.542013Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}