{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"colab":{"provenance":[],"gpuType":"T4"},"accelerator":"GPU","kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431},{"sourceType":"datasetVersion","sourceId":15660789,"datasetId":10026946,"databundleVersionId":16597384},{"sourceType":"modelInstanceVersion","sourceId":826799,"databundleVersionId":16596936,"modelInstanceId":628612,"modelId":640535}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# --- 1. Multi-scale Squeeze and Excitation (MSSE) Block ---\nimport torch\nimport torch.nn as nn\n\nclass MSSEBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, reduction=8):\n        super(MSSEBlock, self).__init__()\n\n        # --- Multi-scale Feature Extraction (Green Boxes) ---\n        # The figure shows a chain: 3x3 -> 3x3 -> 3x3 (effectively 3x3, 5x5, 7x7)\n        self.conv3x3 = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)\n        self.conv5x5 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n        self.conv7x7 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n\n        # --- SE Block (Purple Box) ---\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.se_module = nn.Sequential(\n            nn.Linear(out_channels * 3, (out_channels * 3) // reduction, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear((out_channels * 3) // reduction, out_channels * 3, bias=False),\n            nn.Sigmoid()\n        )\n\n        # --- 1x1 Convolution (Orange Box) ---\n        # This reduces the concatenated 3x channels back to out_channels\n        self.conv1x1 = nn.Conv2d(out_channels * 3, out_channels, kernel_size=1)\n\n        # --- Residual Shortcut (Path b) ---\n        self.shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity()\n\n        # --- Post-Processing (Red Boxes) ---\n        self.relu = nn.ReLU(inplace=True)\n        self.bn = nn.BatchNorm2d(out_channels)\n\n    def forward(self, x):\n        # Path (b): The global residual\n        residual = self.shortcut(x)\n\n        # Multi-scale sequence (Path a)\n        out3 = self.conv3x3(x)\n        out5 = self.conv5x5(out3)\n        out7 = self.conv7x7(out5)\n\n        # Concatenation (Blue Box)\n        # We concatenate the output of the 3x3, the 5x5, and the 7x7\n        concat = torch.cat([out3, out5, out7], dim=1)\n\n        # SE Block calculation\n        b, c, _, _ = concat.size()\n        se_weight = self.avg_pool(concat).view(b, c)\n        se_weight = self.se_module(se_weight).view(b, c, 1, 1)\n\n        # Addition #1: SE Block output + Concatenation output\n        # (This applies the excitation weights to the concatenated map)\n        excited = concat * se_weight + concat\n\n        # 1x1 Convolution\n        fused = self.conv1x1(excited)\n\n        # Addition #2: Result + Global Input (Residual)\n        out = fused + residual\n\n        # Final activations in order: ReLU -> BN\n        out = self.relu(out)\n        out = self.bn(out)\n\n        return out\n# --- 2. Bottleneck Residual (B-Res) Path ---\nclass BResBlock(nn.Module):\n    def __init__(self, channels):\n        super(BResBlock, self).__init__()\n        # Bottleneck: 1x1 -> 3x3 -> 1x1\n        self.branch = nn.Sequential(\n            nn.Conv2d(channels, channels // 4, kernel_size=1),\n            nn.BatchNorm2d(channels // 4),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(channels // 4, channels // 4, kernel_size=3, padding=1),\n            nn.BatchNorm2d(channels // 4),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(channels // 4, channels, kernel_size=1),\n            nn.BatchNorm2d(channels)\n        )\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, x):\n        return self.relu(x + self.branch(x))\n\nclass BResPath(nn.Module):\n    def __init__(self, channels, num_blocks):\n        super(BResPath, self).__init__()\n        layers = [BResBlock(channels) for _ in range(num_blocks)]\n        self.path = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.path(x)\n\n# --- 3. Spatial Attention Block (SAB) ---\nclass SpatialAttentionBlock(nn.Module):\n    def __init__(self, F_g, F_l, F_int):\n        super(SpatialAttentionBlock, self).__init__()\n\n        # Linear transformations for gating (g) and skip (x)\n        self.W_g = nn.Sequential(\n            nn.Conv2d(F_g, F_int, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        self.W_x = nn.Sequential(\n            nn.Conv2d(F_l, F_int, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n\n        # The internal psi path (Conv + BN + Sigmoid)\n        self.psi = nn.Sequential(\n            nn.Conv2d(F_int, 1, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n\n        self.relu = nn.ReLU(inplace=True)\n\n        # --- THE RESAMPLER ---\n        # This is positioned to upscale the attention coefficients\n        # so they match the spatial size of 'x' BEFORE multiplication.\n        #self.resampler = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n\n    def forward(self, g, x):\n        # 1. Processing paths\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n\n        # 2. Additive Attention (Combined at lower resolution)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi) # This is the raw attention mask\n\n        # 3. RESAMPLER applied BEFORE multiplication\n        # Upscales the mask to match x's dimensions\n        #psi_upscaled = self.resampler(psi)\n\n        # 4. Multiplication (Gating)\n        # Now the high-resolution features 'x' are filtered by the upscaled mask\n        return x * psi\n\n# --- 4. MSMA-Net Full Architecture ---\nclass MSMA_Net(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(MSMA_Net, self).__init__()\n\n        # Encoder (Downsampling)\n        self.enc1 = MSSEBlock(in_channels, 32)\n        self.enc2 = MSSEBlock(32, 64)\n        self.enc3 = MSSEBlock(64, 128)\n        self.enc4 = MSSEBlock(128, 256)\n\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n        # Bottleneck (Bridge)\n        self.bottleneck = MSSEBlock(256, 512)\n\n        # B-Res Paths (Replacing standard skip connections)\n        # Lengths as specified in paper: 4, 3, 2, 1\n        self.bres1 = BResPath(32, 4)\n        self.bres2 = BResPath(64, 3)\n        self.bres3 = BResPath(128, 2)\n        self.bres4 = BResPath(256, 1)\n\n        # Decoder (Upsampling)\n        self.up4 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)\n        self.sab4 = SpatialAttentionBlock(F_g=256, F_l=256, F_int=128)\n        self.dec4 = MSSEBlock(512, 256) # Cat(SAB, Up)\n\n        self.up3 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)\n        self.sab3 = SpatialAttentionBlock(F_g=128, F_l=128, F_int=64)\n        self.dec3 = MSSEBlock(256, 128)\n\n        self.up2 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        self.sab2 = SpatialAttentionBlock(F_g=64, F_l=64, F_int=32)\n        self.dec2 = MSSEBlock(128, 64)\n\n        self.up1 = nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2)\n        self.sab1 = SpatialAttentionBlock(F_g=32, F_l=32, F_int=32)\n        self.dec1 = MSSEBlock(64, 32)\n\n        # Final Output\n        self.final_conv = nn.Conv2d(32, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool(e1))\n        e3 = self.enc3(self.pool(e2))\n        e4 = self.enc4(self.pool(e3))\n\n        # Bridge\n        b = self.bottleneck(self.pool(e4))\n\n        # B-Res Paths\n        s1 = self.bres1(e1)\n        s2 = self.bres2(e2)\n        s3 = self.bres3(e3)\n        s4 = self.bres4(e4)\n\n        # Decoder\n        d4_up = self.up4(b)\n        d4_att = self.sab4(g=d4_up, x=s4)\n        d4 = self.dec4(torch.cat([d4_att, d4_up], dim=1))\n\n        d3_up = self.up3(d4)\n        d3_att = self.sab3(g=d3_up, x=s3)\n        d3 = self.dec3(torch.cat([d3_att, d3_up], dim=1))\n\n        d2_up = self.up2(d3)\n        d2_att = self.sab2(g=d2_up, x=s2)\n        d2 = self.dec2(torch.cat([d2_att, d2_up], dim=1))\n\n        d1_up = self.up1(d2)\n        d1_att = self.sab1(g=d1_up, x=s1)\n        d1 = self.dec1(torch.cat([d1_att, d1_up], dim=1))\n\n        #out = torch.sigmoid(self.final_conv(d1))\n        out=self.final_conv(d1)\n        return out","metadata":{"id":"OZ5jocxF2bMi","trusted":true,"execution":{"iopub.status.busy":"2026-06-24T19:40:25.971790Z","iopub.execute_input":"2026-06-24T19:40:25.972631Z","iopub.status.idle":"2026-06-24T19:40:30.000855Z","shell.execute_reply.started":"2026-06-24T19:40:25.972599Z","shell.execute_reply":"2026-06-24T19:40:29.999930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom tkinter import Image\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nimport PIL\nfrom PIL import Image\n\nimport cv2\nimport numpy as np\n\n\n\n","metadata":{"id":"Lbc-jZVN2i5e","outputId":"3e9c88cb-c570-4fac-ec78-d7b6dfe671c1","trusted":true,"execution":{"iopub.status.busy":"2026-06-24T19:42:09.288610Z","iopub.execute_input":"2026-06-24T19:42:09.288918Z","iopub.status.idle":"2026-06-24T19:42:09.303851Z","shell.execute_reply.started":"2026-06-24T19:42:09.288893Z","shell.execute_reply":"2026-06-24T19:42:09.302787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_clean_full_image(img):\n\n    # 1. Gamma Correction\n    gamma = 1.5\n    inv_gamma = 1.0 / gamma\n    table = np.array([((i/255.0)**inv_gamma)*255 for i in np.arange(0,256)]).astype(\"uint8\")\n\n    # 2. Green Channel\n    gray = img[:, :, 1]\n\n    gamma_grey=cv2.LUT(gray, table)\n\n\n\n\n    # 2. Local Noise Reduction\n    # 3. Global Shade Correction\n    h, w = gray.shape[:2] # Correctly use the denoised_gray image's shape\n    blur_size = int(h / 25)\n    if blur_size % 2 == 0: blur_size += 1\n    background = cv2.medianBlur(gray, blur_size) # Blur the denoised image\n    diff = cv2.subtract(background, gray) # High-pass: vessels become bright\n\n    clahe=cv2.createCLAHE(clipLimit=2.0,tileGridSize=(8,8))\n    denoised_gray=clahe.apply(diff)\n\n\n    kernel_denoise = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2, 2))\n    denoised_gray = cv2.morphologyEx(denoised_gray, cv2.MORPH_OPEN, kernel_denoise)\n\n\n    # Generate a binary mask of the actual retina (based on original gray for FOV)\n    _, mask_retina = cv2.threshold(gamma_grey, 15, 255, cv2.THRESH_BINARY)\n\n    # Erode the mask slightly to remove the \"bright ring\" artifact\n    # created by the subtract operation at the FOV boundary\n    erosion_size = max(5, int(h / 50))\n    kernel_border = np.ones((erosion_size, erosion_size), np.uint8)\n    mask_retina = cv2.erode(mask_retina, kernel_border, iterations=1)\n\n    # Final masked image\n    global_cleaned = cv2.bitwise_and(diff, diff, mask=mask_retina)\n    return global_cleaned","metadata":{"id":"8Dc7Q0CHuMbT","trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:17:06.939600Z","iopub.execute_input":"2026-04-11T13:17:06.940365Z","iopub.status.idle":"2026-04-11T13:17:06.946774Z","shell.execute_reply.started":"2026-04-11T13:17:06.940332Z","shell.execute_reply":"2026-04-11T13:17:06.946158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_full_res(model, high_res_img_tensor, patch_size=128, stride=64):\n    # 1. Store the ORIGINAL dimensions (assuming high_res_img_tensor is (B, C, H, W) from val_loader)\n    _, _, h_orig, w_orig = high_res_img_tensor.shape\n\n    # 2. Calculate padding needed to fit the stride/patch\n    pad_h = (stride - (h_orig % stride)) % stride\n    pad_w = (stride - (w_orig % stride)) % stride\n\n    # 3. Apply padding to the IMAGE ONLY\n    # F.pad format: (left, right, top, bottom) for the last two dimensions\n    img_padded = F.pad(high_res_img_tensor, (0, pad_w, 0, pad_h), mode='constant', value=0)\n    _, _, h_pad, w_pad = img_padded.shape # Correctly get padded dimensions from 4D tensor\n\n    # 4. Create accumulation maps for the PADDED size\n    full_prob_map = np.zeros((h_pad, w_pad), dtype=np.float32)\n    count_map = np.zeros((h_pad, w_pad), dtype=np.float32)\n\n    model.eval()\n    with torch.no_grad():\n        for y in range(0, h_pad - patch_size + 1, stride):\n            for x in range(0, w_pad - patch_size + 1, stride):\n                # Extract patch as a tensor (B, C, H, W). Here B=1, C=1\n                patch_tensor = img_padded[:, :, y:y+patch_size, x:x+patch_size]\n\n                # Convert to numpy for local preprocessing (e.g., CLAHE, Z-score)\n                # Squeeze batch and channel dimensions to get a 2D numpy array\n                patch_numpy = patch_tensor.squeeze(0).squeeze(0).cpu().numpy()\n\n                # Apply local preprocessing defined in process_patch\n                #processed_patch_numpy = process_patch(patch_numpy)\n\n                # Convert back to tensor, add batch and channel dimensions, and move to device\n                tensor_input = torch.from_numpy(patch_numpy).unsqueeze(0).unsqueeze(0).to(device).float()\n\n                # Model inference\n                output = torch.sigmoid(model(tensor_input))\n                pred_patch = output.squeeze().cpu().numpy()\n\n                # Stitching into the padded map\n                full_prob_map[y:y+patch_size, x:x+patch_size] += pred_patch\n                count_map[y:y+patch_size, x:x+patch_size] += 1.0\n\n    # 5. Average the overlaps\n    full_prob_map /= (count_map + 1e-6)\n\n    # 6. THE CRITICAL STEP: Crop back to the original size\n    # This removes the padding and makes the map match the original mask perfectly.\n    return full_prob_map[:h_orig, :w_orig]","metadata":{"id":"i9dU9NQHy_WN","trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:17:15.002707Z","iopub.execute_input":"2026-04-11T13:17:15.003338Z","iopub.status.idle":"2026-04-11T13:17:15.011156Z","shell.execute_reply.started":"2026-04-11T13:17:15.003308Z","shell.execute_reply":"2026-04-11T13:17:15.010417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nsave_dir = \"/kaggle/working/output_images\"\n\nos.makedirs(save_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T08:31:27.870278Z","iopub.execute_input":"2026-04-11T08:31:27.870649Z","iopub.status.idle":"2026-04-11T08:31:27.892188Z","shell.execute_reply.started":"2026-04-11T08:31:27.870610Z","shell.execute_reply":"2026-04-11T08:31:27.891029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sortedcontainers import SortedList\nfrom sortedcontainers import SortedDict\nfrom torch.utils.data import DataLoader\n\nimport sys\nimport cv2\nimg_path = \"/kaggle/working/output_images/result.png\"\n\nfolder_path = '/kaggle/input/datasets/mariaherrerot/aptos2019/test_images/test_images'\ni=1\nfor filename in os.listdir(folder_path):\n    if filename.endswith((\".png\", \".jpg\", \".jpeg\")):\n        img_path = os.path.join(folder_path, filename)\n\n        img = cv2.imread(img_path)\n\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # img=circular_crop(img)\n\n       \n        pre_img=get_clean_full_image(img)\n        \n        mean=np.mean(pre_img)\n        std=np.std(pre_img)\n        \n        norm_img=(pre_img-mean)/std\n        \n        \n        \n        \n        device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        model = MSMA_Net(in_channels=1, out_channels=1).to(device)\n        \n        # 3. Load Weights\n        checkpoint_path = '/kaggle/input/datasets/omarahmed014/msma-model/msma_model_BEST.pth'\n        model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n        \n        # 4. VERY IMPORTANT: Set to Evaluation Mode\n        \n        output=predict_full_res(model, torch.from_numpy(norm_img).unsqueeze(0).unsqueeze(0).to(device).float())\n\n        output=(output>0.5).astype('uint8')\n\n        output = (output * 255).astype(\"uint8\")\n\n        output=cv2.resize(output,(256,256))\n\n        \n        cv2.imwrite(f\"/kaggle/working/mask_{i}.png\", output)\n\n        \n\n        print(f\"image_{i} completed\")\n        i=i+1\n\n\n\n\n\n\n","metadata":{"id":"gPkh7Gip5ORc","outputId":"10584ae4-2fe9-4873-d301-a5507a5e4592","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from sortedcontainers import SortedList\n# from sortedcontainers import SortedDict\n# from torch.utils.data import DataLoader\n        \n# import sys\n# import cv2\n# # img_path = os.path.join(folder_path, filename)\n\n# img = cv2.imread('/kaggle/input/competitions/aptos2019-blindness-detection/test_images/003f0afdcd15.png')\n\n# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n# pre_img=get_clean_full_image(img)\n# mean=np.mean(pre_img)\n# std=np.std(pre_img)\n        \n# norm_img=(pre_img-mean)/std\n        \n        \n        \n        \n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# model = MSMA_Net(in_channels=1, out_channels=1).to(device)\n        \n# # 3. Load Weights\n# checkpoint_path = '/kaggle/input/datasets/omarahmed014/msma-model/msma_model_BEST.pth'\n# model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n        \n# # 4. VERY IMPORTANT: Set to Evaluation Mode\n        \n# output=predict_full_res(model, torch.from_numpy(norm_img).unsqueeze(0).unsqueeze(0).to(device).float())\n\n# output=(output>0.5).astype('uint8')\n\n\n# output = (output * 255).astype(\"uint8\")\n\n\n\n\n# print(output)\n\n\n\n\n\n\n    \n\n\n# plt.subplot(1,3,1)\n\n# plt.imshow(img)\n\n\n# plt.subplot(1,3,2)\n# plt.imshow(pre_img,cmap=\"gray\")\n\n\n# plt.subplot(1,3,3)\n# plt.imshow(output,cmap=\"gray\")\n\n\n# cv2.imwrite(f\"/kaggle/working/mask.png\",output)\n\n\n\n\n\n\n\n\n\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:29:50.001668Z","iopub.execute_input":"2026-04-11T13:29:50.002588Z","iopub.status.idle":"2026-04-11T13:29:51.338495Z","shell.execute_reply.started":"2026-04-11T13:29:50.002550Z","shell.execute_reply":"2026-04-11T13:29:51.337850Z"}},"outputs":[],"execution_count":null}]}