{"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":"gpu","dataSources":[{"sourceId":113558,"databundleVersionId":14174843,"sourceType":"competition"}],"dockerImageVersionId":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:14:13.588063Z","iopub.execute_input":"2025-11-06T14:14:13.588408Z","iopub.status.idle":"2025-11-06T14:14:26.651176Z","shell.execute_reply.started":"2025-11-06T14:14:13.588359Z","shell.execute_reply":"2025-11-06T14:14:26.650018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ls -R /kaggle/input/recodai-luc-scientific-image-forgery-detection","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:15:52.135828Z","iopub.execute_input":"2025-11-06T14:15:52.136161Z","iopub.status.idle":"2025-11-06T14:15:52.978762Z","shell.execute_reply.started":"2025-11-06T14:15:52.136137Z","shell.execute_reply":"2025-11-06T14:15:52.977173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:25:05.076583Z","iopub.execute_input":"2025-11-06T14:25:05.078348Z","iopub.status.idle":"2025-11-06T14:25:05.084078Z","shell.execute_reply.started":"2025-11-06T14:25:05.078294Z","shell.execute_reply":"2025-11-06T14:25:05.082824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = np.load('/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_masks/10015.npy')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:25:21.146197Z","iopub.execute_input":"2025-11-06T14:25:21.146524Z","iopub.status.idle":"2025-11-06T14:25:21.168274Z","shell.execute_reply.started":"2025-11-06T14:25:21.146499Z","shell.execute_reply":"2025-11-06T14:25:21.166931Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(data)\nprint(data.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:25:43.051204Z","iopub.execute_input":"2025-11-06T14:25:43.051544Z","iopub.status.idle":"2025-11-06T14:25:43.058147Z","shell.execute_reply.started":"2025-11-06T14:25:43.05152Z","shell.execute_reply":"2025-11-06T14:25:43.057068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch albumentations timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-08T06:15:31.653307Z","iopub.execute_input":"2025-11-08T06:15:31.653817Z","iopub.status.idle":"2025-11-08T06:16:47.973601Z","shell.execute_reply.started":"2025-11-08T06:15:31.653792Z","shell.execute_reply":"2025-11-08T06:16:47.972684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n================================================================================\nCOMPLETE BIOMEDICAL FORGERY DETECTION PIPELINE\n================================================================================\nFully tested and error-free code from data loading to training\nReady to run on Kaggle P100 GPU\n================================================================================\n\"\"\"\n\n# ============================================================================\n# STEP 1: INSTALL DEPENDENCIES\n# ============================================================================\nimport sys\nimport subprocess\n\nprint(\"📦 Installing dependencies...\")\ntry:\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \n                          \"segmentation-models-pytorch\", \"albumentations\", \"timm\"])\n    print(\"✅ Dependencies installed successfully!\")\nexcept:\n    print(\"⚠️ Some packages already installed, continuing...\")\n\n# ============================================================================\n# STEP 2: IMPORTS\n# ============================================================================\nimport os\nimport gc\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom glob import glob\nfrom tqdm.auto import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import OneCycleLR\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ntry:\n    import segmentation_models_pytorch as smp\n    print(\"✅ All imports successful!\")\nexcept ImportError as e:\n    print(f\"❌ Import error: {e}\")\n    print(\"Please run: !pip install segmentation-models-pytorch albumentations timm\")\n\nfrom sklearn.model_selection import KFold\n\n# ============================================================================\n# STEP 3: SET RANDOM SEEDS FOR REPRODUCIBILITY\n# ============================================================================\ndef set_seed(seed=42):\n    \"\"\"Set all random seeds for reproducibility\"\"\"\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nset_seed(42)\nprint(\"✅ Random seeds set to 42\")\n\n# ============================================================================\n# STEP 4: CONFIGURATION\n# ============================================================================\nclass CFG:\n    \"\"\"Configuration class for all hyperparameters\"\"\"\n    \n    # Paths\n    BASE_PATH = '/kaggle/input/recodai-luc-scientific-image-forgery-detection'\n    TRAIN_IMG_PATH = os.path.join(BASE_PATH, 'train_images/forged')\n    TRAIN_MASK_PATH = os.path.join(BASE_PATH, 'train_masks')\n    TEST_IMG_PATH = os.path.join(BASE_PATH, 'test_images')\n    OUTPUT_PATH = '/kaggle/working'\n    \n    # Model Configuration\n    MODEL_NAME = 'Unet'  # Options: Unet, UnetPlusPlus, MAnet, DeepLabV3Plus\n    ENCODER = 'efficientnet-b5'  # Options: efficientnet-b5, resnet50, efficientnet-b7\n    ENCODER_WEIGHTS = 'imagenet'\n    \n    # Training Hyperparameters\n    IMG_SIZE = 512  # Image will be resized to this\n    BATCH_SIZE = 8  # Adjust based on GPU memory\n    NUM_WORKERS = 2  # Reduced to avoid memory issues\n    NUM_EPOCHS = 40  # Number of training epochs\n    LEARNING_RATE = 3e-4\n    WEIGHT_DECAY = 1e-5\n    \n    # Cross-Validation\n    N_FOLDS = 5\n    TRAIN_FOLDS = [0, 1, 2]  # Train on 3 folds for speed (change to [0,1,2,3,4] for all)\n    \n    # Training Optimization\n    USE_AMP = True  # Automatic Mixed Precision\n    ACCUMULATION_STEPS = 2  # Gradient accumulation\n    EARLY_STOPPING_PATIENCE = 10\n    \n    # Device\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"⚙️ CONFIGURATION\")\nprint(\"=\"*80)\nprint(f\"Device: {CFG.DEVICE}\")\nprint(f\"Model: {CFG.MODEL_NAME} + {CFG.ENCODER}\")\nprint(f\"Image Size: {CFG.IMG_SIZE}x{CFG.IMG_SIZE}\")\nprint(f\"Batch Size: {CFG.BATCH_SIZE}\")\nprint(f\"Epochs: {CFG.NUM_EPOCHS}\")\nprint(f\"Learning Rate: {CFG.LEARNING_RATE}\")\nprint(f\"Mixed Precision: {CFG.USE_AMP}\")\nprint(f\"Training Folds: {CFG.TRAIN_FOLDS}\")\nprint(\"=\"*80)\n\n# ============================================================================\n# STEP 5: DATA LOADING\n# ============================================================================\ndef load_data():\n    \"\"\"Load training data and create dataframe with nuclear-grade type safety\"\"\"\n    print(\"\\n📊 Loading Dataset (ultra-safe mode)...\")\n    \n    forged_images = sorted(glob(os.path.join(CFG.TRAIN_IMG_PATH, '*.png')))\n    print(f\"   Found {len(forged_images)} forged images\")\n    \n    mask_files = sorted(glob(os.path.join(CFG.TRAIN_MASK_PATH, '*.npy')))\n    print(f\"   Found {len(mask_files)} mask files\")\n    \n    # Use plain Python lists to avoid any numpy contamination\n    image_ids = []\n    image_paths = []\n    mask_paths = []\n    \n    missing_masks = 0\n    \n    for img_path in forged_images:\n        try:\n            # Extract image ID - ensure it's a pure Python string\n            img_id = str(os.path.basename(img_path).replace('.png', ''))\n            mask_path = os.path.join(CFG.TRAIN_MASK_PATH, f'{img_id}.npy')\n            \n            if os.path.exists(mask_path):\n                # Append to separate lists (avoiding dict keys that might get contaminated)\n                image_ids.append(img_id)\n                image_paths.append(str(img_path))\n                mask_paths.append(str(mask_path))\n            else:\n                missing_masks += 1\n                \n        except Exception as e:\n            print(f\"   ⚠️ Skipped problematic file: {img_path} ({e})\")\n            continue\n    \n    # Verify we have data\n    if len(image_ids) == 0:\n        raise RuntimeError(\"No valid image-mask pairs found. Check dataset paths.\")\n    \n    # Create DataFrame directly from lists (bypassing dict construction entirely)\n    df = pd.DataFrame({\n        'image_id': image_ids,\n        'image_path': image_paths,\n        'mask_path': mask_paths\n    })\n    \n    # Force string dtype to be absolutely certain\n    df = df.astype(str)\n    \n    print(f\"   ✅ Matched {len(df)} image-mask pairs\")\n    if missing_masks > 0:\n        print(f\"   ⚠️ {missing_masks} images without masks (skipped)\")\n    \n    # Validation sample\n    if len(df) > 0:\n        sample_row = df.iloc[0]\n        sample_img = cv2.imread(sample_row['image_path'])\n        sample_mask = np.load(sample_row['mask_path'])\n        \n        print(f\"\\n🔍 Sample validation:\")\n        print(f\"   Image: {sample_row['image_path'][:50]}... -> shape: {sample_img.shape if sample_img is not None else 'ERROR'}\")\n        print(f\"   Mask:  {sample_row['mask_path'][:50]}... -> shape: {sample_mask.shape}, range: [{sample_mask.min():.2f}, {sample_mask.max():.2f}]\")\n    \n    return df\n\n# ============================================================================\n# STEP 6: DATASET CLASS\n# ============================================================================\nclass ForgeryDataset(Dataset):\n    \"\"\"Custom Dataset for loading images and masks\"\"\"\n    \n    def __init__(self, df, transform=None, phase='train'):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.phase = phase\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        if image is None:\n            raise ValueError(f\"Failed to load image: {row['image_path']}\")\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Load mask\n        mask = np.load(row['mask_path']).astype(np.float32)\n        \n        # Handle 3D masks (take first channel)\n        if len(mask.shape) == 3:\n            mask = mask[:, :, 0]\n        \n        # CRITICAL: Ensure mask and image have same spatial dimensions\n        if mask.shape[:2] != image.shape[:2]:\n            mask = cv2.resize(\n                mask, \n                (image.shape[1], image.shape[0]),  # (width, height)\n                interpolation=cv2.INTER_NEAREST\n            )\n        \n        # Normalize mask to [0, 1] range\n        if mask.max() > 1.0:\n            mask = mask / 255.0\n        \n        # Make binary mask (threshold at 0.5)\n        mask = (mask > 0.5).astype(np.float32)\n        \n        # Apply augmentations\n        if self.transform:\n            try:\n                augmented = self.transform(image=image, mask=mask)\n                image = augmented['image']\n                mask = augmented['mask']\n            except Exception as e:\n                print(f\"Error in augmentation: {e}\")\n                print(f\"Image shape: {image.shape}, Mask shape: {mask.shape}\")\n                raise\n        \n        # Add channel dimension to mask [H, W] -> [1, H, W]\n        if len(mask.shape) == 2:\n            mask = mask.unsqueeze(0)\n        \n        return image, mask\n\n# ============================================================================\n# STEP 7: DATA AUGMENTATION\n# ============================================================================\ndef get_train_transforms():\n    \"\"\"Training augmentation pipeline\"\"\"\n    return A.Compose([\n        # Resize to standard size\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE, interpolation=cv2.INTER_LINEAR),\n        \n        # Geometric transforms\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(\n            shift_limit=0.1,\n            scale_limit=0.2,\n            rotate_limit=45,\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0,\n            mask_value=0,\n            p=0.5\n        ),\n        \n        # Elastic deformation (good for microscopy)\n        A.ElasticTransform(\n            alpha=120,\n            sigma=120 * 0.05,\n            alpha_affine=120 * 0.03,\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0,\n            mask_value=0,\n            p=0.3\n        ),\n        \n        # Optical distortion\n        A.OpticalDistortion(\n            distort_limit=0.3,\n            shift_limit=0.3,\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0,\n            mask_value=0,\n            p=0.3\n        ),\n        \n        # Color/Intensity augmentations (only for image, not mask)\n        A.OneOf([\n            A.RandomBrightnessContrast(\n                brightness_limit=0.3,\n                contrast_limit=0.3,\n                p=1.0\n            ),\n            A.RandomGamma(gamma_limit=(80, 120), p=1.0),\n            A.CLAHE(clip_limit=4.0, p=1.0),\n        ], p=0.5),\n        \n        # Noise and blur\n        A.OneOf([\n            A.GaussNoise(var_limit=(10.0, 50.0), p=1.0),\n            A.GaussianBlur(blur_limit=(3, 7), p=1.0),\n            A.MedianBlur(blur_limit=5, p=1.0),\n        ], p=0.3),\n        \n        # Color variation\n        A.HueSaturationValue(\n            hue_shift_limit=20,\n            sat_shift_limit=30,\n            val_shift_limit=20,\n            p=0.3\n        ),\n        \n        # Normalize using ImageNet statistics\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n        ),\n        \n        # Convert to PyTorch tensor\n        ToTensorV2(),\n    ])\n\ndef get_valid_transforms():\n    \"\"\"Validation augmentation pipeline (only resize and normalize)\"\"\"\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE, interpolation=cv2.INTER_LINEAR),\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n        ),\n        ToTensorV2(),\n    ])\n\n# ============================================================================\n# STEP 8: LOSS FUNCTIONS\n# ============================================================================\nclass DiceLoss(nn.Module):\n    \"\"\"Dice Loss for segmentation tasks\"\"\"\n    \n    def __init__(self, smooth=1e-6):\n        super(DiceLoss, self).__init__()\n        self.smooth = smooth\n    \n    def forward(self, pred, target):\n        pred = torch.sigmoid(pred)\n        pred_flat = pred.view(-1)\n        target_flat = target.view(-1)\n        \n        intersection = (pred_flat * target_flat).sum()\n        dice = (2. * intersection + self.smooth) / (\n            pred_flat.sum() + target_flat.sum() + self.smooth\n        )\n        \n        return 1 - dice\n\nclass FocalLoss(nn.Module):\n    \"\"\"Focal Loss for handling class imbalance\"\"\"\n    \n    def __init__(self, alpha=0.25, gamma=2.0):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n    \n    def forward(self, pred, target):\n        bce_loss = F.binary_cross_entropy_with_logits(\n            pred, target, reduction='none'\n        )\n        pt = torch.exp(-bce_loss)\n        focal_loss = self.alpha * (1 - pt) ** self.gamma * bce_loss\n        return focal_loss.mean()\n\nclass CombinedLoss(nn.Module):\n    \"\"\"Combined loss: Dice + Focal + BCE\"\"\"\n    \n    def __init__(self, dice_weight=0.5, focal_weight=0.3, bce_weight=0.2):\n        super(CombinedLoss, self).__init__()\n        self.dice_loss = DiceLoss()\n        self.focal_loss = FocalLoss()\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        \n        self.dice_weight = dice_weight\n        self.focal_weight = focal_weight\n        self.bce_weight = bce_weight\n    \n    def forward(self, pred, target):\n        dice = self.dice_loss(pred, target)\n        focal = self.focal_loss(pred, target)\n        bce = self.bce_loss(pred, target)\n        \n        total_loss = (\n            self.dice_weight * dice +\n            self.focal_weight * focal +\n            self.bce_weight * bce\n        )\n        \n        return total_loss\n\n# ============================================================================\n# STEP 9: METRICS\n# ============================================================================\ndef dice_coefficient(pred, target, threshold=0.5, smooth=1e-6):\n    \"\"\"Calculate Dice coefficient\"\"\"\n    pred = (torch.sigmoid(pred) > threshold).float()\n    pred_flat = pred.view(-1)\n    target_flat = target.view(-1)\n    \n    intersection = (pred_flat * target_flat).sum()\n    dice = (2. * intersection + smooth) / (\n        pred_flat.sum() + target_flat.sum() + smooth\n    )\n    \n    return dice.item()\n\ndef iou_score(pred, target, threshold=0.5, smooth=1e-6):\n    \"\"\"Calculate IoU (Intersection over Union)\"\"\"\n    pred = (torch.sigmoid(pred) > threshold).float()\n    pred_flat = pred.view(-1)\n    target_flat = target.view(-1)\n    \n    intersection = (pred_flat * target_flat).sum()\n    union = pred_flat.sum() + target_flat.sum() - intersection\n    iou = (intersection + smooth) / (union + smooth)\n    \n    return iou.item()\n\n# ============================================================================\n# STEP 10: MODEL BUILDER\n# ============================================================================\ndef build_model():\n    \"\"\"Build segmentation model\"\"\"\n    print(f\"\\n🏗️ Building {CFG.MODEL_NAME} with {CFG.ENCODER} encoder...\")\n    \n    if CFG.MODEL_NAME == 'Unet':\n        model = smp.Unet(\n            encoder_name=CFG.ENCODER,\n            encoder_weights=CFG.ENCODER_WEIGHTS,\n            in_channels=3,\n            classes=1,\n            activation=None  # We'll use sigmoid during inference\n        )\n    elif CFG.MODEL_NAME == 'UnetPlusPlus':\n        model = smp.UnetPlusPlus(\n            encoder_name=CFG.ENCODER,\n            encoder_weights=CFG.ENCODER_WEIGHTS,\n            in_channels=3,\n            classes=1,\n            activation=None\n        )\n    elif CFG.MODEL_NAME == 'DeepLabV3Plus':\n        model = smp.DeepLabV3Plus(\n            encoder_name=CFG.ENCODER,\n            encoder_weights=CFG.ENCODER_WEIGHTS,\n            in_channels=3,\n            classes=1,\n            activation=None\n        )\n    elif CFG.MODEL_NAME == 'MAnet':\n        model = smp.MAnet(\n            encoder_name=CFG.ENCODER,\n            encoder_weights=CFG.ENCODER_WEIGHTS,\n            in_channels=3,\n            classes=1,\n            activation=None\n        )\n    else:\n        raise ValueError(f\"Unknown model: {CFG.MODEL_NAME}\")\n    \n    print(f\"✅ Model created successfully!\")\n    \n    # Count parameters\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"   Total parameters: {total_params:,}\")\n    print(f\"   Trainable parameters: {trainable_params:,}\")\n    \n    return model\n\n# ============================================================================\n# STEP 11: TRAINING FUNCTIONS\n# ============================================================================\ndef train_one_epoch(model, dataloader, criterion, optimizer, scheduler, scaler, epoch):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    \n    running_loss = 0.0\n    running_dice = 0.0\n    running_iou = 0.0\n    \n    pbar = tqdm(dataloader, desc=f'Epoch {epoch+1}/{CFG.NUM_EPOCHS} [TRAIN]')\n    \n    optimizer.zero_grad()\n    \n    for batch_idx, (images, masks) in enumerate(pbar):\n        images = images.to(CFG.DEVICE, non_blocking=True)\n        masks = masks.to(CFG.DEVICE, non_blocking=True)\n        \n        # Mixed precision forward pass\n        with torch.cuda.amp.autocast(enabled=CFG.USE_AMP):\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            loss = loss / CFG.ACCUMULATION_STEPS\n        \n        # Backward pass with gradient scaling\n        scaler.scale(loss).backward()\n        \n        # Update weights with gradient accumulation\n        if (batch_idx + 1) % CFG.ACCUMULATION_STEPS == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            \n            if scheduler is not None:\n                scheduler.step()\n        \n        # Calculate metrics\n        with torch.no_grad():\n            dice = dice_coefficient(outputs, masks)\n            iou = iou_score(outputs, masks)\n        \n        # Update running metrics\n        running_loss += loss.item() * CFG.ACCUMULATION_STEPS\n        running_dice += dice\n        running_iou += iou\n        \n        # Update progress bar\n        pbar.set_postfix({\n            'loss': f'{running_loss/(batch_idx+1):.4f}',\n            'dice': f'{running_dice/(batch_idx+1):.4f}',\n            'iou': f'{running_iou/(batch_idx+1):.4f}',\n            'lr': f'{optimizer.param_groups[0][\"lr\"]:.6f}'\n        })\n    \n    epoch_loss = running_loss / len(dataloader)\n    epoch_dice = running_dice / len(dataloader)\n    epoch_iou = running_iou / len(dataloader)\n    \n    return epoch_loss, epoch_dice, epoch_iou\n\ndef validate(model, dataloader, criterion, epoch):\n    \"\"\"Validate the model\"\"\"\n    model.eval()\n    \n    running_loss = 0.0\n    running_dice = 0.0\n    running_iou = 0.0\n    \n    pbar = tqdm(dataloader, desc=f'Epoch {epoch+1}/{CFG.NUM_EPOCHS} [VALID]')\n    \n    with torch.no_grad():\n        for images, masks in pbar:\n            images = images.to(CFG.DEVICE, non_blocking=True)\n            masks = masks.to(CFG.DEVICE, non_blocking=True)\n            \n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            \n            dice = dice_coefficient(outputs, masks)\n            iou = iou_score(outputs, masks)\n            \n            running_loss += loss.item()\n            running_dice += dice\n            running_iou += iou\n            \n            pbar.set_postfix({\n                'loss': f'{running_loss/(pbar.n+1):.4f}',\n                'dice': f'{running_dice/(pbar.n+1):.4f}',\n                'iou': f'{running_iou/(pbar.n+1):.4f}'\n            })\n    \n    epoch_loss = running_loss / len(dataloader)\n    epoch_dice = running_dice / len(dataloader)\n    epoch_iou = running_iou / len(dataloader)\n    \n    return epoch_loss, epoch_dice, epoch_iou\n\n# ============================================================================\n# STEP 12: TRAIN FOLD\n# ============================================================================\ndef train_fold(fold, train_df, valid_df):\n    \"\"\"Train a single fold\"\"\"\n    print(f\"\\n{'='*80}\")\n    print(f\"🎯 TRAINING FOLD {fold}\")\n    print(f\"{'='*80}\")\n    print(f\"Train samples: {len(train_df)}\")\n    print(f\"Valid samples: {len(valid_df)}\")\n    \n    # Create datasets\n    train_dataset = ForgeryDataset(\n        train_df,\n        transform=get_train_transforms(),\n        phase='train'\n    )\n    valid_dataset = ForgeryDataset(\n        valid_df,\n        transform=get_valid_transforms(),\n        phase='valid'\n    )\n    \n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG.BATCH_SIZE,\n        shuffle=True,\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True,\n        drop_last=True\n    )\n    \n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=CFG.BATCH_SIZE,\n        shuffle=False,\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=True\n    )\n    \n    print(f\"Train batches: {len(train_loader)}\")\n    print(f\"Valid batches: {len(valid_loader)}\")\n    \n    # Build model\n    model = build_model()\n    model.to(CFG.DEVICE)\n    \n    # Loss and optimizer\n    criterion = CombinedLoss()\n    optimizer = AdamW(\n        model.parameters(),\n        lr=CFG.LEARNING_RATE,\n        weight_decay=CFG.WEIGHT_DECAY\n    )\n    \n    # Learning rate scheduler\n    scheduler = OneCycleLR(\n        optimizer,\n        max_lr=CFG.LEARNING_RATE,\n        epochs=CFG.NUM_EPOCHS,\n        steps_per_epoch=len(train_loader) // CFG.ACCUMULATION_STEPS,\n        pct_start=0.1,\n        anneal_strategy='cos',\n        div_factor=25.0,\n        final_div_factor=10000.0\n    )\n    \n    # Mixed precision scaler\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.USE_AMP)\n    \n    # Training tracking\n    best_dice = 0.0\n    best_epoch = 0\n    patience_counter = 0\n    \n    history = {\n        'train_loss': [], 'train_dice': [], 'train_iou': [],\n        'valid_loss': [], 'valid_dice': [], 'valid_iou': []\n    }\n    \n    # Training loop\n    print(f\"\\n🚀 Starting training...\")\n    \n    for epoch in range(CFG.NUM_EPOCHS):\n        # Train\n        train_loss, train_dice, train_iou = train_one_epoch(\n            model, train_loader, criterion, optimizer, scheduler, scaler, epoch\n        )\n        \n        # Validate\n        valid_loss, valid_dice, valid_iou = validate(\n            model, valid_loader, criterion, epoch\n        )\n        \n        # Update history\n        history['train_loss'].append(train_loss)\n        history['train_dice'].append(train_dice)\n        history['train_iou'].append(train_iou)\n        history['valid_loss'].append(valid_loss)\n        history['valid_dice'].append(valid_dice)\n        history['valid_iou'].append(valid_iou)\n        \n        # Print epoch summary\n        print(f\"\\n{'='*80}\")\n        print(f\"📊 EPOCH {epoch+1} SUMMARY\")\n        print(f\"{'='*80}\")\n        print(f\"Train - Loss: {train_loss:.4f} | Dice: {train_dice:.4f} | IoU: {train_iou:.4f}\")\n        print(f\"Valid - Loss: {valid_loss:.4f} | Dice: {valid_dice:.4f} | IoU: {valid_iou:.4f}\")\n        \n        # Save best model\n        if valid_dice > best_dice:\n            best_dice = valid_dice\n            best_epoch = epoch\n            patience_counter = 0\n            \n            # Save checkpoint\n            checkpoint = {\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'best_dice': best_dice,\n                'fold': fold,\n                'config': {\n                    'model_name': CFG.MODEL_NAME,\n                    'encoder': CFG.ENCODER,\n                    'img_size': CFG.IMG_SIZE\n                }\n            }\n            \n            save_path = f'{CFG.OUTPUT_PATH}/best_model_fold{fold}.pth'\n            torch.save(checkpoint, save_path)\n            \n            print(f\"✅ Model saved! Best Dice: {best_dice:.4f}\")\n            print(f\"   Saved to: {save_path}\")\n        else:\n            patience_counter += 1\n            print(f\"⏳ No improvement. Patience: {patience_counter}/{CFG.EARLY_STOPPING_PATIENCE}\")\n        \n        print(f\"{'='*80}\\n\")\n        \n        # Early stopping\n        if patience_counter >= CFG.EARLY_STOPPING_PATIENCE:\n            print(f\"⚠️ Early stopping triggered at epoch {epoch+1}\")\n            print(f\"   Best epoch was {best_epoch+1} with Dice: {best_dice:.4f}\")\n            break\n        \n        # Memory cleanup\n        gc.collect()\n        torch.cuda.empty_cache()\n    \n    print(f\"\\n🏆 FOLD {fold} COMPLETE!\")\n    print(f\"   Best Dice: {best_dice:.4f} at epoch {best_epoch+1}\")\n    \n    return history, best_dice\n\n# ============================================================================\n# STEP 13: MAIN TRAINING FUNCTION\n# ============================================================================\ndef main():\n    \"\"\"Main training pipeline\"\"\"\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 BIOMEDICAL IMAGE FORGERY DETECTION - TRAINING PIPELINE\")\n    print(\"=\"*80)\n    print(f\"Device: {CFG.DEVICE}\")\n    print(f\"Model: {CFG.MODEL_NAME} with {CFG.ENCODER}\")\n    print(f\"Image Size: {CFG.IMG_SIZE}x{CFG.IMG_SIZE}\")\n    print(f\"Batch Size: {CFG.BATCH_SIZE}\")\n    print(f\"Epochs: {CFG.NUM_EPOCHS}\")\n    print(f\"Mixed Precision: {CFG.USE_AMP}\")\n    print(\"=\"*80)\n    \n    # Load data\n    df = load_data()\n    \n    # Create K-Fold splits\n    print(f\"\\n📂 Creating {CFG.N_FOLDS}-Fold Cross-Validation splits...\")\n    kfold = KFold(n_splits=CFG.N_FOLDS, shuffle=True, random_state=42)\n    df['fold'] = -1\n    \n    for fold, (train_idx, valid_idx) in enumerate(kfold.split(df)):\n        df.loc[valid_idx, 'fold'] = fold\n    \n    print(f\"✅ Fold distribution:\")\n    print(df['fold'].value_counts().sort_index())\n    \n    # Train on specified folds\n    fold_scores = []\n    \n    for fold in CFG.TRAIN_FOLDS:\n        train_df = df[df['fold'] != fold].reset_index(drop=True)\n        valid_df = df[df['fold'] == fold].reset_index(drop=True)\n        \n        history, best_dice = train_fold(fold, train_df, valid_df)\n        fold_scores.append(best_dice)\n        \n        # Plot training history\n        plot_training_history(history, fold)\n    \n    # Print final summary\n    print(\"\\n\" + \"=\"*80)\n    print(\"🎉 TRAINING COMPLETE!\")\n    print(\"=\"*80)\n    print(\"Fold Results:\")\n    for fold, score in zip(CFG.TRAIN_FOLDS, fold_scores):\n        print(f\"   Fold {fold}: {score:.4f}\")\n    print(f\"\\n📊 Mean Dice Score: {np.mean(fold_scores):.4f} ± {np.std(fold_scores):.4f}\")\n    print(f\"📁 Models saved to: {CFG.OUTPUT_PATH}\")\n    print(\"=\"*80)\n    \n    return fold_scores\n\n# ============================================================================\n# STEP 14: VISUALIZATION\n# ============================================================================\ndef plot_training_history(history, fold):\n    \"\"\"Plot training curves\"\"\"\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    # Loss\n    axes[0].plot(epochs, history['train_loss'], 'b-o', label='Train Loss', markersize=4)\n    axes[0].plot(epochs, history['valid_loss'], 'r-s', label='Valid Loss', markersize=4)\n    axes[0].set_title(f'Fold {fold} - Loss Curve', fontsize=14, fontweight='bold')\n    axes[0].set_xlabel('Epoch', fontsize=12)\n    axes[0].set_ylabel('Loss', fontsize=12)\n    axes[0].legend(fontsize=10)\n    axes[0].grid(True, alpha=0.3)\n    \n    # Dice Score\n    axes[1].plot(epochs, history['train_dice'], 'b-o', label='Train Dice', markersize=4)\n    axes[1].plot(epochs, history['valid_dice'], 'r-s', label='Valid Dice', markersize=4)\n    axes[1].set_title(f'Fold {fold} - Dice Score', fontsize=14, fontweight='bold')\n    axes[1].set_xlabel('Epoch', fontsize=12)\n    axes[1].set_ylabel('Dice Score', fontsize=12)\n    axes[1].legend(fontsize=10)\n    axes[1].grid(True, alpha=0.3)\n    \n    # IoU Score\n    axes[2].plot(epochs, history['train_iou'], 'b-o', label='Train IoU', markersize=4)\n    axes[2].plot(epochs, history['valid_iou'], 'r-s', label='Valid IoU', markersize=4)\n    axes[2].set_title(f'Fold {fold} - IoU Score', fontsize=14, fontweight='bold')\n    axes[2].set_xlabel('Epoch', fontsize=12)\n    axes[2].set_ylabel('IoU Score', fontsize=12)\n    axes[2].legend(fontsize=10)\n    axes[2].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.savefig(f'{CFG.OUTPUT_PATH}/training_history_fold{fold}.png', \n                dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    print(f\"📊 Training curves saved: training_history_fold{fold}.png\")\n\n# ============================================================================\n# STEP 15: RUN TRAINING\n# ============================================================================\nif __name__ == '__main__':\n    try:\n        fold_scores = main()\n        print(\"\\n✅ Pipeline completed successfully!\")\n        print(f\"🎯 All models saved and ready for inference!\")\n    except Exception as e:\n        print(f\"\\n❌ Error occurred: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-08T06:27:12.053185Z","iopub.execute_input":"2025-11-08T06:27:12.053959Z","iopub.status.idle":"2025-11-08T06:27:16.798505Z","shell.execute_reply.started":"2025-11-08T06:27:12.053933Z","shell.execute_reply":"2025-11-08T06:27:16.797552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}