{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":113558,"databundleVersionId":14174843,"sourceType":"competition"},{"sourceId":265158055,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n# sys.path.append(\"/kaggle/input/smp-lib-zip/\") # Kept as-is for Kaggle\nimport warnings\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nimport cv2\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport segmentation_models_pytorch as smp\nimport time\nfrom collections import defaultdict\nimport copy\nimport random\n\n# ===================================================================\n# CONFIGURATION\n# ===================================================================\nclass Config:\n    \"\"\"\n    Configuration class for all hyperparameters and paths.\n    \"\"\"\n    # Paths\n    BASE_PATH = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection\"\n    TRAIN_AUTHENTIC_PATH = os.path.join(BASE_PATH, \"train_images/authentic\")\n    TRAIN_FORGED_PATH = os.path.join(BASE_PATH, \"train_images/forged\")\n    TRAIN_MASKS_PATH = os.path.join(BASE_PATH, \"train_masks\")\n    TEST_IMAGES_PATH = os.path.join(BASE_PATH, \"test_images\")\n    \n    # Model & Training\n    TARGET_SIZE = (512, 512)\n    MODEL_NAME = 'resnet34'\n    \n    # --- IMPROVEMENT: Using 'imagenet' weights for transfer learning ---\n    # This usually leads to faster convergence and better performance\n    # than training from scratch (None).\n    ENCODER_WEIGHTS = 'imagenet' \n    \n    LEARNING_RATE = 3e-4\n    WEIGHT_DECAY = 1e-5\n    NUM_EPOCHS = 25 # You can adjust this\n    BATCH_SIZE = 6\n    NUM_WORKERS = 2\n    \n    # General\n    SEED = 42\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ===================================================================\n# UTILS\n# ===================================================================\n\ndef set_seed(seed=42):\n    \"\"\"\n    Sets the seed for reproducibility.\n    \"\"\"\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed) # if using multi-GPU\n    # Ensure deterministic behavior\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\ndef get_image_dimensions(image_path):\n    \"\"\"Helper to get image dimensions.\"\"\"\n    try:\n        with Image.open(image_path) as img:\n            return img.size\n    except Exception as e:\n        print(f\"Error opening image {image_path}: {e}\")\n        return (0, 0)\n\n# ===================================================================\n# DATA LOADING AND EDA\n# ===================================================================\n\ndef run_eda(cfg):\n    \"\"\"\n    Encapsulates the Data Loading and EDA steps.\n    \"\"\"\n    print(\"=\"*60)\n    print(\" DATA LOADING AND BASIC STATISTICS\")\n    print(\"=\"*60)\n    \n    authentic_images = sorted([f for f in os.listdir(cfg.TRAIN_AUTHENTIC_PATH) if f.endswith('.png')])\n    forged_images = sorted([f for f in os.listdir(cfg.TRAIN_FORGED_PATH) if f.endswith('.png')])\n    mask_files = sorted([f for f in os.listdir(cfg.TRAIN_MASKS_PATH) if f.endswith('.npy')])\n    \n    print(f\"Number of AUTHENTIC images: {len(authentic_images)}\")\n    print(f\"Number of FORGED images: {len(forged_images)}\")\n    print(f\"Number of MASK files: {len(mask_files)}\")\n    print(f\"Total TRAIN images: {len(authentic_images) + len(forged_images)}\")\n    \n    # --- Analyzing Image Dimensions (Sampled) ---\n    print(\"\\nAnalyzing image dimensions (sampling 50 of each)...\")\n    authentic_dims = []\n    for img_name in tqdm(authentic_images[:50], desc=\"Authentic Dims\"):\n        img_path = os.path.join(cfg.TRAIN_AUTHENTIC_PATH, img_name)\n        w, h = get_image_dimensions(img_path)\n        authentic_dims.append((w, h, w*h))\n\n    forged_dims = []\n    for img_name in tqdm(forged_images[:50], desc=\"Forged Dims\"):\n        img_path = os.path.join(cfg.TRAIN_FORGED_PATH, img_name)\n        w, h = get_image_dimensions(img_path)\n        forged_dims.append((w, h, w*h))\n\n    # --- VISUALIZING FORGED REGIONS WITH OVERLAY ---\n    print(\"\\nVisualizing sample forged images...\")\n    visualize_forgery_overlay(forged_images, cfg.TRAIN_MASKS_PATH, n_samples=3)\n\n    # --- ANALYZING MASKS (Sampled) ---\n    print(\"\\nAnalyzing mask statistics (sampling 100 masks)...\")\n    mask_stats = analyze_masks(mask_files[:100], cfg.TRAIN_MASKS_PATH)\n    \n    print(f\"\\n--- MASK STATISTICS (Sample) ---\")\n    print(f\"  Mean regions: {np.mean(mask_stats['n_regions']):.2f}\")\n    print(f\"  Mean forged area: {np.mean(mask_stats['forged_percentage']):.2f}%\")\n    \n    return authentic_images, forged_images\n\ndef visualize_forgery_overlay(forged_list, masks_path, n_samples=3): \n    fig, axes = plt.subplots(n_samples, 3, figsize=(18, 5*n_samples))\n    fig.suptitle('Forged Images with Mask Overlay', fontsize=16, fontweight='bold')\n    \n    for i in range(n_samples):\n        img_path = os.path.join(Config.TRAIN_FORGED_PATH, forged_list[i])\n        img = np.array(Image.open(img_path))\n        mask_name = forged_list[i].replace('.png', '.npy')\n        mask_path = os.path.join(masks_path, mask_name)\n        \n        if os.path.exists(mask_path):\n            mask = np.load(mask_path)\n            if len(mask.shape) == 3:\n                combined_mask = np.max(mask, axis=0)\n                n_regions = mask.shape[0]\n            else:\n                combined_mask = mask\n                n_regions = 1\n            \n            axes[i, 0].imshow(img)\n            axes[i, 0].set_title(f'Original Image\\n{forged_list[i]}', fontsize=10)\n            axes[i, 0].axis('off')\n\n            axes[i, 1].imshow(combined_mask, cmap='hot')\n            axes[i, 1].set_title(f'Forgery Mask\\n{n_regions} region(s)', fontsize=10)\n            axes[i, 1].axis('off')\n            \n            overlay = img.copy()\n            if len(img.shape) == 2:  \n                overlay = cv2.cvtColor(overlay, cv2.COLOR_GRAY2RGB)\n            \n            red_mask = np.zeros_like(overlay)\n            if len(red_mask.shape) == 3:\n                red_mask[:,:,0] = combined_mask * 255\n            \n            alpha = 0.4\n            blended = cv2.addWeighted(overlay, 1-alpha, red_mask, alpha, 0)\n            \n            axes[i, 2].imshow(blended)\n            axes[i, 2].set_title('Overlay (Red = Forged)', fontsize=10)\n            axes[i, 2].axis('off')\n        else:\n            for j in range(3):\n                axes[i, j].text(0.5, 0.5, 'MASK NOT FOUND', ha='center', va='center', fontsize=12)\n                axes[i, j].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\ndef analyze_masks(mask_files_sample, masks_path):\n    mask_stats = defaultdict(list)\n    for mask_file in tqdm(mask_files_sample, desc=\"Processing masks\"): \n        mask_path = os.path.join(masks_path, mask_file)\n        mask = np.load(mask_path)\n        \n        if len(mask.shape) == 3:\n            n_regions = mask.shape[0]\n            combined_mask = np.max(mask, axis=0)\n        else:\n            n_regions = 1\n            combined_mask = mask\n        \n        mask_stats['n_regions'].append(n_regions)\n        \n        total_pixels = combined_mask.shape[0] * combined_mask.shape[1]\n        if total_pixels == 0: continue\n        \n        forged_pixels = np.sum(combined_mask > 0)\n        forged_pct = (forged_pixels / total_pixels) * 100\n        \n        mask_stats['total_forged_pixels'].append(forged_pixels)\n        mask_stats['forged_percentage'].append(forged_pct)\n    \n    return mask_stats\n\n# ===================================================================\n# DATA PREPARATION\n# ===================================================================\n\ndef get_dataframes(authentic_images, forged_images, cfg):\n    \"\"\"\n    Creates and splits the training dataframe.\n    \"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\" PREPARING DATASETS\")\n    print(\"=\"*60)\n    \n    image_names, image_paths, mask_paths, labels, categories = [], [], [], [], []\n\n    for img_name in authentic_images:\n        image_names.append(img_name)\n        image_paths.append(os.path.join(cfg.TRAIN_AUTHENTIC_PATH, img_name))\n        mask_paths.append(None)\n        labels.append(0)\n        categories.append('authentic')\n\n    for img_name in forged_images:\n        image_names.append(img_name)\n        image_paths.append(os.path.join(cfg.TRAIN_FORGED_PATH, img_name))\n        mask_name = img_name.replace('.png', '.npy')\n        mask_path = os.path.join(cfg.TRAIN_MASKS_PATH, mask_name)\n        mask_paths.append(mask_path if os.path.exists(mask_path) else None)\n        labels.append(1)\n        categories.append('forged')\n\n    train_df = pd.DataFrame({\n        'image_name': image_names,\n        'image_path': image_paths,\n        'mask_path': mask_paths,\n        'label': labels,\n        'category': categories\n    })\n\n    # Split the data\n    from sklearn.model_selection import train_test_split\n    train_data, val_data = train_test_split(\n        train_df, \n        test_size=0.2, \n        random_state=Config.SEED,\n        stratify=train_df['label']\n    )\n    \n    print(f\"Total samples: {len(train_df)}\")\n    print(f\"Train set size: {len(train_data)}\")\n    print(f\"Validation set size: {len(val_data)}\")\n    print(f\"Train distribution:\\n{train_data['category'].value_counts(normalize=True)}\")\n    print(f\"Validation distribution:\\n{val_data['category'].value_counts(normalize=True)}\")\n    \n    return train_data.reset_index(drop=True), val_data.reset_index(drop=True)\n\n# ===================================================================\n# AUGMENTATIONS & DATASET\n# ===================================================================\n\ndef get_transforms(cfg):\n    \"\"\"\n    Returns augmentation pipelines for training and validation.\n    \"\"\"\n    train_transform = A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.Rotate(limit=30, p=0.5, border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.ShiftScaleRotate(\n            shift_limit=0.1,\n            scale_limit=0.15,\n            rotate_limit=15,\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0,\n            p=0.5\n        ),\n        A.OneOf([\n            A.OpticalDistortion(distort_limit=0.1, p=1.0),\n            A.GridDistortion(num_steps=5, distort_limit=0.1, p=1.0),\n            A.ElasticTransform(alpha=1, sigma=50, p=1.0),\n        ], p=0.3),\n        A.OneOf([\n            A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=1.0),\n            A.HueSaturationValue(hue_shift_limit=10, sat_shift_limit=20, val_shift_limit=20, p=1.0),\n            A.RandomGamma(gamma_limit=(80, 120), p=1.0),\n        ], p=0.5),\n        A.OneOf([\n            A.GaussNoise(var_limit=(10.0, 50.0), p=1.0),\n            A.GaussianBlur(blur_limit=(3, 5), p=1.0),\n            A.MotionBlur(blur_limit=5, p=1.0),\n        ], p=0.3),\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=1.0,\n        ),\n        ToTensorV2(),\n    ])\n\n    val_transform = A.Compose([\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=1.0,\n        ),\n        ToTensorV2(),\n    ])\n    \n    return train_transform, val_transform\n\nclass ForgeryDetectionDataset(Dataset):\n    def __init__(self, dataframe, transform=None, target_size=(512, 512)):\n        self.df = dataframe\n        self.transform = transform\n        self.target_size = target_size\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_path = row['image_path']\n        mask_path = row['mask_path']\n        label = row['label']\n        \n        # Load image\n        image = Image.open(image_path).convert('RGB')\n        image = image.resize(self.target_size, Image.BILINEAR)\n        image = np.array(image, dtype=np.float32) / 255.0\n        \n        # Load mask\n        if mask_path is not None and os.path.exists(mask_path):\n            mask = np.load(mask_path)\n            if len(mask.shape) == 3:\n                mask = np.max(mask, axis=0) # Combine regions\n            \n            mask = cv2.resize(\n                mask.astype(np.float32),\n                self.target_size,\n                interpolation=cv2.INTER_NEAREST\n            )\n            mask = (mask > 0.5).astype(np.float32)\n        else:\n            mask = np.zeros(self.target_size, dtype=np.float32)\n        \n        # Apply transforms\n        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image = augmented['image']\n            mask = augmented['mask']\n        \n        # Ensure mask has a channel dimension\n        if len(mask.shape) == 2:\n            mask = mask.unsqueeze(0)\n        \n        return {\n            'image': image,\n            'mask': mask,\n            'label': torch.tensor(label, dtype=torch.float32) # Also return label for potential classification loss\n        }\n\ndef get_dataloaders(train_df, val_df, cfg):\n    \"\"\"\n    Creates and returns train and validation dataloaders.\n    \"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\" CREATING DATALOADERS\")\n    print(\"=\"*60)\n    \n    train_transform, val_transform = get_transforms(cfg)\n    \n    train_dataset = ForgeryDetectionDataset(\n        dataframe=train_df,\n        transform=train_transform,\n        target_size=cfg.TARGET_SIZE\n    )\n    val_dataset = ForgeryDetectionDataset(\n        dataframe=val_df,\n        transform=val_transform,\n        target_size=cfg.TARGET_SIZE\n    )\n    \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 if cfg.DEVICE == 'cuda' else False,\n        drop_last=True\n    )\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=cfg.BATCH_SIZE,\n        shuffle=False,\n        num_workers=cfg.NUM_WORKERS,\n        pin_memory=True if cfg.DEVICE == 'cuda' else False\n    )\n    \n    print(f\"✓ Dataloaders created\")\n    print(f\"  - Batch size: {cfg.BATCH_SIZE}\")\n    print(f\"  - Training batches: {len(train_loader)}\")\n    print(f\"  - Validation batches: {len(val_loader)}\")\n    \n    return train_loader, val_loader\n\n# ===================================================================\n# MODEL, LOSS, METRICS\n# ===================================================================\n\nclass UNetForgeryDetector(nn.Module):\n    def __init__(self, encoder_name, encoder_weights, in_channels=3, classes=1):\n        super(UNetForgeryDetector, self).__init__()\n        self.model = smp.Unet(\n            encoder_name=encoder_name,\n            encoder_weights=encoder_weights, \n            in_channels=in_channels,\n            classes=classes,\n            activation=None # Will apply sigmoid in loss/metrics\n        )\n    \n    def forward(self, x):\n        return self.model(x)\n\nclass DiceLoss(nn.Module):  \n    def __init__(self, smooth=1.0):\n        super(DiceLoss, self).__init__()\n        self.smooth = smooth\n    \n    def forward(self, predictions, targets):\n        predictions = torch.sigmoid(predictions)\n        predictions = predictions.view(-1)\n        targets = targets.view(-1)\n        intersection = (predictions * targets).sum()\n        dice = (2. * intersection + self.smooth) / (predictions.sum() + targets.sum() + self.smooth)\n        return 1 - dice\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, bce_weight=0.5, dice_weight=0.5):\n        super(CombinedLoss, self).__init__()\n        self.bce_weight = bce_weight\n        self.dice_weight = dice_weight\n        self.bce = nn.BCEWithLogitsLoss()\n        self.dice = DiceLoss()\n    \n    def forward(self, predictions, targets):\n        bce_loss = self.bce(predictions, targets)\n        dice_loss = self.dice(predictions, targets)\n        combined = self.bce_weight * bce_loss + self.dice_weight * dice_loss\n        return combined, bce_loss, dice_loss\n\n# --- Metrics ---\ndef dice_coefficient(predictions, targets, threshold=0.5, smooth=1.0):\n    predictions = (torch.sigmoid(predictions) > threshold).float()\n    predictions = predictions.view(-1)\n    targets = targets.view(-1)\n    intersection = (predictions * targets).sum()\n    dice = (2. * intersection + smooth) / (predictions.sum() + targets.sum() + smooth)\n    return dice.item()\n\ndef iou_score(predictions, targets, threshold=0.5, smooth=1.0):\n    predictions = (torch.sigmoid(predictions) > threshold).float()\n    predictions = predictions.view(-1)\n    targets = targets.view(-1)\n    intersection = (predictions * targets).sum()\n    union = predictions.sum() + targets.sum() - intersection\n    iou = (intersection + smooth) / (union + smooth)\n    return iou.item()\n\ndef pixel_accuracy(predictions, targets, threshold=0.5):\n    predictions = (torch.sigmoid(predictions) > threshold).float()\n    correct = (predictions == targets).float().sum()\n    total = targets.numel()\n    return (correct / total).item()\n\n# ===================================================================\n# --- ADDED SECTION: TRAINING & VALIDATION FUNCTIONS ---\n# ===================================================================\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    \"\"\"\n    Performs one training epoch.\n    \"\"\"\n    model.train()\n    \n    total_loss = 0.0\n    total_dice = 0.0\n    total_iou = 0.0\n    \n    progress_bar = tqdm(loader, desc=\"Training\", leave=False)\n    for batch in progress_bar:\n        images = batch['image'].to(device)\n        masks = batch['mask'].to(device)\n        \n        # Zero gradients\n        optimizer.zero_grad()\n        \n        # Forward pass\n        outputs = model(images)\n        \n        # Calculate loss\n        loss, bce_loss, dice_loss = criterion(outputs, masks)\n        \n        # Backward pass and optimization\n        loss.backward()\n        optimizer.step()\n        \n        # Update metrics\n        batch_size = images.size(0)\n        total_loss += loss.item() * batch_size\n        total_dice += dice_coefficient(outputs.detach(), masks) * batch_size\n        total_iou += iou_score(outputs.detach(), masks) * batch_size\n        \n        progress_bar.set_postfix(\n            loss=f\"{loss.item():.4f}\",\n            bce=f\"{bce_loss.item():.4f}\",\n            dice=f\"{dice_loss.item():.4f}\"\n        )\n        \n    avg_loss = total_loss / len(loader.dataset)\n    avg_dice = total_dice / len(loader.dataset)\n    avg_iou = total_iou / len(loader.dataset)\n    \n    return avg_loss, avg_dice, avg_iou\n\n\ndef validate_one_epoch(model, loader, criterion, device):\n    \"\"\"\n    Performs one validation epoch.\n    \"\"\"\n    model.eval()\n    \n    total_loss = 0.0\n    total_dice = 0.0\n    total_iou = 0.0\n    \n    with torch.no_grad():\n        progress_bar = tqdm(loader, desc=\"Validation\", leave=False)\n        for batch in progress_bar:\n            images = batch['image'].to(device)\n            masks = batch['mask'].to(device)\n            \n            # Forward pass\n            outputs = model(images)\n            \n            # Calculate loss\n            loss, bce_loss, dice_loss = criterion(outputs, masks)\n            \n            # Update metrics\n            batch_size = images.size(0)\n            total_loss += loss.item() * batch_size\n            total_dice += dice_coefficient(outputs, masks) * batch_size\n            total_iou += iou_score(outputs, masks) * batch_size\n            \n            progress_bar.set_postfix(\n                loss=f\"{loss.item():.4f}\"\n            )\n\n    avg_loss = total_loss / len(loader.dataset)\n    avg_dice = total_dice / len(loader.dataset)\n    avg_iou = total_iou / len(loader.dataset)\n    \n    return avg_loss, avg_dice, avg_iou\n\n# ===================================================================\n# --- ADDED SECTION: MAIN EXECUTION & TRAINING LOOP ---\n# ===================================================================\n\ndef main():\n    # 1. Setup\n    cfg = Config()\n    set_seed(cfg.SEED)\n    warnings.filterwarnings(\"ignore\")\n    plt.style.use('seaborn-v0_8-darkgrid')\n    sns.set_palette(\"husl\")\n    print(f\"Using device: {cfg.DEVICE}\")\n\n    # 2. EDA\n    authentic_images, forged_images = run_eda(cfg)\n    \n    # 3. Data Preparation\n    train_df, val_df = get_dataframes(authentic_images, forged_images, cfg)\n    \n    # 4. Dataloaders\n    train_loader, val_loader = get_dataloaders(train_df, val_df, cfg)\n    \n    # 5. Model, Loss, Optimizer\n    print(\"\\n\" + \"=\"*60)\n    print(\" INITIALIZING MODEL & COMPONENTS\")\n    print(\"=\"*60)\n    \n    model = UNetForgeryDetector(\n        encoder_name=cfg.MODEL_NAME,\n        encoder_weights=cfg.ENCODER_WEIGHTS,\n    ).to(cfg.DEVICE)\n    \n    criterion = CombinedLoss(bce_weight=0.5, dice_weight=0.5)\n    \n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=cfg.LEARNING_RATE,\n        weight_decay=cfg.WEIGHT_DECAY\n    )\n    \n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        mode='min',\n        factor=0.5,\n        patience=3, # Reduced patience a bit\n        verbose=True,\n        min_lr=1e-7\n    )\n    \n    print(f\"✓ Model: {cfg.MODEL_NAME} (Weights: {cfg.ENCODER_WEIGHTS})\")\n    print(\"✓ Loss: Combined (0.5*BCE + 0.5*Dice)\")\n    print(\"✓ Optimizer: AdamW\")\n    print(\"✓ Scheduler: ReduceLROnPlateau\")\n\n    # 6. Training Loop\n    print(\"\\n\" + \"=\"*60)\n    print(\" STARTING MODEL TRAINING\")\n    print(\"=\"*60)\n    \n    best_val_loss = float('inf')\n    best_model_state = None\n    history = defaultdict(list)\n    \n    for epoch in range(1, cfg.NUM_EPOCHS + 1):\n        start_time = time.time()\n        \n        # Train\n        train_loss, train_dice, train_iou = train_one_epoch(\n            model, train_loader, optimizer, criterion, cfg.DEVICE\n        )\n        \n        # Validate\n        val_loss, val_dice, val_iou = validate_one_epoch(\n            model, val_loader, criterion, cfg.DEVICE\n        )\n        \n        # Update scheduler\n        scheduler.step(val_loss)\n        \n        # Log metrics\n        history['train_loss'].append(train_loss)\n        history['train_dice'].append(train_dice)\n        history['val_loss'].append(val_loss)\n        history['val_dice'].append(val_dice)\n        \n        end_time = time.time()\n        epoch_mins = (end_time - start_time) / 60\n        \n        print(f\"\\nEpoch {epoch:02}/{cfg.NUM_EPOCHS} | Time: {epoch_mins:.2f}m\")\n        print(f\"\\tTrain Loss: {train_loss:.4f} | Train Dice: {train_dice:.4f} | Train IoU: {train_iou:.4f}\")\n        print(f\"\\t Val. Loss: {val_loss:.4f} |  Val. Dice: {val_dice:.4f} |  Val. IoU: {val_iou:.4f}\")\n        \n        # Save best model\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            best_model_state = copy.deepcopy(model.state_dict())\n            torch.save(best_model_state, 'best_model.pth')\n            print(f\"\\t>> New best model saved! (Val Loss: {best_val_loss:.4f})\")\n        \n        print(\"-\"*60)\n        \n    print(\"Training finished!\")\n    \n    # 7. Plotting Results\n    print(\"Plotting training history...\")\n    plt.figure(figsize=(12, 5))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(history['train_loss'], label='Train Loss')\n    plt.plot(history['val_loss'], label='Val Loss')\n    plt.title('Loss per Epoch')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(history['train_dice'], label='Train Dice')\n    plt.plot(history['val_dice'], label='Val Dice')\n    plt.title('Dice Score per Epoch')\n    plt.xlabel('Epoch')\n    plt.ylabel('Dice Score')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.savefig('training_history.png')\n    plt.show()\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}