{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":14456136,"sourceType":"competition"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Scientific Forgery Classification and Segmentation","metadata":{}},{"cell_type":"code","source":"!pip install -U git+https://github.com/qubvel/segmentation_models.pytorch > /dev/null\n!pip install numba Pillow > /dev/null","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:16:04.864336Z","iopub.execute_input":"2025-11-26T19:16:04.864954Z","iopub.status.idle":"2025-11-26T19:17:30.015089Z","shell.execute_reply.started":"2025-11-26T19:16:04.864922Z","shell.execute_reply":"2025-11-26T19:17:30.014004Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import csv\nimport functools\nimport json\nimport os\nimport pathlib\nimport random\nimport time\nimport traceback\nimport warnings\n\nimport albumentations as A\nimport matplotlib.pyplot as plt\nimport numba\nimport numpy as np\nimport numpy.typing as npt\nimport pandas as pd\nimport PIL.Image as Image\nimport segmentation_models_pytorch as smp\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nimport torch.utils.data as data\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:20:36.888018Z","iopub.execute_input":"2025-11-26T19:20:36.888738Z","iopub.status.idle":"2025-11-26T19:20:42.595181Z","shell.execute_reply.started":"2025-11-26T19:20:36.88871Z","shell.execute_reply":"2025-11-26T19:20:42.594568Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class Config:\n    RANDOM_SEED = 42\n\n    IS_KAGGLE = os.path.exists(\"/kaggle/input/\")\n\n    # HARDWARE_RELATED\n    DEVICE = (\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"mps\" if torch.backends.mps.is_available() else \"cpu\"\n    )\n\n    # User torch.float16 for MacBooks with MPS chips\n    # DTYPE = torch.float32 if (IS_KAGGLE or (DEVICE == \"cuda\")) else torch.float16\n\n    # DATASET_RELATED\n    TRAIN_TRANSFORM = A.Compose(\n        [\n            A.Resize(512, 512, p=1.0),\n            A.HorizontalFlip(),\n            A.VerticalFlip(),\n            A.RandomBrightnessContrast(),\n            A.ChannelShuffle(),\n            A.GaussNoise(),\n            A.Normalize(),\n            A.ToTensorV2(),\n        ]\n    )\n    INFERENCE_TRANSFORM = A.Compose([A.Resize(512, 512), A.Normalize(), A.ToTensorV2()])\n    TRAIN_VAL_SPLIT_RATIO = [0.85, 0.15]\n\n    BASE_DIR = (\n        pathlib.Path(\".\")\n        if not IS_KAGGLE\n        else pathlib.Path(\n            \"/kaggle/input/recodai-luc-scientific-image-forgery-detection\"\n        )\n    )\n    TRAIN_IMAGES_DIR = BASE_DIR / \"train_images\"\n    AUTHENTIC_IMAGES_DIR = TRAIN_IMAGES_DIR / \"authentic\"\n    FORGED_IMAGES_DIR = TRAIN_IMAGES_DIR / \"forged\"\n    MASKS_IMAGES_DIR = BASE_DIR / \"train_masks\"\n    SUPPLEMENTAL_IMAGES_DIR = BASE_DIR / \"supplemental_images\"\n    SUPPLEMENTAL_MASKS_DIR = BASE_DIR / \"supplemental_masks\"\n    INFERENCE_IMAGES_DIR = BASE_DIR / \"test_images\"\n\n    # MODEL_RELATED\n    ENCODER_NAME = \"efficientnet-b4\"\n    ENCODER_WEIGHTS = (\n        None  # No weights needed because for the competition, internet is shut off\n    )\n    IN_CHANNELS = 3\n    SEGMENTATION_CLASSES = 1\n    CLASSIFICATION_CLASSES = 1\n    SEGMENTATION_ACTIVATION = None\n    CLASSIFICATION_ACTIVATION = None\n    SEGMENTATION_INFERENCE_ACTIVATION = torch.sigmoid\n    CLASSIFICATION_INFERENCE_ACTIVATION = torch.sigmoid\n\n    # TRAINING & INFERENCE RELATED\n    LEARNING_RATE = 0.0005\n    EPOCHS = 10 * 10 if IS_KAGGLE else 2\n    CHECK_VAL_EVERY_N_EPOCH = 1\n    TRAIN_BATCH_SIZE = 8\n    VAL_BATCH_SIZE = 8\n    INFERENCE_BATCH_SIZE = 8\n    SHOULD_SHUFFLE_TRAIN_DATALOADER = True\n    SHOULD_SHUFFLE_VAL_DATALOADER = False\n    SHOULD_SHUFFLE_INFERENCE_DATALOADER = False\n    INFERENCE_CLASSFICATION_THRESHOLD = 0.5\n    INFERENCE_SEGMENTATION_THRESHOLD = 0.5\n\n    # ARTIFACT RELATED\n    MODEL_SAVE_DIR = BASE_DIR if not IS_KAGGLE else \"/kaggle/working\"\n\n    # SUBMISSIONS_RELATED\n    SUBMISSIONS_DIR = BASE_DIR if not IS_KAGGLE else \"/kaggle/working\"\n\n\ndef clear_acclerator_cache(Config):\n    device = Config.DEVICE\n\n    funcs = {\n        \"mps\": torch.mps.empty_cache,\n        \"cuda\": torch.cuda.empty_cache,\n    }\n\n    funcs[device]()\n\n\nclear_acclerator_cache(Config=Config)\n\nprint(f\"IS_KAGGLE: {Config.IS_KAGGLE}\")\nprint(f\"DEVICE: {Config.DEVICE}\")\nprint(f\"INFERENCE_IMAGES_DIR: {Config.INFERENCE_IMAGES_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:20:49.046031Z","iopub.execute_input":"2025-11-26T19:20:49.046822Z","iopub.status.idle":"2025-11-26T19:20:49.090305Z","shell.execute_reply.started":"2025-11-26T19:20:49.046796Z","shell.execute_reply":"2025-11-26T19:20:49.089505Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utils Code","metadata":{}},{"cell_type":"code","source":"def save_model(model, optimizer, location):\n\n    torch.save(\n        {\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n        },\n        location,\n    )\n\n\ndef load_model(location, model, optimizer):\n    with open(location, \"rb\") as checkpoint:\n        state_dicts = torch.load(checkpoint)\n\n        model.load_state_dict(state_dicts[\"model_state_dict\"])\n        optimizer.load_state_dict(state_dicts[\"optimizer_state_dict\"])\n\n\ndef print_model_summary(model):\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    non_trainable_params = total_params - trainable_params\n    model_size = sum(p.numel() for p in model.parameters()) * 4 / 1024 / 1024\n\n    print(\"=\" * 80)\n    print(f\"Total Parameters: {total_params:,}\")\n    print(f\"Trainable Parameters: {trainable_params:,}\")\n    print(f\"Non-Trainable Parameters: {non_trainable_params:,}\")\n    print(f\"Model Size (in MB): {model_size:,.2f}\")\n    print(\"=\" * 80)\n\n\ndef write_submission_csv(directory, submissions, fieldnames):\n    filename = os.path.join(directory, \"submission.csv\")\n\n    with open(filename, \"w\", newline=\"\") as submission_file:\n        writer = csv.DictWriter(submission_file, fieldnames)\n\n        writer.writeheader()\n\n        writer.writerows(submissions)\n\n        # log the creation of submission file\n        print(f\"Submission file created at: {filename}\")\n        print(f\"Number of dicts written: {len(submissions)}\")\n        print(f\"Wrote {len(submissions) + 1} rows to submission file.\")\n\n\n@numba.jit(nopython=True)\ndef _rle_encode_jit(x: npt.NDArray, fg_val: int = 1) -> list[int]:\n    \"\"\"Numba-jitted RLE encoder.\"\"\"\n\n    dots = np.where(x.T.flatten() == fg_val)[0]\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return run_lengths\n\n\ndef seed_everything(seed: int):\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    torch.backends.cudnn.benchmark = True\n\n\nseed_everything(Config.RANDOM_SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:20:54.442659Z","iopub.execute_input":"2025-11-26T19:20:54.442967Z","iopub.status.idle":"2025-11-26T19:20:54.465511Z","shell.execute_reply.started":"2025-11-26T19:20:54.442944Z","shell.execute_reply":"2025-11-26T19:20:54.464853Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training and Plotting Codes","metadata":{}},{"cell_type":"code","source":"# TRAIN FUNCTION TO TRAIN THE MODEL\ndef show_training_duration(n_epochs):\n\n    def outer_wrapper(fn):\n        @functools.wraps(fn)\n        def wrapper(*args, **kwargs):\n            start_time = time.time()\n\n            result = fn(*args, **kwargs)\n\n            end_time = time.time()\n\n            # Calculate duration\n            total_seconds = end_time - start_time\n            hours = int(total_seconds // 3600)\n            minutes = int((total_seconds % 3600) // 60)\n            seconds = int(total_seconds % 60)\n\n            print(f\"Total Training Time: {hours}h {minutes}m {seconds}s\")\n            print(f\"Training Time Per Epoch: {total_seconds / n_epochs:,.3f} s\")\n\n            return result\n\n        return wrapper\n\n    return outer_wrapper\n\n\ndef _print_batch_message(\n    c_epoch, t_epoch, c_batch, t_batch, tot_loss, class_loss, seg_loss\n):\n    max_batch_padding = len(str(t_batch))\n    max_epoch_padding = len(str(t_epoch))\n\n    message = f\"Epoch: [{c_epoch:>{max_epoch_padding}}/{t_epoch}] Batch: [{c_batch:>{max_batch_padding}}/{t_batch}] Tot. Loss: {tot_loss:.4f} | Class. Loss: {class_loss:.4f} | Seg. Loss: {seg_loss:.4f}\"\n\n    print(message, end=\"\\r\")\n\n\ndef _print_epoch_message(c_epoch, t_epoch, t_batch, tot_loss, class_loss, seg_loss):\n    max_epoch_padding = len(str(t_epoch))\n\n    message = f\"Epoch [{c_epoch:>{max_epoch_padding}}/{t_epoch}] Batch: [{t_batch}/{t_batch}] Tot. Loss: {tot_loss:.4f} | Class. Loss: {class_loss:.4f} | Seg. Loss: {seg_loss:.4f}\"\n\n    print(message, end=\"\\n\")\n\n\n@show_training_duration(n_epochs=Config.EPOCHS)\ndef train(\n    model,\n    train_loader,\n    criterion,\n    optimizer,\n    val_loader=None,\n    epochs=1,\n    scheduler=None,\n    val_criterion=None,\n    device=\"cpu\",\n    model_save_location=\"/data\",\n):\n    if val_criterion is None:\n        val_criterion = criterion\n\n    n_batches = len(train_loader)\n    training_history = {\n        \"train_total_loss\": [],\n        \"val_total_loss\": [],\n        \"train_classification_loss\": [],\n        \"val_classification_loss\": [],\n        \"train_segmentation_loss\": [],\n        \"val_segmentation_loss\": [],\n    }\n\n    try:\n        model = model.to(device)\n        print(f\"++++++++++ MODEL MOVED TO DEVICE: {device} ++++++++++\")\n\n        print(f\"++++++++++ TRAINING START ++++++++++\")\n\n        for epoch in range(1, epochs + 1):\n\n            # PHASE: TRAINING\n            model.train()\n\n            batch_train_loss = []\n            batch_classfication_loss = []\n            batch_segmentation_loss = []\n            for batch_idx, batch in enumerate(train_loader, 1):\n                images, masks, labels = batch\n\n                # Move to Device\n                images = images.to(device)\n                masks = masks.to(device)\n                labels = labels.to(device)\n\n                # Input Masks Shape: (B, H, W, C)\n                # Model Output Masks Shape: (B, C, H, W)\n                # Solution: Permute Input Masks to match model output shape\n                masks = masks.permute(0, 3, 1, 2)\n\n                # Forward Pass\n                classification_labels, segmentation_labels = model(images)\n\n                # Loss Calculation\n                total_loss, classification_loss, segmentation_loss = criterion(\n                    classification_labels, labels, segmentation_labels, masks\n                )\n\n                # Backward Pass\n                optimizer.zero_grad()\n                total_loss.backward()\n                optimizer.step()\n\n                # Add metrics\n                loss_sum = classification_loss + segmentation_loss\n                batch_train_loss.append(loss_sum)\n                batch_classfication_loss.append(classification_loss)\n                batch_segmentation_loss.append(segmentation_loss)\n\n                _print_batch_message(\n                    epoch,\n                    epochs,\n                    batch_idx,\n                    n_batches,\n                    batch_train_loss[-1],\n                    batch_classfication_loss[-1],\n                    batch_segmentation_loss[-1],\n                )\n\n                # Delete data to free acclerator memory\n                del images, masks, labels\n                del classification_labels, segmentation_labels\n                del total_loss\n\n            training_history[\"train_total_loss\"].append(\n                sum(batch_train_loss) / n_batches\n            )\n            training_history[\"train_classification_loss\"].append(\n                sum(batch_classfication_loss) / n_batches\n            )\n            training_history[\"train_segmentation_loss\"].append(\n                sum(batch_segmentation_loss) / n_batches\n            )\n\n            # PHASE: VALIDATION\n            if val_loader is not None:\n                model.eval()\n\n                with torch.no_grad():\n                    n_val_batches = len(val_loader)\n\n                    batch_val_loss = []\n                    batch_classfication_loss = []\n                    batch_segmentation_loss = []\n                    for batch_idx, batch in enumerate(val_loader, 1):\n                        images, masks, labels = batch\n\n                        # Move to Device\n                        images = images.to(device)\n                        masks = masks.to(device)\n                        labels = labels.to(device)\n\n                        # Input Masks Shape: (B, H, W, C)\n                        # Model Output Masks Shape: (B, C, H, W)\n                        # Solution: Permute Input Masks to match model output shape\n                        masks = masks.permute(0, 3, 1, 2)\n\n                        # Forward Pass\n                        classification_labels, segmentation_labels = model(images)\n\n                        # Loss Calculation\n                        total_loss, classification_loss, segmentation_loss = (\n                            val_criterion(\n                                classification_labels,\n                                labels,\n                                segmentation_labels,\n                                masks,\n                            )\n                        )\n\n                        # Add Metrics\n                        batch_val_loss.append(classification_loss + segmentation_loss)\n                        batch_classfication_loss.append(classification_loss)\n                        batch_segmentation_loss.append(segmentation_loss)\n\n                        # Delete data to free acclerator memory\n                        del images, masks, labels\n                        del classification_labels, segmentation_labels\n                        del total_loss\n\n                training_history[\"val_total_loss\"].append(\n                    sum(batch_val_loss) / n_val_batches\n                )\n                training_history[\"val_classification_loss\"].append(\n                    sum(batch_classfication_loss) / n_val_batches\n                )\n                training_history[\"val_segmentation_loss\"].append(\n                    sum(batch_segmentation_loss) / n_val_batches\n                )\n\n            _print_epoch_message(\n                epoch,\n                epochs,\n                n_batches,\n                training_history[\n                    \"val_total_loss\" if val_loader is not None else \"train_total_loss\"\n                ][-1],\n                training_history[\n                    (\n                        \"val_classification_loss\"\n                        if val_loader is not None\n                        else \"train_classification_loss\"\n                    )\n                ][-1],\n                training_history[\n                    (\n                        \"val_segmentation_loss\"\n                        if val_loader is not None\n                        else \"train_segmentation_loss\"\n                    )\n                ][-1],\n            )\n\n            if scheduler is not None:\n                scheduler.step()\n\n        print(f\"++++++++++ TRAINING ENDED ++++++++++\")\n        print(f\"++++++++++ SAVING MODEL ++++++++++\")\n        save_model(\n            model, optimizer, os.path.join(model_save_location, \"model_ckpt_final.pth\")\n        )\n        print(f\"++++++++++ SAVED MODEL ++++++++++\")\n    except Exception as e:\n        print(f\"Training failed with error: {e}\")\n        print(traceback.format_exc())\n    except KeyboardInterrupt as ke:\n        filename = f\"model_ckpt_interrupted_epoch_{epoch}.pth\"\n        model_save_location = os.path.join(model_save_location, filename)\n\n        print(f\"++++++++++ TRAINING INTERUPTED ++++++++++\")\n        print(f\"++++++++++ SAVING MODEL TO {model_save_location} ++++++++++\")\n\n        save_model(model, optimizer, model_save_location)\n\n    return training_history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:20:57.337347Z","iopub.execute_input":"2025-11-26T19:20:57.33806Z","iopub.status.idle":"2025-11-26T19:20:57.357173Z","shell.execute_reply.started":"2025-11-26T19:20:57.338039Z","shell.execute_reply":"2025-11-26T19:20:57.35641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CODE TO PLOT MODEL TRAINING HISTORY\ndef plot_tranining_history(training_history):\n\n    n_epochs = range(1, len(training_history[\"train_total_loss\"]) + 1)\n\n    fig, axes = plt.subplots(2, 3, figsize=(18, 15))\n\n    fig.suptitle(\"HYDRANET TRANING HISTORY\", fontsize=16, fontweight=\"bold\")\n\n    # Plot total loss\n    axes[0, 0].plot(\n        n_epochs,\n        training_history[\"train_total_loss\"],\n        \"b-\",\n        label=\"Training Total Loss\",\n        linewidth=2,\n    )\n    axes[0, 0].plot(\n        n_epochs,\n        training_history[\"val_total_loss\"],\n        \"o-\",\n        label=\"Validation Total Loss\",\n        linewidth=2,\n    )\n    axes[0, 0].set_title(\"Total Loss\", fontweight=\"bold\")\n    axes[0, 0].set_xlabel(\"Epochs\")\n    axes[0, 0].set_ylabel(\"Total Loss\")\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.5)\n\n    # Plot Classification Loss\n    axes[0, 1].plot(\n        n_epochs,\n        training_history[\"train_classification_loss\"],\n        \"y-\",\n        label=\"Training Classification Loss\",\n        linewidth=2,\n    )\n    axes[0, 1].plot(\n        n_epochs,\n        training_history[\"val_classification_loss\"],\n        \"o-\",\n        label=\"Validation Classification Loss\",\n        linewidth=2,\n    )\n    axes[0, 1].set_title(\"Classification Loss\", fontweight=\"bold\")\n    axes[0, 1].set_xlabel(\"Epochs\")\n    axes[0, 1].set_ylabel(\"Classification Loss\")\n    axes[0, 1].legend()\n    axes[0, 1].grid(True, alpha=0.5)\n\n    # Plot Segmentation Loss\n    axes[0, 2].plot(\n        n_epochs,\n        training_history[\"train_segmentation_loss\"],\n        \"y-\",\n        label=\"Training Segmentation Loss\",\n        linewidth=2,\n    )\n    axes[0, 2].plot(\n        n_epochs,\n        training_history[\"val_segmentation_loss\"],\n        \"o-\",\n        label=\"Validation Segmentation Loss\",\n        linewidth=2,\n    )\n    axes[0, 2].set_title(\"Segmentation Loss\", fontweight=\"bold\")\n    axes[0, 2].set_xlabel(\"Epochs\")\n    axes[0, 2].set_ylabel(\"Segmentation Loss\")\n    axes[0, 2].legend()\n    axes[0, 2].grid(True, alpha=0.5)\n\n    plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n\n    plt.show()\n\n\ndef plot_mask_from_tensor_batch(masks):\n    n_masks, n_channels, height, width = masks.shape\n\n    masks = masks.cpu().numpy()\n\n    if n_channels > 1:\n        print(\"More than 1 channel for a mask. Can't visualize mask\")\n\n    if n_masks == 1:\n        mask = np.squeeze(masks, axis=0)\n        mask = np.transpose(mask, (1, 2, 0))[:, :, 0]\n\n        plt.imshow(mask, cmap=\"grey\")\n\n        return\n\n    n_rows = (n_masks // 3) + 1\n    n_cols = 3\n\n    fig, axs = plt.subplots(n_rows, n_cols)\n\n    axs = axs.flatten()\n    n_axs = len(axs)\n\n    for idx in range(n_masks):\n        mask = masks[idx, 0, :, :]\n\n        axs[idx].imshow(mask, cmap=\"grey\")\n        axs[idx].set_axis_off()\n        axs[idx].set_title(f\"Sample: {idx + 1}\")\n\n    for ax_idx in range(idx, n_axs):\n        axs[ax_idx].set_axis_off()\n\n    fig.suptitle(\"Predicted Masks for Test Dataset\")\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:21:02.964844Z","iopub.execute_input":"2025-11-26T19:21:02.96557Z","iopub.status.idle":"2025-11-26T19:21:02.976298Z","shell.execute_reply.started":"2025-11-26T19:21:02.965546Z","shell.execute_reply":"2025-11-26T19:21:02.975548Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## IO","metadata":{}},{"cell_type":"code","source":"def list_all_files_in_dir(dir_path):\n    return [os.path.join(dir_path, f) for f in sorted(os.listdir(dir_path))]\n\n\ndef get_file_stats(file_path):\n    directory, filename = os.path.split(file_path)\n    filename, extension = os.path.splitext(filename)\n\n    return directory, filename, extension\n\n\ndef read_png_image(file_path):\n    with Image.open(file_path) as img:\n        numpy_array = np.array(img)\n\n        shape = numpy_array.shape\n        n_dims = len(shape)\n\n        if n_dims == 2:\n            numpy_array = np.stack([numpy_array] * 3, axis=-1)\n        elif n_dims == 3:\n            n_channels = numpy_array.shape[2]\n\n            if n_channels == 1:\n                numpy_array = np.repeat(\n                    [numpy_array[:, :, 0], numpy_array[:, :, 0], numpy_array[:, :, 0]]\n                )\n            elif n_channels == 2:\n                zeros = np.zeros_like(numpy_array[:, :, 0])\n                numpy_array = np.stack(\n                    [numpy_array[:, :, 0], numpy_array[:, :, 1], zeros], axis=-1\n                )\n            elif n_channels >= 3:\n                numpy_array = numpy_array[:, :, :3]\n\n        shape = numpy_array.shape\n        n_dims = len(shape)\n\n        return numpy_array\n\n\ndef read_npy_image(file_path):\n    data = np.load(file_path)\n\n    shape = data.shape\n    n_dims = len(shape)\n\n    if n_dims == 2:\n        # There is no channel image\n        data = data[np.newaxis, :, :]\n    elif n_dims == 3:\n        n_channels, _, _ = shape\n\n        if n_channels > 1:\n            data = np.max(data, axis=0)\n            data = data[np.newaxis, :, :]\n\n    data = data.transpose(1, 2, 0)\n\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:21:06.712793Z","iopub.execute_input":"2025-11-26T19:21:06.713429Z","iopub.status.idle":"2025-11-26T19:21:06.721663Z","shell.execute_reply.started":"2025-11-26T19:21:06.713406Z","shell.execute_reply":"2025-11-26T19:21:06.720664Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class ForgeryImagesTrainDataset(data.Dataset):\n    def __init__(self, Config):\n        self.authentic_images_dir = Config.AUTHENTIC_IMAGES_DIR\n        self.forged_images_dir = Config.FORGED_IMAGES_DIR\n        self.masks_images_dir = Config.MASKS_IMAGES_DIR\n        self.supplemental_images_dir = Config.SUPPLEMENTAL_IMAGES_DIR\n        self.supplemental_masks_dir = Config.SUPPLEMENTAL_MASKS_DIR\n        self.transform = Config.TRAIN_TRANSFORM\n\n        self.samples = []\n        self.should_use_float32 = Config.DEVICE == \"mps\"\n\n        # Process Authentic Images\n        authentic_images = list_all_files_in_dir(self.authentic_images_dir)\n        for image_path in authentic_images:\n            self.samples.append(\n                {\n                    \"image_path\": image_path,\n                    \"mask_path\": None,\n                    \"label\": 0,\n                }\n            )\n\n        # Process Forged Images\n        forged_images = list_all_files_in_dir(self.forged_images_dir)\n        for image_path in forged_images:\n            parent_dir, filename, extension = get_file_stats(image_path)\n            mask_path = os.path.join(self.masks_images_dir, filename + \".npy\")\n            self.samples.append(\n                {\n                    \"image_path\": image_path,\n                    \"mask_path\": mask_path,\n                    \"label\": 1,\n                }\n            )\n\n        # Process Supplemental Images\n        supplemental_images = list_all_files_in_dir(self.supplemental_images_dir)\n        for image_path in supplemental_images:\n            parent_dir, filename, extension = get_file_stats(image_path)\n            mask_path = os.path.join(self.supplemental_masks_dir, filename + \".npy\")\n            self.samples.append(\n                {\n                    \"image_path\": image_path,\n                    \"mask_path\": mask_path,\n                    \"label\": 1,\n                }\n            )\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n\n        image = read_png_image(sample[\"image_path\"])\n\n        image_h, image_w, _ = image.shape\n\n        if sample[\"mask_path\"] is None:\n            mask = np.zeros((image_h, image_w, 1))\n        else:\n            mask = read_npy_image(sample[\"mask_path\"])\n\n        augmented_sample = self.transform(image=image, mask=mask)\n\n        image = augmented_sample[\"image\"]\n        mask = augmented_sample[\"mask\"]\n        label = torch.tensor([sample[\"label\"]])\n\n        image=image.to(dtype=torch.float32)\n        mask = mask.to(dtype=torch.float32)\n        label = label.to(dtype=torch.float32)\n\n        return image, mask, label\n\n\nclass ForgeryImagesInferenceDataset(data.Dataset):\n    def __init__(self, Config):\n        self.images_dir = Config.INFERENCE_IMAGES_DIR\n        self.transform = Config.INFERENCE_TRANSFORM\n        self.images = list_all_files_in_dir(self.images_dir)\n\n        self.should_use_float32 = Config.DEVICE == \"mps\"\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image_path = self.images[idx]\n        image = read_png_image(image_path)\n\n        augment = self.transform(image=image)\n        image = augment[\"image\"]\n\n        image = image.to(dtype=torch.float32)\n\n        return image, image_path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:27:10.598082Z","iopub.execute_input":"2025-11-26T19:27:10.598438Z","iopub.status.idle":"2025-11-26T19:27:10.610526Z","shell.execute_reply.started":"2025-11-26T19:27:10.598412Z","shell.execute_reply":"2025-11-26T19:27:10.609609Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = ForgeryImagesTrainDataset(\n    Config=Config,\n)\n\ntrain_loader = data.DataLoader(\n    train_ds,\n    batch_size=Config.TRAIN_BATCH_SIZE,\n    shuffle=Config.SHOULD_SHUFFLE_TRAIN_DATALOADER,\n)\n\nimages, masks, labels = next(iter(train_loader))\n\nprint(f\"Batch Images Shape: {images.shape}\")\nprint(f\"Batch Masks Shape: {masks.shape}\")\nprint(f\"Batch Labels Shape: {labels.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:27:14.203114Z","iopub.execute_input":"2025-11-26T19:27:14.203421Z","iopub.status.idle":"2025-11-26T19:27:14.995049Z","shell.execute_reply.started":"2025-11-26T19:27:14.203403Z","shell.execute_reply":"2025-11-26T19:27:14.994163Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss Function","metadata":{}},{"cell_type":"code","source":"class DualLoss(nn.Module):\n    def __init__(\n        self,\n        classification_criterion,\n        segmentation_criterion,\n        classification_weight=0.5,\n        segmentation_weight=0.5,\n    ):\n        super(DualLoss, self).__init__()\n\n        self.classification_criterion = classification_criterion\n\n        self.segmentation_criterion = segmentation_criterion\n\n        self.classification_weight = classification_weight\n        self.segmentation_weight = segmentation_weight\n\n    def forward(\n        self,\n        classification_pred,\n        classification_target,\n        segmentation_pred,\n        segmentation_target,\n    ):\n        classification_loss = self.classification_criterion(\n            classification_pred, classification_target\n        )\n\n        segmentation_loss = self.segmentation_criterion(\n            segmentation_pred, segmentation_target\n        )\n\n        total_loss = (\n            self.classification_weight * classification_loss\n            + self.segmentation_weight * segmentation_loss\n        )\n\n        classification_loss = classification_loss.item()\n        segmentation_loss = segmentation_loss.item()\n\n        return total_loss, classification_loss, segmentation_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:27:17.877158Z","iopub.execute_input":"2025-11-26T19:27:17.877796Z","iopub.status.idle":"2025-11-26T19:27:17.883432Z","shell.execute_reply.started":"2025-11-26T19:27:17.877774Z","shell.execute_reply":"2025-11-26T19:27:17.882634Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class ForgeryDetectionModel(nn.Module):\n    def __init__(self, Config):\n        super(ForgeryDetectionModel, self).__init__()\n\n        # Use Hydranet Architecture to build the model\n        # Simultaneously perform classification and segmentation so that we can use the same encoder\n        # for both tasks.\n        # Doing this gives us which segmentation masks to run RLE encoding on.\n        # This saves compute and also makes the model more compact.\n        self.model = smp.Unet(\n            encoder_name=Config.ENCODER_NAME,\n            encoder_weights=Config.ENCODER_WEIGHTS,\n            activation=None,\n            in_channels=Config.IN_CHANNELS,\n            classes=Config.SEGMENTATION_CLASSES,\n            aux_params={\n                \"classes\": Config.CLASSIFICATION_CLASSES,\n                \"activation\": None,\n            },\n        )\n\n    def forward(self, x):\n        segmentation_result, classification_result = self.model(x)\n\n        classification_result = classification_result.to(dtype=torch.float)\n\n        return classification_result, segmentation_result\n\n\nmodel = ForgeryDetectionModel(Config=Config)\n\nmodel.train()\n\nwith torch.no_grad():\n    sample_tensor = torch.rand((Config.INFERENCE_BATCH_SIZE, 3, 512, 512))\n\n    classification_result, segmentation_result = model(sample_tensor)\n\n    print(f\"Inference Input Image Shape: {sample_tensor.shape}\")\n    print(f\"Inference Classification Shape: {classification_result.shape}\")\n    print(f\"Inference Segmentation Shape: {segmentation_result.shape}\")\n\ndel model\ndel sample_tensor\ndel classification_result, segmentation_result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:27:21.216076Z","iopub.execute_input":"2025-11-26T19:27:21.216733Z","iopub.status.idle":"2025-11-26T19:27:26.972261Z","shell.execute_reply.started":"2025-11-26T19:27:21.21671Z","shell.execute_reply":"2025-11-26T19:27:26.971505Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training and Plotting","metadata":{}},{"cell_type":"code","source":"# CODE TO TRAIN THE MODEL\ntotal_ds = ForgeryImagesTrainDataset(Config=Config)\ntrain_ds, val_ds = data.random_split(total_ds, lengths=Config.TRAIN_VAL_SPLIT_RATIO)\n\ntrain_loader = data.DataLoader(\n    train_ds,\n    batch_size=Config.TRAIN_BATCH_SIZE,\n    shuffle=Config.SHOULD_SHUFFLE_TRAIN_DATALOADER,\n)\nval_loader = data.DataLoader(\n    val_ds,\n    batch_size=Config.VAL_BATCH_SIZE,\n    shuffle=Config.SHOULD_SHUFFLE_VAL_DATALOADER,\n)\n\nmodel = ForgeryDetectionModel(Config=Config)\n\noptimizer = optim.Adam(model.parameters(), lr=Config.LEARNING_RATE)\n# scheduler = lr_scheduler.MultiplicativeLR(optimizer, lr_lambda=lambda epoch: 0.9 if epoch % 5 == 0 else 1.0)\n\ncriterion = DualLoss(\n    nn.BCEWithLogitsLoss(), smp.losses.DiceLoss(mode=\"binary\", from_logits=True),\n    classification_weight=1.0,segmentation_weight=1.5,\n)\n\ntraining_history = train(\n    model=model,\n    train_loader=train_loader,\n    criterion=criterion,\n    optimizer=optimizer,\n    val_loader=val_loader,\n    epochs=Config.EPOCHS,\n    # scheduler=scheduler,\n    device=Config.DEVICE,\n    model_save_location=Config.MODEL_SAVE_DIR,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T19:27:26.973228Z","iopub.execute_input":"2025-11-26T19:27:26.973642Z","execution_failed":"2025-11-26T21:13:28.657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_tranining_history(training_history=training_history)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-26T21:13:28.657Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"clear_acclerator_cache(Config=Config)\n\nmodel = model.to(Config.DEVICE)\nmodel.eval()\n\ntest_ds = ForgeryImagesInferenceDataset(Config=Config)\ntest_dl = data.DataLoader(\n    test_ds,\n    batch_size=Config.INFERENCE_BATCH_SIZE,\n    shuffle=Config.SHOULD_SHUFFLE_INFERENCE_DATALOADER,\n)\n\n\nresults = []\nwith torch.no_grad():\n    for batch_idx, batch in enumerate(test_dl):\n        images, filenames = batch\n\n        images = images.to(Config.DEVICE)\n\n        # Get result\n        classification_result, segmentation_result = model(images)\n\n        # Add activation\n        classification_result = Config.CLASSIFICATION_INFERENCE_ACTIVATION(\n            classification_result\n        )\n        segmentation_result = Config.SEGMENTATION_INFERENCE_ACTIVATION(\n            segmentation_result\n        )\n\n        # Apply thresholding\n        classification_result = torch.where(classification_result > 0.5, 1, 0)\n        segmentation_result = torch.where(segmentation_result > 0.5, 1, 0)\n\n        # Squeeze results so that it becomes easy for rle_encoding\n        classification_result = classification_result.squeeze(-1)\n        segmentation_result = segmentation_result.squeeze(1)\n\n        # Convert to numpy array for rle_encoding\n        classification_result = classification_result.cpu().numpy()\n        segmentation_result = segmentation_result.cpu().numpy()\n\n        for idx in range(len(classification_result)):\n            filename = filenames[idx]\n            _, filename, _ = get_file_stats(filename)\n\n            val = classification_result[idx]\n\n            annotation = \"authentic\"\n\n            if val == 1:\n                mask = segmentation_result[idx, :, :].astype(np.uint8)\n                mask = _rle_encode_jit(mask)\n                mask = json.dumps(mask)\n                annotation = mask\n\n            result = {\n                \"case_id\": filename,\n                \"annotation\": annotation,\n            }\n\n            results.append(result)\n\nwrite_submission_csv(Config.SUBMISSIONS_DIR, results, [\"case_id\", \"annotation\"])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Quit Interpreter and Kaggle Session to save GPU Quota\nexit(0)","metadata":{},"outputs":[],"execution_count":null}]}