{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":132732,"databundleVersionId":16583342}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"49615b93","cell_type":"markdown","source":"# Synthetic Image Source Attribution\n\nThis notebook solves a 10-class image attribution problem: given a synthetic face image, identify which of 10 text-to-image models generated it. The dataset comes from the ICANN 2026 DLMMDD Workshop Challenge.\n\nWe use a dual-stream architecture that processes each image twice once as RGB and once through SRM noise filters then fuses both feature sets for classification. By the end of this notebook we'll have a trained model and a `submission.csv` ready to upload.","metadata":{}},{"id":"6ae8f935","cell_type":"markdown","source":"## Setup\n\nLet's start by downloading the competition data and importing everything we need.","metadata":{}},{"id":"004c760b","cell_type":"code","source":"import kagglehub\n\npath = kagglehub.competition_download('dlmmdd-workshop-synthetic-source-attribution-challenge')\n\nprint(\"Path to competition files:\", path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.254569Z","iopub.execute_input":"2026-05-27T22:56:49.255291Z","iopub.status.idle":"2026-05-27T22:56:49.793320Z","shell.execute_reply.started":"2026-05-27T22:56:49.255259Z","shell.execute_reply":"2026-05-27T22:56:49.792305Z"}},"outputs":[],"execution_count":null},{"id":"9423e72b","cell_type":"code","source":"import os\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nfrom PIL import Image\nimport cv2\n%matplotlib inline\n\n# PyTorch core\nimport torch\nimport torch.nn as nn\n\n# TorchVision\nimport torchvision.transforms.v2 as v2\nfrom torchvision.models import efficientnet_b4, EfficientNet_B4_Weights\nfrom torchvision.utils import make_grid\n\n# Data utilities\nfrom torch.utils.data.dataloader import DataLoader\nfrom torch.utils.data import Dataset\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.794816Z","iopub.execute_input":"2026-05-27T22:56:49.795128Z","iopub.status.idle":"2026-05-27T22:56:49.802127Z","shell.execute_reply.started":"2026-05-27T22:56:49.795096Z","shell.execute_reply":"2026-05-27T22:56:49.801484Z"}},"outputs":[],"execution_count":null},{"id":"95f07454","cell_type":"markdown","source":"## Configuration\n\nAll hyperparameters and paths live in one place so they're easy to change.","metadata":{}},{"id":"f8ddb03d","cell_type":"code","source":"MAIN_PATH = \"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data\"\nTRAIN_PATH = os.path.join(MAIN_PATH, 'Training')\nTEST_PATH  = os.path.join(MAIN_PATH, 'Test')\nTRAIN_CSV_PATH     = os.path.join(MAIN_PATH, 'training.csv')\nTEST_CSV_PATH      = os.path.join(MAIN_PATH, 'test.csv')\nLABEL_TO_NAMES_PATH = os.path.join(MAIN_PATH, 'sources.txt')\n\nNUM_OF_CLASSES = 10\n\n# 384 is EfficientNet-B4's native resolution, giving richer feature maps than 224\nSIZE       = 384\nBATCH_SIZE = 64\nEPOCHS     = 20\nPATIENCE   = 5\nLR         = 1e-4\n\n# ImageNet statistics used to normalize inputs for pretrained models\nMEAN_NORM = [0.485, 0.456, 0.406]\nSTD_NORM  = [0.229, 0.224, 0.225]\n\nNUM_WORKERS = 4\nPIN_MEMORY  = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.803205Z","iopub.execute_input":"2026-05-27T22:56:49.803808Z","iopub.status.idle":"2026-05-27T22:56:49.817161Z","shell.execute_reply.started":"2026-05-27T22:56:49.803772Z","shell.execute_reply":"2026-05-27T22:56:49.816528Z"}},"outputs":[],"execution_count":null},{"id":"14b77d1a","cell_type":"markdown","source":"## Exploratory Data Analysis\n\nBefore building anything, let's get a feel for the data class distribution, image formats, and what the images actually look like.","metadata":{}},{"id":"6d627f67","cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV_PATH)\ntrain_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.818825Z","iopub.execute_input":"2026-05-27T22:56:49.819099Z","iopub.status.idle":"2026-05-27T22:56:49.846166Z","shell.execute_reply.started":"2026-05-27T22:56:49.819077Z","shell.execute_reply":"2026-05-27T22:56:49.845515Z"}},"outputs":[],"execution_count":null},{"id":"21b8557e","cell_type":"code","source":"train_df['y'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.846878Z","iopub.execute_input":"2026-05-27T22:56:49.847168Z","iopub.status.idle":"2026-05-27T22:56:49.852849Z","shell.execute_reply.started":"2026-05-27T22:56:49.847146Z","shell.execute_reply":"2026-05-27T22:56:49.852280Z"}},"outputs":[],"execution_count":null},{"id":"efda14f5","cell_type":"code","source":"def show_all_extensions(path):\n    \"\"\"\n    Collect all unique file extensions found in a directory.\n\n    Args:\n        path: directory path to scan\n\n    Returns:\n        list of unique extension strings (e.g. ['jpg', 'png'])\n    \"\"\"\n    imgs_names = os.listdir(path)\n    result = []\n    for img_name in imgs_names:\n        x = img_name.split('.')[-1]\n        if x not in result:\n            result.append(x)\n    return result\n\nshow_all_extensions(TRAIN_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.853704Z","iopub.execute_input":"2026-05-27T22:56:49.853980Z","iopub.status.idle":"2026-05-27T22:56:49.873090Z","shell.execute_reply.started":"2026-05-27T22:56:49.853946Z","shell.execute_reply":"2026-05-27T22:56:49.872544Z"}},"outputs":[],"execution_count":null},{"id":"f1957aea-16b5-4c97-b755-937b4319854f","cell_type":"markdown","source":"I Commented here because it takes some time 3 - 5 minutes","metadata":{}},{"id":"98cd1e13","cell_type":"code","source":"# def show_all_imgs_sizes(path):\n#     \"\"\"\n#     Collect all unique image shapes found in a directory.\n\n#     Args:\n#         path: directory path to scan\n\n#     Returns:\n#         list of unique shape tuples (H, W, C)\n#     \"\"\"\n#     imgs_names = os.listdir(path)\n#     all_sizes = []\n#     for img_name in imgs_names:\n#         img_shape = plt.imread(os.path.join(path, img_name)).shape\n#         if img_shape not in all_sizes:\n#             all_sizes.append(img_shape)\n#     return all_sizes\n\n# show_all_imgs_sizes(TRAIN_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.873758Z","iopub.execute_input":"2026-05-27T22:56:49.874056Z","iopub.status.idle":"2026-05-27T22:56:49.877962Z","shell.execute_reply.started":"2026-05-27T22:56:49.874007Z","shell.execute_reply":"2026-05-27T22:56:49.877375Z"}},"outputs":[],"execution_count":null},{"id":"d8b8a97c","cell_type":"code","source":"def show_img_per_class():\n    \"\"\"Display one random image from each class.\"\"\"\n    unique_classes = sorted(train_df['y'].unique())\n    num_classes = len(unique_classes)\n\n    cols = min(4, num_classes)\n    rows = (num_classes + cols - 1) // cols\n\n    plt.figure(figsize=(cols * 4, rows * 4))\n\n    for idx, cls in enumerate(unique_classes):\n        class_df = train_df[train_df['y'] == cls]\n        img_path = np.random.choice(class_df[\"path\"].values)\n        img_full_path = os.path.join(TRAIN_PATH, img_path.split('/')[-1])\n        img = plt.imread(img_full_path)\n\n        # grayscale images need to be converted before display\n        if img.ndim == 2:\n            img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n\n        plt.subplot(rows, cols, idx + 1)\n        plt.imshow(img)\n        plt.title(f\"Class/Label: {cls}\")\n        plt.axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n\nshow_img_per_class()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:49.878725Z","iopub.execute_input":"2026-05-27T22:56:49.879001Z","iopub.status.idle":"2026-05-27T22:56:52.291084Z","shell.execute_reply.started":"2026-05-27T22:56:49.878968Z","shell.execute_reply":"2026-05-27T22:56:52.290240Z"}},"outputs":[],"execution_count":null},{"id":"37dd73f9","cell_type":"markdown","source":"## Transforms and Dataset\n\nNow let's define how images get prepared before entering the model.\n\nThe training transform includes light augmentation to make the model more robust to the post-processing operations present in the test set. The val/test transform only resizes and normalizes no augmentation.","metadata":{}},{"id":"dd62a7a7","cell_type":"code","source":"def get_transforms(size=SIZE, apply_on_train=False):\n    \"\"\"\n    Build a transform pipeline for training or validation/test.\n\n    Args:\n        size: target square image size in pixels\n        apply_on_train: if True, adds augmentation steps before normalization\n\n    Returns:\n        v2.Compose transform pipeline\n    \"\"\"\n    base = [\n        v2.Resize((size, size)),\n        v2.ToImage(),\n        v2.ToDtype(torch.float32, scale=True),\n    ]\n\n    augmentation = [\n        v2.RandomHorizontalFlip(p=0.5),\n        # simulate grayscale conversion, which is one of the test post-processing ops\n        v2.RandomGrayscale(p=0.1),\n        # keep values small to avoid destroying forensic color signatures\n        v2.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.05),\n    ]\n\n    # ImageNet mean/std normalization (required for pretrained models)\n    tail = [v2.Normalize(MEAN_NORM, STD_NORM)]\n\n    if apply_on_train:\n        return v2.Compose(base + augmentation + tail)\n    return v2.Compose(base + tail)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.292169Z","iopub.execute_input":"2026-05-27T22:56:52.292428Z","iopub.status.idle":"2026-05-27T22:56:52.299494Z","shell.execute_reply.started":"2026-05-27T22:56:52.292404Z","shell.execute_reply":"2026-05-27T22:56:52.298684Z"}},"outputs":[],"execution_count":null},{"id":"7ad99040","cell_type":"code","source":"class SyntheticDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for the synthetic image attribution challenge.\n\n    Handles both train/val (returns image + label) and test (returns image + ID)\n    via the is_test flag.\n\n    Args:\n        dataframe: pandas DataFrame with columns [ID, path, y] for train or [ID, path] for test\n        transform: torchvision transform pipeline to apply to each image\n        is_test: if True, returns (image, ID) instead of (image, label)\n    \"\"\"\n    def __init__(self, dataframe, transform=None, is_test=False):\n        self.df = dataframe\n        self.transform = transform\n        self.is_test = is_test\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        image_path = self.df.iloc[idx, 1]\n        img = Image.open(image_path).convert('RGB')\n        if self.transform:\n            img = self.transform(img)\n\n        if self.is_test:\n            # test set has no labels, return the image ID for submission matching\n            id_ = self.df.iloc[idx, 0]\n            return img, id_\n        else:\n            label = self.df.iloc[idx, 2]\n            return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.301695Z","iopub.execute_input":"2026-05-27T22:56:52.302092Z","iopub.status.idle":"2026-05-27T22:56:52.316808Z","shell.execute_reply.started":"2026-05-27T22:56:52.302061Z","shell.execute_reply":"2026-05-27T22:56:52.316043Z"}},"outputs":[],"execution_count":null},{"id":"8f8b066a","cell_type":"markdown","source":"## Data Loading\n\nNow we load the CSVs, fix the image paths to point to the correct directory, and split the training set into train and validation.","metadata":{}},{"id":"daebeebb","cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV_PATH)\ntest_df  = pd.read_csv(TEST_CSV_PATH)\n\n# replace relative paths in the CSV with absolute paths on disk\ntrain_df['path'] = train_df['path'].apply(lambda x: os.path.join(TRAIN_PATH, x.split('/')[-1]))\ntest_df['path']  = test_df['path'].apply(lambda x: os.path.join(TEST_PATH,  x.split('/')[-1]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.317884Z","iopub.execute_input":"2026-05-27T22:56:52.318142Z","iopub.status.idle":"2026-05-27T22:56:52.365590Z","shell.execute_reply.started":"2026-05-27T22:56:52.318119Z","shell.execute_reply":"2026-05-27T22:56:52.365077Z"}},"outputs":[],"execution_count":null},{"id":"3dc9bf9b","cell_type":"code","source":"train_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.366420Z","iopub.execute_input":"2026-05-27T22:56:52.366729Z","iopub.status.idle":"2026-05-27T22:56:52.374916Z","shell.execute_reply.started":"2026-05-27T22:56:52.366705Z","shell.execute_reply":"2026-05-27T22:56:52.374102Z"}},"outputs":[],"execution_count":null},{"id":"9af3a163","cell_type":"code","source":"test_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.375962Z","iopub.execute_input":"2026-05-27T22:56:52.376851Z","iopub.status.idle":"2026-05-27T22:56:52.390343Z","shell.execute_reply.started":"2026-05-27T22:56:52.376825Z","shell.execute_reply":"2026-05-27T22:56:52.389646Z"}},"outputs":[],"execution_count":null},{"id":"00aad236","cell_type":"code","source":"# stratify ensures each class keeps the same proportion in train and val\ntrain_df, val_df = train_test_split(\n    train_df,\n    test_size=0.2,\n    random_state=42,\n    stratify=train_df[\"y\"]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.391279Z","iopub.execute_input":"2026-05-27T22:56:52.391832Z","iopub.status.idle":"2026-05-27T22:56:52.409297Z","shell.execute_reply.started":"2026-05-27T22:56:52.391796Z","shell.execute_reply":"2026-05-27T22:56:52.408644Z"}},"outputs":[],"execution_count":null},{"id":"c16bb0db","cell_type":"code","source":"train_dataset = SyntheticDataset(train_df, get_transforms(SIZE, True))\nval_dataset   = SyntheticDataset(val_df,   get_transforms(SIZE))\ntest_dataset  = SyntheticDataset(test_df,  get_transforms(SIZE), is_test=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.410393Z","iopub.execute_input":"2026-05-27T22:56:52.410685Z","iopub.status.idle":"2026-05-27T22:56:52.414872Z","shell.execute_reply.started":"2026-05-27T22:56:52.410662Z","shell.execute_reply":"2026-05-27T22:56:52.414138Z"}},"outputs":[],"execution_count":null},{"id":"dbd4e839","cell_type":"code","source":"train_dl = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,  num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY)\n# shuffle=False for val and test to keep ID order consistent\nval_dl   = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY)\ntest_dl  = DataLoader(test_dataset,  batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=PIN_MEMORY)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.415657Z","iopub.execute_input":"2026-05-27T22:56:52.416015Z","iopub.status.idle":"2026-05-27T22:56:52.427836Z","shell.execute_reply.started":"2026-05-27T22:56:52.415986Z","shell.execute_reply":"2026-05-27T22:56:52.427068Z"}},"outputs":[],"execution_count":null},{"id":"219cd185","cell_type":"markdown","source":"Let's do a quick sanity check and visualize one batch to confirm the transforms are working correctly.","metadata":{}},{"id":"a24b6c37","cell_type":"code","source":"def denorm(imgs):\n    \"\"\"\n    Reverse ImageNet normalization so images can be displayed correctly.\n\n    Args:\n        imgs: normalized tensor of shape (B, C, H, W) or (C, H, W)\n\n    Returns:\n        tensor in [0, 1] range with same shape as input\n    \"\"\"\n    mean = torch.tensor(MEAN_NORM).view(1, 3, 1, 1).to(imgs.device)\n    std  = torch.tensor(STD_NORM).view(1, 3, 1, 1).to(imgs.device)\n\n    if imgs.dim() == 3:\n        imgs = imgs.unsqueeze(0)\n\n    return imgs * std + mean","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.428782Z","iopub.execute_input":"2026-05-27T22:56:52.429083Z","iopub.status.idle":"2026-05-27T22:56:52.442608Z","shell.execute_reply.started":"2026-05-27T22:56:52.429049Z","shell.execute_reply":"2026-05-27T22:56:52.441861Z"}},"outputs":[],"execution_count":null},{"id":"f482916e","cell_type":"code","source":"def show_grid_images(dataloader):\n    \"\"\"Display a grid of images from one batch (useful for sanity-checking augmentation).\"\"\"\n    imgs, labels = next(iter(dataloader))\n    plt.figure(figsize=(30, 30))\n    plt.imshow(make_grid(denorm(imgs), nrow=16).permute(1, 2, 0))\n    plt.axis('off')\n    plt.tight_layout()\n    plt.show()\n\nshow_grid_images(train_dl)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:56:52.443566Z","iopub.execute_input":"2026-05-27T22:56:52.443983Z","iopub.status.idle":"2026-05-27T22:57:06.517381Z","shell.execute_reply.started":"2026-05-27T22:56:52.443948Z","shell.execute_reply":"2026-05-27T22:57:06.516340Z"}},"outputs":[],"execution_count":null},{"id":"1c3045de","cell_type":"markdown","source":"## Model Architecture\n\nWe use two components before building the full model.\n\n**SRMConv** applies Spatial Rich Model filters fixed convolutional kernels from the image forensics literature. They extract high-frequency noise residuals that reveal traces left by each generative model, patterns invisible to the naked eye.\n\n**DualStreamModel** runs two parallel EfficientNet-B4 backbones: one on the raw RGB image, one on the SRM-filtered version. The features from both streams are concatenated and passed to a classification head.","metadata":{}},{"id":"6a9449a5","cell_type":"code","source":"class SRMConv(nn.Module):\n    \"\"\"\n    Fixed convolutional layer applying three Spatial Rich Model (SRM) filters.\n\n    These filters extract noise residuals from images forensic artifacts that\n    reveal which generative model produced the image. Weights are fixed (not trained).\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n\n        # high-pass filter sensitive to local pixel inconsistencies\n        f1 = np.array([[ 0,  0,  0,  0,  0],\n                       [ 0, -1,  2, -1,  0],\n                       [ 0,  2, -4,  2,  0],\n                       [ 0, -1,  2, -1,  0],\n                       [ 0,  0,  0,  0,  0]]) / 4.0\n\n        # captures second-order noise statistics across a wider neighborhood\n        f2 = np.array([[-1,  2, -2,  2, -1],\n                       [ 2, -6,  8, -6,  2],\n                       [-2,  8,-12,  8, -2],\n                       [ 2, -6,  8, -6,  2],\n                       [-1,  2, -2,  2, -1]]) / 12.0\n\n        # simple edge-residual filter\n        f3 = np.array([[ 0,  0,  0,  0,  0],\n                       [ 0,  0,  0,  0,  0],\n                       [ 0,  1, -2,  1,  0],\n                       [ 0,  0,  0,  0,  0],\n                       [ 0,  0,  0,  0,  0]]) / 2.0\n\n        filters = np.stack([f1, f2, f3], axis=0)       # (3, 5, 5)\n        filters = np.expand_dims(filters, axis=1)       # (3, 1, 5, 5)\n        # apply each filter to all 3 RGB channels independently\n        filters = np.repeat(filters, 3, axis=1)         # (3, 3, 5, 5)\n\n        # requires_grad=False keeps these weights frozen during training\n        self.weight = nn.Parameter(\n            torch.tensor(filters, dtype=torch.float32),\n            requires_grad=False\n        )\n\n    def forward(self, x):\n        return nn.functional.conv2d(x, self.weight, padding=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:57:06.519304Z","iopub.execute_input":"2026-05-27T22:57:06.519710Z","iopub.status.idle":"2026-05-27T22:57:06.528969Z","shell.execute_reply.started":"2026-05-27T22:57:06.519659Z","shell.execute_reply":"2026-05-27T22:57:06.528398Z"}},"outputs":[],"execution_count":null},{"id":"9c930ae2","cell_type":"code","source":"class DualStreamModel(nn.Module):\n    \"\"\"\n    Two-stream classification model for synthetic image attribution.\n\n    Stream 1 processes the raw RGB image through EfficientNet-B4.\n    Stream 2 applies SRM noise filters first, then passes the result through\n    a second EfficientNet-B4. Both feature vectors are concatenated and fed\n    into a fully connected classification head.\n\n    Args:\n        num_classes: number of output classes (10 for this challenge)\n    \"\"\"\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        # stream 1: standard RGB features\n        self.rgb_backbone = efficientnet_b4(weights=EfficientNet_B4_Weights.DEFAULT)\n        rgb_features = self.rgb_backbone.classifier[1].in_features\n        # replace classifier with identity to get raw feature vectors\n        self.rgb_backbone.classifier = nn.Identity()\n        self.backbone_name = self.rgb_backbone.__class__.__name__\n\n        # stream 2: forensic noise features via SRM\n        self.srm = SRMConv()\n        self.srm_backbone = efficientnet_b4(weights=EfficientNet_B4_Weights.DEFAULT)\n        self.srm_backbone.classifier = nn.Identity()\n\n        # fusion head: takes concatenated features from both streams\n        self.classifier = nn.Sequential(\n            nn.Linear(rgb_features * 2, 512),\n            nn.ReLU(),\n            nn.Dropout(p=0.3),\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        rgb_feat = self.rgb_backbone(x)\n        srm_feat = self.srm_backbone(self.srm(x))\n        combined = torch.cat([rgb_feat, srm_feat], dim=1)\n        return self.classifier(combined)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:57:06.529899Z","iopub.execute_input":"2026-05-27T22:57:06.530146Z","iopub.status.idle":"2026-05-27T22:57:06.543760Z","shell.execute_reply.started":"2026-05-27T22:57:06.530125Z","shell.execute_reply":"2026-05-27T22:57:06.542829Z"}},"outputs":[],"execution_count":null},{"id":"6e6133a4","cell_type":"markdown","source":"## Training Setup\n\nWe instantiate the model and configure the optimizer, loss function, learning rate scheduler, and mixed precision scaler.\n\nWe use AdamW with OneCycleLR the scheduler ramps the learning rate up then down over the course of training, which often converges faster than a fixed LR. Mixed precision (float16) cuts memory usage and speeds up training on modern GPUs.","metadata":{}},{"id":"98330b71","cell_type":"code","source":"model = DualStreamModel(num_classes=NUM_OF_CLASSES).to(device)\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-3)\n\n# wrap in DataParallel if multiple GPUs are available\nif torch.cuda.device_count() > 1:\n    model = nn.DataParallel(model)\n    # preserve backbone_name attribute through the DataParallel wrapper\n    model.backbone_name = model.module.backbone_name\n\ncriterion = nn.CrossEntropyLoss()\n\n# OneCycleLR steps every batch, not every epoch\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=LR,\n    steps_per_epoch=len(train_dl),\n    epochs=EPOCHS\n)\n\n# GradScaler enables mixed precision training\nscaler = torch.amp.GradScaler(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:57:06.544772Z","iopub.execute_input":"2026-05-27T22:57:06.545201Z","iopub.status.idle":"2026-05-27T22:57:07.979626Z","shell.execute_reply.started":"2026-05-27T22:57:06.545163Z","shell.execute_reply":"2026-05-27T22:57:07.979053Z"}},"outputs":[],"execution_count":null},{"id":"31d3fc52","cell_type":"markdown","source":"## Training Loop\n\nThe `fit` function runs the full training and validation loop. It saves the best checkpoint whenever validation accuracy improves, and stops early if there's no improvement for `PATIENCE` epochs.","metadata":{}},{"id":"ac9ccc32","cell_type":"code","source":"def fit(model):\n    \"\"\"\n    Train the model with early stopping based on validation accuracy.\n\n    Each epoch runs the training loop, evaluates on validation, saves the\n    best checkpoint, and stops if val accuracy doesn't improve for PATIENCE epochs.\n\n    Args:\n        model: nn.Module to train (uses global optimizer, criterion, scheduler, scaler)\n\n    Returns:\n        history: list of dicts with train_loss, train_acc, val_loss, val_acc per epoch\n    \"\"\"\n    history  = []\n    best_acc = 0.0\n    counter  = 0\n\n    for epoch in range(EPOCHS):\n        train_all_loss = []\n        train_all_acc  = []\n\n        model.train()\n\n        for batch in tqdm(train_dl, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n            imgs, labels = batch\n            imgs   = imgs.to(device)\n            labels = labels.to(device)\n\n            # zero gradients before the forward pass\n            optimizer.zero_grad()\n\n            # autocast runs the forward pass in float16 for speed\n            with torch.amp.autocast(device_type='cuda'):\n                outputs = model(imgs)\n                loss    = criterion(outputs, labels)\n\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            # OneCycleLR must step every batch\n            scheduler.step()\n\n            _, preds = torch.max(outputs, dim=1)\n            train_all_loss.append(loss.detach().cpu().item())\n            train_all_acc.append((preds == labels).float().mean().item())\n\n        train_loss = sum(train_all_loss) / len(train_all_loss)\n        train_acc  = sum(train_all_acc)  / len(train_all_acc)\n\n        val_all_loss = []\n        val_all_acc  = []\n\n        model.eval()\n        with torch.no_grad():\n            for batch in val_dl:\n                imgs, labels = batch\n                imgs   = imgs.to(device)\n                labels = labels.to(device)\n\n                outputs = model(imgs)\n                loss    = criterion(outputs, labels)\n\n                _, preds = torch.max(outputs, dim=1)\n                val_all_loss.append(loss.detach().cpu().item())\n                val_all_acc.append((preds == labels).float().mean().item())\n\n        val_loss = sum(val_all_loss) / len(val_all_loss)\n        val_acc  = sum(val_all_acc)  / len(val_all_acc)\n\n        print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | \"\n              f\"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}\")\n\n        history.append({\n            \"train_loss\": train_loss, \"train_acc\": train_acc,\n            \"val_loss\":   val_loss,   \"val_acc\":   val_acc\n        })\n\n        if val_acc > best_acc:\n            best_acc = val_acc\n            counter  = 0\n            torch.save(model.state_dict(), f\"best_{model.backbone_name}_model.pth\")\n            print(\"  >> best model saved\")\n        else:\n            counter += 1\n            print(f\"  No improvement ({counter}/{PATIENCE})\")\n            if counter >= PATIENCE:\n                print(\"Early stopping triggered.\")\n                break\n\n    return history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:57:07.980545Z","iopub.execute_input":"2026-05-27T22:57:07.980832Z","iopub.status.idle":"2026-05-27T22:57:07.991092Z","shell.execute_reply.started":"2026-05-27T22:57:07.980788Z","shell.execute_reply":"2026-05-27T22:57:07.990464Z"}},"outputs":[],"execution_count":null},{"id":"98d03e2a","cell_type":"markdown","source":"## Inference with Test-Time Augmentation\n\nInstead of predicting once per image, we run each test image through the model `n_augments` times with slightly different augmentations and average the softmax probabilities. This reduces variance and usually improves accuracy by a small margin.","metadata":{}},{"id":"cb2aa0bd","cell_type":"code","source":"def predict_with_tta(model, loader, device, n_augments=5):\n    \"\"\"\n    Run inference with Test-Time Augmentation (TTA).\n\n    Each image is augmented n_augments times and the softmax probabilities\n    are averaged before taking the argmax prediction.\n\n    Args:\n        model: trained nn.Module in eval mode\n        loader: DataLoader that returns (imgs, ids) pairs\n        device: torch device to run inference on\n        n_augments: number of augmented views to average per image\n\n    Returns:\n        all_ids: list of image IDs matching the test CSV\n        all_preds: list of predicted class labels (integers 0-9)\n    \"\"\"\n    model.eval()\n\n    # light augmentation enough to add diversity without destroying forensic features\n    tta_transform = v2.Compose([\n        v2.Resize((SIZE, SIZE)),\n        v2.RandomHorizontalFlip(p=0.5),\n        v2.ColorJitter(brightness=0.1, contrast=0.1),\n        v2.ToImage(),\n        v2.ToDtype(torch.float32, scale=True),\n        v2.Normalize(MEAN_NORM, STD_NORM)\n    ])\n\n    all_ids, all_preds = [], []\n\n    with torch.no_grad():\n        for imgs, ids in tqdm(loader):\n            imgs = imgs.to(device)\n\n            # accumulate softmax probabilities across all augmented views\n            probs = torch.zeros(imgs.size(0), NUM_OF_CLASSES).to(device)\n            for _ in range(n_augments):\n                aug_imgs = torch.stack([tta_transform(img) for img in imgs])\n                aug_imgs = aug_imgs.to(device)\n                outputs  = model(aug_imgs)\n                probs   += torch.softmax(outputs, dim=1)\n\n            preds = probs.argmax(dim=1).cpu().tolist()\n            all_preds.extend(preds)\n            all_ids.extend(ids.tolist() if torch.is_tensor(ids) else ids)\n\n    return all_ids, all_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-27T22:57:07.992095Z","iopub.execute_input":"2026-05-27T22:57:07.992441Z","iopub.status.idle":"2026-05-27T22:57:08.009454Z","shell.execute_reply.started":"2026-05-27T22:57:07.992419Z","shell.execute_reply":"2026-05-27T22:57:08.008620Z"}},"outputs":[],"execution_count":null},{"id":"6b4b18c0","cell_type":"markdown","source":"## Run Training\n\nLet's train the model. The best checkpoint is saved automatically.","metadata":{}},{"id":"ba64985e","cell_type":"code","source":"# No need to train again and the GPU cand handle with EffecientNetB4 with Size=384\n# You can Un-Comment here and try!\n# history = fit(model)","metadata":{},"outputs":[],"execution_count":null},{"id":"00ff7b12","cell_type":"markdown","source":"## Generate Submission\n\nWe load the best checkpoint, run inference with TTA on the test set, and save the predictions as `submission.csv`.","metadata":{}},{"id":"35d5b190","cell_type":"code","source":"# load the best checkpoint saved during training\nmodel.load_state_dict(torch.load(f\"best_{model.backbone_name}_model.pth\", map_location=device))\n\nall_ids, all_preds = predict_with_tta(model, test_dl, device, n_augments=5)\n\nsubmission = pd.DataFrame({\"ID\": all_ids, \"TARGET\": all_preds})\nsubmission = submission.sort_values(\"ID\").reset_index(drop=True)\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(submission.head(10))\nprint(f\"Total predictions: {len(submission)}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"01e3ebf4-78a1-4c09-aca3-834f6c31ae9b","cell_type":"markdown","source":"THE END...","metadata":{}}]}