{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12299634,"sourceType":"datasetVersion","datasetId":7752320},{"sourceId":12299721,"sourceType":"datasetVersion","datasetId":7752387},{"sourceId":418031,"sourceType":"datasetVersion","datasetId":131128}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<a id=\"1\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b> APTOS 2019 Blindness Detection </b></div>\n<h2><center>Detect diabetic retinopathy to stop blindness before it's too late</center></h2>\n<center><img src=\"https://raw.githubusercontent.com/dimitreOliveira/MachineLearning/master/Kaggle/APTOS%202019%20Blindness%20Detection/aux_img.png\"></center>\n\nIn this notebook, we address the urgent issue of diabetic retinopathy detection, a major cause of blindness among working-aged adults. Aravind Eye Hospital in India aims to improve healthcare access in rural areas, where technicians capture retinal images, but trained doctors are needed to manually review them. While effective, this method is time-consuming and limits scalability.\n\nOur goal is to build a machine learning model that automates the detection of diabetic retinopathy, enabling faster and more accurate diagnoses. By leveraging state-of-the-art deep learning models, we aim to improve disease detection and expand this technology's potential to identify other conditions, like glaucoma and macular degeneration.\n\nModels Utilized\nWe combine the strengths of multiple cutting-edge models for effective image classification:\n\n**ResNetV2_50**: A deep convolutional network that specializes in feature extraction, allowing it to detect subtle features in retinal images with high precision.\n\n**DeiT_base_patch16** (Vision Transformer): A transformer-based architecture that excels at capturing global dependencies within images, making it ideal for complex medical images.\n\n**FastViT_s12**: A faster version of Vision Transformers, optimized for speed without sacrificing accuracy, ensuring quick inference times in real-time diagnostic settings.\n\n**swinv2_small_window16_256**: A Shifted Window Transformer capable of handling images at multiple scales, providing fine-grained detail for more accurate detection of retinopathy.\n","metadata":{}},{"cell_type":"markdown","source":"- <a href=\"#libraries\">1. Importing Required Libraries</a>\n- <a href=\"#EDA\">2. EDA</a>\n- <a href=\"#Data Preprocssing\">3. Data Preprecssing</a>\n- <a href=\"#Transformation\">4. Data Splitting and Transformation </a>\n- <a href=\"#Tuning\">5. Fine Tuning The Models </a>\n    - <a href=\"#function\">5.1. Train and validation function   </a> \n    - <a href=\"#Models\">5.2. Fine tuning  swinv2_small_window16_256 </a> \n- <a href=\"#Evalution\">6. Evalution</a>","metadata":{}},{"cell_type":"markdown","source":"<a id=\"libraries\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b> Importing Required Libraries </b></div>","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nimport pandas as pd\nimport os\nimport timm\nimport random\nimport time\nfrom collections import OrderedDict\nfrom torch.cuda import amp\nimport numpy as np\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.optimizer\nfrom torchvision import transforms as T\nimport matplotlib.pyplot as plt\nfrom torchvision.io import read_image\nimport cv2\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix\nfrom sklearn.metrics import f1_score\nimport seaborn as sns\nfrom tqdm import tqdm\nimport concurrent.futures\nfrom torch.utils.data import random_split\nfrom sklearn.metrics import classification_report, f1_score\n\nprint(torch.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:10:26.285081Z","iopub.execute_input":"2025-06-28T07:10:26.285789Z","iopub.status.idle":"2025-06-28T07:10:26.292517Z","shell.execute_reply.started":"2025-06-28T07:10:26.285754Z","shell.execute_reply":"2025-06-28T07:10:26.291574Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"EDA\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b> EDA </b></div>","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport os\n\n# Path to your dataset\ndata_dir = '/kaggle/input/diabetic-retinopathy-resized/resized_train/resized_train'\n\n# Let's create the DataFrame manually, assuming filenames are like: 10_left.jpeg, 10_right.jpeg, etc.\nimage_files = os.listdir(data_dir)\n\n# Assuming label is in filename, or a labels.csv is present (this is important!)\n# If label info is not in filename, load the label file instead:\n# labels_df = pd.read_csv('../input/diabetic-retinopathy-resized/trainLabels.csv')\n\n# Load train labels (for Diabetic Retinopathy Resized dataset)\ntrain = pd.read_csv('/kaggle/input/diabetic-retinopathy-resized/trainLabels.csv')  # Contains 'image' and 'level' columns\n\n# Append .jpeg to match image filenames\ntrain['image'] = train['image'].apply(lambda x: x + '.jpeg')\n\nprint('Number of train samples: ', train.shape[0])\ndisplay(train.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:10:31.063119Z","iopub.execute_input":"2025-06-28T07:10:31.063491Z","iopub.status.idle":"2025-06-28T07:10:31.116577Z","shell.execute_reply.started":"2025-06-28T07:10:31.063456Z","shell.execute_reply":"2025-06-28T07:10:31.115532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Number of train samples: ', train.shape[0])\ndisplay(train.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:10:36.235488Z","iopub.execute_input":"2025-06-28T07:10:36.23612Z","iopub.status.idle":"2025-06-28T07:10:36.244781Z","shell.execute_reply.started":"2025-06-28T07:10:36.236085Z","shell.execute_reply":"2025-06-28T07:10:36.243848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Corrected code\nf, ax = plt.subplots(figsize=(14, 8.7))\nax = sns.countplot(x=\"level\", data=train, palette=\"GnBu_d\")  # 'level' is the correct column\nsns.despine()\nplt.title(\"Distribution of Diabetic Retinopathy Severity Levels\", fontsize=16)\nplt.xlabel(\"DR Severity Level\", fontsize=14)\nplt.ylabel(\"Count\", fontsize=14)\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:12:40.555848Z","iopub.execute_input":"2025-06-28T07:12:40.556206Z","iopub.status.idle":"2025-06-28T07:12:40.830362Z","shell.execute_reply.started":"2025-06-28T07:12:40.556175Z","shell.execute_reply":"2025-06-28T07:12:40.829673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\nimport cv2\n\n# Set the style for the plot\nsns.set_style(\"white\")\n\n# Mapping class labels to their corresponding categories\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Plotting the first 15 images along with their labels\ncount = 1\nplt.figure(figsize=[20, 20])\n\n# Loop through first 15 images from 'image' column (not 'id_code')\nfor img_name in train['image'][:15]:\n    # Corrected path to read image\n    img_path = f\"/kaggle/input/diabetic-retinopathy-resized/resized_train/resized_train/10003_left.jpeg\"\n    img = cv2.imread(img_path)[..., [2, 1, 0]]  # Convert BGR to RGB\n\n    # Get the corresponding label\n    label = train[train['image'] == img_name]['level'].values[0]\n\n    # Plotting\n    plt.subplot(5, 5, count)\n    plt.imshow(img)\n    plt.title(f\"Image {count}: {level_to_category[label]}\")\n    plt.axis('off')\n    count += 1\n\n# Save and show the plot\nplt.savefig('/kaggle/working/imagebeforepreprocessing.png')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:17:12.614522Z","iopub.execute_input":"2025-06-28T07:17:12.614874Z","iopub.status.idle":"2025-06-28T07:17:16.875559Z","shell.execute_reply.started":"2025-06-28T07:17:12.614844Z","shell.execute_reply":"2025-06-28T07:17:16.874667Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"EDA\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b> Data Preprocessing </b></div>","metadata":{}},{"cell_type":"code","source":"# Function to crop the image based on grayscale threshold\ndef crop_image_from_gray(img, tol=7):\n    if img.ndim == 2:\n        mask = img > tol\n        return img[np.ix_(mask.any(1), mask.any(0))]\n    elif img.ndim == 3:\n        gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        mask = gray_img > tol\n        \n        check_shape = img[:,:,0][np.ix_(mask.any(1), mask.any(0))].shape[0]\n        if check_shape == 0:  # Image is too dark so that we crop out everything\n            return img  # Return original image\n        else:\n            img1 = img[:,:,0][np.ix_(mask.any(1), mask.any(0))]\n            img2 = img[:,:,1][np.ix_(mask.any(1), mask.any(0))]\n            img3 = img[:,:,2][np.ix_(mask.any(1), mask.any(0))]\n            img = np.stack([img1, img2, img3], axis=-1)\n        return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:18:47.198151Z","iopub.execute_input":"2025-06-28T07:18:47.19853Z","iopub.status.idle":"2025-06-28T07:18:47.205654Z","shell.execute_reply.started":"2025-06-28T07:18:47.198499Z","shell.execute_reply":"2025-06-28T07:18:47.204695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport pandas as pd\nfrom tqdm import tqdm\nimport concurrent.futures\n\n# Set input and output directories\ninput_dir = '/kaggle/input/diabetic-retinopathy-resized/resized_train/resized_train'\noutput_dir = '/kaggle/working/processed_images/'\nos.makedirs(output_dir, exist_ok=True)\n\n# Load the CSV\ncsv_path = '/kaggle/input/diabetic-retinopathy-resized/trainLabels.csv'\ndf = pd.read_csv(csv_path)\n\n# Update column name if needed\ndf['image'] = df['image'].apply(lambda x: x + '.jpeg')  # Add extension if not present\n\n# Create CLAHE object\nclahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n\n# Define cropping function if not already\ndef crop_image_from_gray(img, tol=7):\n    if img.ndim == 2:\n        mask = img > tol\n        return img[np.ix_(mask.any(1), mask.any(0))]\n    elif img.ndim == 3:\n        gray_img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        mask = gray_img > tol\n        if mask.any():\n            img1 = img[:, :, 0][np.ix_(mask.any(1), mask.any(0))]\n            img2 = img[:, :, 1][np.ix_(mask.any(1), mask.any(0))]\n            img3 = img[:, :, 2][np.ix_(mask.any(1), mask.any(0))]\n            return cv2.merge([img1, img2, img3])\n        else:\n            return img\n    return img\n\n# Processing function\ndef process_image(row):\n    sample_image_id = row['image']  # Correct column\n    sample_image_path = os.path.join(input_dir, sample_image_id)\n\n    if os.path.exists(sample_image_path):\n        image = cv2.imread(sample_image_path)\n        if image is None:\n            return\n\n        image_cropped = crop_image_from_gray(image)\n        image_resized = cv2.resize(image_cropped, (256, 256))\n        b, g, r = cv2.split(image_resized)\n\n        b_clahe = clahe.apply(b)\n        g_clahe = clahe.apply(g)\n        r_clahe = clahe.apply(r)\n\n        result_image = cv2.merge([b_clahe, g_clahe, r_clahe])\n        output_path = os.path.join(output_dir, sample_image_id)\n        cv2.imwrite(output_path, result_image)\n\n# Process with multithreading\nwith concurrent.futures.ThreadPoolExecutor(max_workers=8) as executor:\n    list(tqdm(executor.map(process_image, [row for _, row in df.iterrows()]), total=df.shape[0], desc=\"Processing images\", unit=\"image\"))\n\nprint(\"✅ Processing complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:22:29.585039Z","iopub.execute_input":"2025-06-28T07:22:29.585416Z","iopub.status.idle":"2025-06-28T07:27:11.860175Z","shell.execute_reply.started":"2025-06-28T07:22:29.585381Z","shell.execute_reply":"2025-06-28T07:27:11.859255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nimport os\n\n# Setting the style for the plot\nsns.set_style(\"white\")\n\n# Mapping class labels to their corresponding categories\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Plotting the first 15 processed images\ncount = 1\nplt.figure(figsize=[20, 20])\n\n# Loop through first 15 images from the DataFrame\nfor img_name in train['image'][:15]:  # Use 'image' column instead of 'id_code'\n    img_path = f\"/kaggle/working/processed_images/{img_name}\"  # Processed images are saved as .jpeg or .png\n\n    if not os.path.exists(img_path.replace('.jpeg', '.png')):\n        continue  # Skip if the processed image is not available\n\n    # Read and convert image\n    img = cv2.imread(img_path.replace('.jpeg', '.png'))\n    if img is None:\n        continue  # Skip if the image cannot be loaded\n\n    img = img[..., [2, 1, 0]]  # BGR to RGB\n\n    # Get label\n    label = train[train['image'] == img_name]['level'].values[0]\n\n    # Plot\n    plt.subplot(5, 5, count)\n    plt.imshow(img)\n    plt.title(f\"Image {count}: {level_to_category[label]}\")\n    plt.axis('off')\n    count += 1\n\n    if count > 15:\n        break\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:29:09.124154Z","iopub.execute_input":"2025-06-28T07:29:09.124574Z","iopub.status.idle":"2025-06-28T07:29:09.138086Z","shell.execute_reply.started":"2025-06-28T07:29:09.124541Z","shell.execute_reply":"2025-06-28T07:29:09.137289Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"Transformation\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b> Data Splitting and Transformation </b></div>","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/working/processed_images\"\nTRAIN_DIR = \"/kaggle/working/processed_images\"\nCSV_PATH = \"/kaggle/input/diabetic-retinopathy-resized/trainLabels.csv\"\nMODEL_PATH = \"./kaggle/working/\"\nLEARNING_RATE = 1e-4\nTRAIN_BATCH_SIZE = 32\nVALID_BATCH_SIZE = 32\nTRAIN_SPLIT = 0.8\nNUM_WORKERS = 2\nUSE_AMP = True\nEPOCHS=20","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:29:18.284734Z","iopub.execute_input":"2025-06-28T07:29:18.285048Z","iopub.status.idle":"2025-06-28T07:29:18.28966Z","shell.execute_reply.started":"2025-06-28T07:29:18.28502Z","shell.execute_reply":"2025-06-28T07:29:18.288802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RetinopathyDataset(Dataset):\n    def __init__(self, image_dir, csv_file, transforms=None):\n        self.data = pd.read_csv(csv_file)\n        self.transforms = transforms\n        self.image_dir = image_dir\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        # img_name = os.path.join('../input/aptos2019-blindness-detection/train_images',\n        #                         self.data.loc[idx, 'id_code'] + '.png')\n\n        img_name = os.path.join(self.image_dir, self.data.loc[idx, 'id_code'] + '.png')\n\n        tensor_image = read_image(img_name)\n        label = torch.tensor(self.data.loc[idx, 'diagnosis'], dtype=torch.long)\n\n        if self.transforms is not None:\n            tensor_image = self.transforms(tensor_image)\n\n        return (tensor_image, label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:29:46.211414Z","iopub.execute_input":"2025-06-28T07:29:46.212176Z","iopub.status.idle":"2025-06-28T07:29:46.21845Z","shell.execute_reply.started":"2025-06-28T07:29:46.21214Z","shell.execute_reply":"2025-06-28T07:29:46.21757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_trasforms_DeiT_base_patch16= T.Compose([\n    T.ConvertImageDtype(torch.float32),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\nfull_dataset = RetinopathyDataset(TRAIN_DIR, CSV_PATH, transforms=train_trasforms_DeiT_base_patch16)\n\ntrain_size = int(TRAIN_SPLIT * len(full_dataset))\ntest_size = len(full_dataset) - train_size\n\ntrain_dataset, val_dataset = torch.utils.data.random_split(full_dataset, [train_size, test_size])\n\ntrain_loader = DataLoader(full_dataset, batch_size=TRAIN_BATCH_SIZE,\n                          shuffle=False, num_workers=NUM_WORKERS, drop_last=True, pin_memory=False)\n\nval_loader = DataLoader(val_dataset, batch_size=VALID_BATCH_SIZE, shuffle=False,\n                        num_workers=NUM_WORKERS, drop_last=True, pin_memory=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:29:57.890795Z","iopub.execute_input":"2025-06-28T07:29:57.891676Z","iopub.status.idle":"2025-06-28T07:29:57.919187Z","shell.execute_reply.started":"2025-06-28T07:29:57.891624Z","shell.execute_reply":"2025-06-28T07:29:57.918266Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Helper Functions and Utilities for Training and Evaluation","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef accuracy(output, target):\n    \"\"\"\n    Computes the overall accuracy of predictions.\n\n    Args:\n        output (torch.Tensor): Model predictions with shape (batch_size, num_classes).\n        target (torch.Tensor): True labels with shape (batch_size,).\n\n    Returns:\n        float: Accuracy as a percentage.\n    \"\"\"\n    with torch.no_grad():\n        _, predicted = output.max(1)  # Get the class index with the highest score\n        correct = predicted.eq(target).sum().item()  # Count correct predictions\n        total = target.size(0)  # Total number of samples\n        accuracy = 100.0 * correct / total  # Compute accuracy percentage\n\n    return accuracy\n\n\ndef set_debug_apis(state: bool = False):\n    \"\"\"\n    Configures PyTorch debugging tools.\n\n    Args:\n        state (bool): If True, enables debugging tools for profiling and anomaly detection.\n    \"\"\"\n    torch.autograd.profiler.profile(enabled=state)\n    torch.autograd.profiler.emit_nvtx(enabled=state)\n    torch.autograd.set_detect_anomaly(mode=state)\n\n\ndef seed_everything(seed):\n    \"\"\"\n    Sets seeds for reproducibility in training.\n\n    Args:\n        seed (int): Seed value to ensure determinism.\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)  # Seed for hash-based operations\n    np.random.seed(seed)  # Seed for NumPy\n    torch.manual_seed(seed)  # Seed for PyTorch (CPU)\n    torch.cuda.manual_seed(seed)  # Seed for PyTorch (GPU)\n    torch.backends.cudnn.deterministic = True  # Make CuDNN deterministic\n    torch.backends.cudnn.benchmark = True  # Enable benchmark mode for CuDNN\n\n\ndef print_size_of_model(model):\n    \"\"\"\n    Calculates and prints the size of a PyTorch model.\n\n    Args:\n        model (torch.nn.Module): The model whose size is to be calculated.\n    \"\"\"\n    torch.save(model.state_dict(), \"temp.p\")  # Save model state\n    print(\"Size (MB):\", os.path.getsize(\"temp.p\") / 1e6)  # Convert bytes to MB\n    os.remove(\"temp.p\")  # Clean up temporary file\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:30:08.039883Z","iopub.execute_input":"2025-06-28T07:30:08.040248Z","iopub.status.idle":"2025-06-28T07:30:08.048115Z","shell.execute_reply.started":"2025-06-28T07:30:08.040192Z","shell.execute_reply":"2025-06-28T07:30:08.047205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(torch.cuda.is_available())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:30:14.547299Z","iopub.execute_input":"2025-06-28T07:30:14.547976Z","iopub.status.idle":"2025-06-28T07:30:14.552317Z","shell.execute_reply.started":"2025-06-28T07:30:14.547942Z","shell.execute_reply":"2025-06-28T07:30:14.551482Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"Tuning\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b> Fine Tuning The Models </b></div>","metadata":{}},{"cell_type":"markdown","source":"<a id=\"function\"></a>\n### <div style=\"text-align:center; border-radius:15px; padding:5px; color:white; margin:0; font-size:100%; font-family:Arial; background-color:#f39c12;\"><b>Training function and validations</b></div>\n","metadata":{}},{"cell_type":"code","source":"def train_step(model: torch.nn.Module, train_loader, criterion, device: str, optimizer, scheduler=None, num_batches: int = None, log_interval: int = 100, scaler=None):\n    \"\"\"\n    Performs one step of training with progress tracking using tqdm, and updates progress at each epoch.\n    \"\"\"\n    model = model.to(device)\n    model.train()\n\n    start_train_step = time.time()\n    metrics = OrderedDict()\n\n    total_loss = 0\n    correct_predictions = 0\n    total_samples = 0\n\n    # Initialize tqdm progress bar for the training loop (per epoch)\n    with tqdm(train_loader, desc=\"Training\", unit=\"batch\") as pbar:\n        for batch_idx, (inputs, target) in enumerate(pbar):\n            inputs = inputs.to(device)\n            target = target.to(device)\n\n            # Zero the parameter gradients\n            optimizer.zero_grad()\n\n            # Mixed precision training if scaler is provided\n            if scaler is not None:\n                with torch.cuda.amp.autocast():\n                    output = model(inputs)\n                    loss = criterion(output, target)\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                output = model(inputs)\n                loss = criterion(output, target)\n                loss.backward()\n                optimizer.step()\n\n            # Update the scheduler if it's provided\n            if scheduler is not None:\n                scheduler.step()\n\n            # Update metrics\n            total_loss += loss.item() * inputs.size(0)\n            _, predicted = output.max(1)\n            correct_predictions += predicted.eq(target).sum().item()\n            total_samples += inputs.size(0)\n\n            # Update progress bar (per batch)\n            pbar.set_postfix({\n                \"Loss\": f\"{total_loss / total_samples:.4f}\",\n                \"Accuracy\": f\"{100.0 * correct_predictions / total_samples:.2f}%\"\n            })\n\n            # Break if num_batches is set and we've reached the limit\n            if num_batches is not None and batch_idx + 1 >= num_batches:\n                break\n\n    # Final metrics for the epoch\n    end_train_step = time.time()\n    metrics[\"loss\"] = total_loss / total_samples\n    metrics[\"accuracy\"] = 100.0 * correct_predictions / total_samples\n\n    # Print the time taken for the train step and the summary of the metrics\n    print(f\"\\nEpoch Summary: Time taken for train step = {end_train_step - start_train_step:.2f} sec\")\n    print(f\"Training loss = {metrics['loss']:.4f}, Training accuracy = {metrics['accuracy']:.2f}%\")\n\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:30:36.290813Z","iopub.execute_input":"2025-06-28T07:30:36.291163Z","iopub.status.idle":"2025-06-28T07:30:36.301895Z","shell.execute_reply.started":"2025-06-28T07:30:36.29113Z","shell.execute_reply":"2025-06-28T07:30:36.300966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()  # No gradient calculation needed during validation\ndef val_step(model: torch.nn.Module, val_loader, criterion, device: str, num_batches=None, log_interval: int = 100):\n    \"\"\"\n    Performs one step of validation with progress tracking using tqdm.\n\n    Args:\n        model: A PyTorch CNN Model.\n        val_loader: DataLoader for the validation set.\n        criterion: Loss function to evaluate.\n        device: \"cuda\" or \"cpu\".\n        num_batches: (optional) Limit validation to a certain number of batches.\n        log_interval: (optional) Log after every specified number of batches.\n    \"\"\"\n    \n    model = model.to(device)\n    model.eval()  # Set the model to evaluation mode\n\n    start_val_step = time.time()  # Track the start time of the validation step\n    metrics = OrderedDict()\n\n    total_loss = 0\n    correct_predictions = 0\n    total_samples = 0\n\n    # Initialize tqdm progress bar for the validation loop (per epoch)\n    with tqdm(val_loader, desc=\"Validation\", unit=\"batch\") as pbar:\n        for batch_idx, (inputs, target) in enumerate(pbar):\n            inputs = inputs.to(device)\n            target = target.to(device)\n\n            # Forward pass (no gradient computation)\n            output = model(inputs)\n            loss = criterion(output, target)\n\n            # Update metrics\n            total_loss += loss.item() * inputs.size(0)\n            _, predicted = output.max(1)  # Get predictions\n            correct_predictions += predicted.eq(target).sum().item()  # Compare predictions with targets\n            total_samples += inputs.size(0)\n\n            # Update progress bar with the current loss and accuracy\n            pbar.set_postfix({\n                \"Loss\": f\"{total_loss / total_samples:.4f}\",\n                \"Accuracy\": f\"{100.0 * correct_predictions / total_samples:.2f}%\"\n            })\n\n            # Break if num_batches is specified and we've reached the limit\n            if num_batches is not None and batch_idx + 1 >= num_batches:\n                break\n\n    # Final metrics for the epoch\n    end_val_step = time.time()  # Track the end time of the validation step\n    metrics[\"loss\"] = total_loss / total_samples\n    metrics[\"accuracy\"] = 100.0 * correct_predictions / total_samples\n\n    # Print the time taken and the final validation metrics\n    print(f\"\\nValidation Summary: Time taken for validation step = {end_val_step - start_val_step:.2f} sec\")\n    print(f\"Validation loss = {metrics['loss']:.4f}, Validation accuracy = {metrics['accuracy']:.2f}%\")\n\n    return metrics\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:31:05.882276Z","iopub.execute_input":"2025-06-28T07:31:05.882616Z","iopub.status.idle":"2025-06-28T07:31:05.89137Z","shell.execute_reply.started":"2025-06-28T07:31:05.882586Z","shell.execute_reply":"2025-06-28T07:31:05.89047Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"Models\"></a>\n### <div style=\"text-align:center; border-radius:15px; padding:5px; color:white; margin:0; font-size:100%; font-family:Arial; background-color:#f39c12;\"><b> Fine Tuning  SWIN </b></div>\n","metadata":{}},{"cell_type":"code","source":"\nMODEL_NAME = \"swinv2_small_window16_256\"\nMODEL_SAVE=  \"/kaggle/working/swinv2.Csv\"\nseed_everything(42)\nset_debug_apis(False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:31:36.267699Z","iopub.execute_input":"2025-06-28T07:31:36.268395Z","iopub.status.idle":"2025-06-28T07:31:36.273769Z","shell.execute_reply.started":"2025-06-28T07:31:36.268363Z","shell.execute_reply":"2025-06-28T07:31:36.272816Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Handle imbalnced Classes ","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:31:40.703832Z","iopub.execute_input":"2025-06-28T07:31:40.70457Z","iopub.status.idle":"2025-06-28T07:31:40.708674Z","shell.execute_reply.started":"2025-06-28T07:31:40.704535Z","shell.execute_reply":"2025-06-28T07:31:40.707735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.utils.class_weight import compute_class_weight\nimport numpy as np\n\n# Extract the class labels (correct column name is 'level')\nclass_labels = train['level'].values\n\n# Compute class weights using sklearn\nunique_classes = np.unique(class_labels)\nclass_weights = compute_class_weight(class_weight='balanced', classes=unique_classes, y=class_labels)\n\n# Create a dictionary mapping each class to its weight\nclass_weight_dict = {class_label: weight for class_label, weight in zip(unique_classes, class_weights)}\n\nprint(\"Class Weights:\", class_weight_dict)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:33:02.030092Z","iopub.execute_input":"2025-06-28T07:33:02.030902Z","iopub.status.idle":"2025-06-28T07:33:02.04455Z","shell.execute_reply.started":"2025-06-28T07:33:02.030867Z","shell.execute_reply":"2025-06-28T07:33:02.043682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model= timm.create_model(MODEL_NAME, pretrained=True, num_classes=5)\nweights_tensor = torch.tensor(list(class_weight_dict.values())).float()\ncriterion= nn.CrossEntropyLoss(weight=weights_tensor.to(device))\noptimizer= optim.AdamW(model.parameters(), lr=LEARNING_RATE)\n\nif torch.cuda.is_available():\n    device = \"cuda\"\nelse:\n    device = \"cpu\"\n\nif USE_AMP:\n    from torch.cuda import amp\n    scaler = amp.GradScaler()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:33:20.730556Z","iopub.execute_input":"2025-06-28T07:33:20.731182Z","iopub.status.idle":"2025-06-28T07:33:21.948574Z","shell.execute_reply.started":"2025-06-28T07:33:20.731148Z","shell.execute_reply":"2025-06-28T07:33:21.94773Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset\n\nclass DRDataset(Dataset):\n    def __init__(self, dataframe, image_dir, transform=None):\n        self.data = dataframe.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        # Get image file name (without extension) and append `.jpeg`\n        image_name = self.data.loc[idx, 'image'] + '.jpeg'\n        image_path = os.path.join(self.image_dir, image_name)\n        \n        # Read image and convert BGR to RGB\n        image = cv2.imread(image_path)\n        if image is None:\n            raise FileNotFoundError(f\"Image not found: {image_path}\")\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        # Apply image transformation if provided\n        if self.transform:\n            image = self.transform(image=image)['image']\n\n        # Get corresponding label\n        label = self.data.loc[idx, 'level']\n        return image, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:42:09.848564Z","iopub.execute_input":"2025-06-28T07:42:09.848946Z","iopub.status.idle":"2025-06-28T07:42:09.857252Z","shell.execute_reply.started":"2025-06-28T07:42:09.848904Z","shell.execute_reply":"2025-06-28T07:42:09.856277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Load CSV\nimport pandas as pd\ndf = pd.read_csv('/kaggle/input/diabetic-retinopathy-resized/trainLabels.csv')\n\n# Split into train and val\ntrain_df, val_df = train_test_split(df, test_size=0.2, stratify=df['level'], random_state=42)\n\n# Image directory path\nimage_dir = '/kaggle/input/diabetic-retinopathy-resized/train_resized/'\n\n# Define transforms\ntrain_transform = A.Compose([\n    A.Resize(256, 256),\n    A.HorizontalFlip(),\n    A.VerticalFlip(),\n    A.Normalize(),\n    ToTensorV2(),\n])\n\nval_transform = A.Compose([\n    A.Resize(256, 256),\n    A.Normalize(),\n    ToTensorV2(),\n])\n\n# Create datasets\ntrain_dataset = DRDataset(train_df, image_dir, transform=train_transform)\nval_dataset = DRDataset(val_df, image_dir, transform=val_transform)\n\n# Create dataloaders\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:42:24.132941Z","iopub.execute_input":"2025-06-28T07:42:24.133667Z","iopub.status.idle":"2025-06-28T07:42:24.75214Z","shell.execute_reply.started":"2025-06-28T07:42:24.133631Z","shell.execute_reply":"2025-06-28T07:42:24.751271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport time\nimport torch\nimport pandas as pd\n\n# Initialize lists to store metrics\ntrain_loss = []\ntrain_accuracy = []\nval_loss = []\nval_accuracy = []\n\n# Record start time\nstart_time = time.time()\n\n# Total number of batches for training and validation\ntrain_batches = len(train_loader)\nval_batches = len(val_loader)\n                \n# Start epoch loop with tqdm tracking total iterations per epoch\nfor epoch in tqdm(range(EPOCHS), desc=\"Epochs\"):\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n\n    # Initialize tqdm for the training loop (track overall progress for the epoch)\n    with tqdm(total=train_batches, desc=f\"Training Epoch {epoch+1}\", unit=\"batch\") as pbar_train:\n        # Perform training step\n        train_metrics = train_step(model, train_loader, criterion, device, optimizer, scaler=scaler)\n        train_loss.append(train_metrics[\"loss\"])\n        train_accuracy.append(train_metrics[\"accuracy\"])\n\n        # Update progress bar after training\n        pbar_train.set_postfix({\n            \"Loss\": f\"{train_metrics['loss']:.4f}\",\n            \"Accuracy\": f\"{train_metrics['accuracy']:.2f}%\"\n        })\n        pbar_train.update(train_batches)  # Update progress bar to the total number of batches\n\n    # Initialize tqdm for the validation loop (track overall progress for the epoch)\n    with tqdm(total=val_batches, desc=f\"Validation Epoch {epoch+1}\", unit=\"batch\") as pbar_val:\n        # Perform validation step\n        val_metrics = val_step(model, val_loader, criterion, device)\n        val_loss.append(val_metrics[\"loss\"])\n        val_accuracy.append(val_metrics[\"accuracy\"])\n\n        # Update progress bar after validation\n        pbar_val.set_postfix({\n            \"Loss\": f\"{val_metrics['loss']:.4f}\",\n            \"Accuracy\": f\"{val_metrics['accuracy']:.2f}%\" \n        })\n        pbar_val.update(val_batches)  # Update progress bar to the total number of batches\n\n    # Print epoch summary (not per-batch)\n    print(f\"\\nEpoch {epoch+1} Summary:\")\n    print(f\"Training loss = {train_metrics['loss']:.4f}, Training accuracy = {train_metrics['accuracy']:.2f}%\")\n    print(f\"Validation loss = {val_metrics['loss']:.4f}, Validation accuracy = {val_metrics['accuracy']:.2f}%\")\n\n    # Save model checkpoint\n    checkpoint_path = f\"{MODEL_NAME}_epoch_{epoch+1}.pt\"\n    torch.save(model.state_dict(), checkpoint_path)\n\n# Record end time\nend_time = time.time()\n\n# Calculate total training time\ntotal_time = end_time - start_time\nprint(f\"\\nTotal training time: {total_time:.2f} seconds\")\n\n# Create a DataFrame to store the metrics\nmetrics_df = pd.DataFrame({\n    \"epoch\": range(1, EPOCHS + 1),\n    \"train_loss\": train_loss,\n    \"train_accuracy\": train_accuracy,\n    \"val_loss\": val_loss,\n    \"val_accuracy\": val_accuracy,\n})\n\n# Save the DataFrame to CSV\nmetrics_df.to_csv(MODEL_SAVE, index=False)\nprint(\"Training metrics saved to CSV.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:55:15.028732Z","iopub.execute_input":"2025-06-28T07:55:15.029104Z","iopub.status.idle":"2025-06-28T07:55:15.100546Z","shell.execute_reply.started":"2025-06-28T07:55:15.029071Z","shell.execute_reply":"2025-06-28T07:55:15.099382Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"Evalution\"></a>\n# <div style=\"text-align:center; border-radius:30px 30px; padding:7px; color:white; margin:0; font-size:110%; font-family:Pacifico; background-color:#5b81d4; overflow:hidden\"><b>Evalution </b></div>","metadata":{}},{"cell_type":"code","source":"metrics_path = f\"{MODEL_NAME}_metrics.csv\"  # Path to metrics CSV\nmetrics_df = pd.read_csv(MODEL_SAVE)\n\n# Set Seaborn style for better aesthetics\nsns.set_theme(style=\"whitegrid\")\n\n# Create a figure for visualizations\nplt.figure(figsize=(12, 8))\n\n# Plot training and validation loss\nplt.subplot(2, 1, 1)\nsns.lineplot(x='epoch', y='train_loss', data=metrics_df, label='Train Loss', color='blue', marker=\"o\")\nsns.lineplot(x='epoch', y='val_loss', data=metrics_df, label='Validation Loss', color='orange', marker=\"o\")\nplt.title(\"Training and Validation Loss per Epoch\", fontsize=14, fontweight='bold')\nplt.xlabel(\"Epoch\", fontsize=12)\nplt.ylabel(\"Loss\", fontsize=12)\nplt.legend()\nplt.grid(alpha=0.3)\n\n# Plot training and validation accuracy\nplt.subplot(2, 1, 2)\nsns.lineplot(x='epoch', y='train_accuracy', data=metrics_df, label='Train Accuracy', color='green', marker=\"o\")\nsns.lineplot(x='epoch', y='val_accuracy', data=metrics_df, label='Validation Accuracy', color='red', marker=\"o\")\nplt.title(\"Training and Validation Accuracy per Epoch\", fontsize=14, fontweight='bold')\nplt.xlabel(\"Epoch\", fontsize=12)\nplt.ylabel(\"Accuracy\", fontsize=12)\nplt.legend()\nplt.grid(alpha=0.3)\n\n# Adjust spacing between plots\nplt.tight_layout()\n\n# Save the visualization as a file (optional)\noutput_plot_path = f\"{MODEL_NAME}_training_visualization.png\"\nplt.savefig(output_plot_path, dpi=300)\nprint(f\"Visualization saved to {output_plot_path}\")\n\n# Show the plots\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:21.233544Z","iopub.execute_input":"2025-06-28T06:56:21.233814Z","iopub.status.idle":"2025-06-28T06:56:22.88493Z","shell.execute_reply.started":"2025-06-28T06:56:21.233786Z","shell.execute_reply":"2025-06-28T06:56:22.884095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load your model from checkpoint\ndef load_model_from_checkpoint(checkpoint_path, model_class, device):\n    model.load_state_dict(torch.load(checkpoint_path, map_location=device))\n    model.to(device)\n    model.eval()  # Set to evaluation mode\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:22.886251Z","iopub.execute_input":"2025-06-28T06:56:22.88659Z","iopub.status.idle":"2025-06-28T06:56:22.891821Z","shell.execute_reply.started":"2025-06-28T06:56:22.886553Z","shell.execute_reply":"2025-06-28T06:56:22.890981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Predict on validation data\ndef predict_on_validation(model, dataloader, device):\n    model.eval()  # Set the model to evaluation mode\n    y_true = []\n    y_pred = []\n\n    with torch.no_grad():  # No need to compute gradients for validation\n        for inputs, labels in val_loader:  # Assuming dataloader is a dictionary with 'val' key\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n            outputs = model(inputs)\n            _, preds = torch.max(outputs, 1)\n\n            # Collect true and predicted labels\n            y_true.extend(labels.cpu().numpy())\n            y_pred.extend(preds.cpu().numpy())\n\n    return np.array(y_true), np.array(y_pred)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:22.89284Z","iopub.execute_input":"2025-06-28T06:56:22.893125Z","iopub.status.idle":"2025-06-28T06:56:22.902993Z","shell.execute_reply.started":"2025-06-28T06:56:22.893098Z","shell.execute_reply":"2025-06-28T06:56:22.902348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define paths and device\ncheckpoint_path = \"/kaggle/working/swinv2_small_window16_256_epoch_19.pt\"\nmodel_class=5\n# Initialize and load the model (replace YourModel with your actual model class)\nmodel = load_model_from_checkpoint(checkpoint_path, model_class, device)\n\n# Assuming `dataloaders` is a dictionary with 'val' key containing the validation dataloader\n# Predict on validation set\ny_true, y_pred = predict_on_validation(model, val_loader, device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:22.90401Z","iopub.execute_input":"2025-06-28T06:56:22.904291Z","iopub.status.idle":"2025-06-28T06:56:29.080361Z","shell.execute_reply.started":"2025-06-28T06:56:22.904249Z","shell.execute_reply":"2025-06-28T06:56:29.079118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate Classification Report\nreport = classification_report(y_true, y_pred, digits=2)\nprint(\"\\nClassification Report:\\n\", report)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.081946Z","iopub.execute_input":"2025-06-28T06:56:29.082278Z","iopub.status.idle":"2025-06-28T06:56:29.095466Z","shell.execute_reply.started":"2025-06-28T06:56:29.082242Z","shell.execute_reply":"2025-06-28T06:56:29.094679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Compute confusion matrix\ncm = confusion_matrix(y_true, y_pred)\n\n# Plot confusion matrix using seaborn\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n            xticklabels=np.arange(cm.shape[0]), yticklabels=np.arange(cm.shape[0]))\nplt.xlabel('Predicted')\nplt.ylabel('True')\nplt.title('Confusion Matrix')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.09663Z","iopub.execute_input":"2025-06-28T06:56:29.096947Z","iopub.status.idle":"2025-06-28T06:56:29.449959Z","shell.execute_reply.started":"2025-06-28T06:56:29.09691Z","shell.execute_reply":"2025-06-28T06:56:29.449089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the level-to-category mapping\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Calculate errors per class (misclassifications)\nclass_errors = {}\nnum_classes = cm.shape[0]\n\nfor i in range(num_classes):\n    total = np.sum(cm[i, :])  # Total number of instances of class i\n    incorrect = total - cm[i, i]  # Misclassified instances of class i\n    error_rate = incorrect / total if total != 0 else 0\n    class_errors[level_to_category[i]] = round(error_rate, 4)  # Round error_rate and map to class name\n\n# Print class errors\nfor class_name, error_rate in class_errors.items():\n    print(f\"Class {class_name}: Error Rate = {error_rate}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.451326Z","iopub.execute_input":"2025-06-28T06:56:29.451705Z","iopub.status.idle":"2025-06-28T06:56:29.458831Z","shell.execute_reply.started":"2025-06-28T06:56:29.451667Z","shell.execute_reply":"2025-06-28T06:56:29.457852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a DataFrame for the error rates\nerror_data = pd.DataFrame(list(class_errors.items()), columns=['Class', 'Error Rate'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.459815Z","iopub.execute_input":"2025-06-28T06:56:29.460053Z","iopub.status.idle":"2025-06-28T06:56:29.472298Z","shell.execute_reply.started":"2025-06-28T06:56:29.460029Z","shell.execute_reply":"2025-06-28T06:56:29.471572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot only the table, fitting the entire figure\nfig, ax = plt.subplots(figsize=(8, 4))  # Adjust size as needed\nax.axis('off')  # Turn off axis for the table\n\n# Create and display the table, making it fill the figure\ntable = ax.table(cellText=error_data.values, colLabels=error_data.columns, cellLoc='center', loc='center')\ntable.auto_set_font_size(False)\ntable.set_fontsize(12)\ntable.auto_set_column_width(col=list(range(len(error_data.columns))))\ntable.scale(20, 5)  # Scale table size (adjust the values as needed)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.473332Z","iopub.execute_input":"2025-06-28T06:56:29.473985Z","iopub.status.idle":"2025-06-28T06:56:29.825676Z","shell.execute_reply.started":"2025-06-28T06:56:29.473958Z","shell.execute_reply":"2025-06-28T06:56:29.824506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_image(image_path):\n    # Load the image\n    image = cv2.imread(image_path)\n    \n    # Crop the image based on gray threshold (assuming crop_image_from_gray is defined elsewhere)\n    image_cropped = crop_image_from_gray(image)  \n    \n    # Resize the image to 256x256\n    image_resized = cv2.resize(image_cropped, (256, 256))\n    \n    # Convert the image from BGR to RGB (OpenCV loads images as BGR)\n    image_rgb = cv2.cvtColor(image_resized, cv2.COLOR_BGR2RGB)\n    \n    # Split the image into its channels (BGR format)\n    blue, green, red = cv2.split(image_rgb)\n    \n    # Apply CLAHE to all three channels\n    blue_clahe = clahe.apply(blue)\n    green_clahe = clahe.apply(green)\n    red_clahe = clahe.apply(red)\n    \n    # Merge the CLAHE-enhanced channels back together\n    result_image = cv2.merge([blue_clahe, green_clahe, red_clahe])\n    \n    # Convert image to a PyTorch tensor\n    result_image = torch.tensor(result_image, dtype=torch.float32).permute(2, 0, 1) / 255.0  # CHW format, normalize between [0, 1]\n    \n    # Apply the transformations\n    result_image = train_transforms_DeiT_base_patch16(result_image)\n    \n    return result_image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.827108Z","iopub.execute_input":"2025-06-28T06:56:29.82764Z","iopub.status.idle":"2025-06-28T06:56:29.841036Z","shell.execute_reply.started":"2025-06-28T06:56:29.827584Z","shell.execute_reply":"2025-06-28T06:56:29.839835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\nimport cv2\nfrom tqdm import tqdm\n\n# Initialize lists to store image names and predictions\nimage_names = []\npredictions = []\n\n# Define the preprocess function (if needed)\ndef process_image(image_path):\n    # Load image\n    image = cv2.imread(image_path)\n    \n    # Resize the image to 256x256\n    image_resized = cv2.resize(image, (256, 256))  # Resize image to 256x256\n    \n    # Convert to RGB (from BGR)\n    image_resized = cv2.cvtColor(image_resized, cv2.COLOR_BGR2RGB)\n    \n    # Normalize the image and convert to tensor\n    image = image_resized.transpose((2, 0, 1))  # Convert HWC to CHW format\n    image = torch.tensor(image, dtype=torch.float32) / 255.0  # Normalize image\n    \n    return image.unsqueeze(0)  # Add batch dimension\n\n# Load the test CSV\ndf_test = pd.read_csv(\"/kaggle/input/aptos2019-blindness-detection/test.csv\")\n\n# Loop through the test images and make predictions\nfor image_name in tqdm(df_test['id_code'], desc=\"Predicting on test images\"):\n    image_path = os.path.join('/kaggle/input/aptos2019-blindness-detection/test_images', f\"{image_name}.png\")\n    \n    # Preprocess the image\n    image = process_image(image_path).to(device)\n\n    # Make prediction using the model\n    with torch.no_grad():\n        output = model(image)\n        _, predicted_label = torch.max(output, 1)\n\n    # Store the image name and predicted label\n    image_names.append(image_name)\n    predictions.append(predicted_label.item())  # Get the label as a scalar\n\n# Create DataFrame with image names and predicted labels\nresults_df = pd.DataFrame({\n    'id_code': image_names,\n    'diagnosis': predictions\n})\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:56:29.842332Z","iopub.execute_input":"2025-06-28T06:56:29.842765Z","iopub.status.idle":"2025-06-28T06:59:01.847566Z","shell.execute_reply.started":"2025-06-28T06:56:29.842713Z","shell.execute_reply":"2025-06-28T06:59:01.846615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save the DataFrame to a CSV file\nresults_df.to_csv('/kaggle/working/submission.csv', index=False)\n# Display the first few rows of the DataFrame\nprint(results_df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:59:01.849082Z","iopub.execute_input":"2025-06-28T06:59:01.849462Z","iopub.status.idle":"2025-06-28T06:59:01.859325Z","shell.execute_reply.started":"2025-06-28T06:59:01.849423Z","shell.execute_reply":"2025-06-28T06:59:01.858476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), f\"/kaggle/working/swinv2_small_window16_256_epoch_20.pt\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T06:59:01.860528Z","iopub.execute_input":"2025-06-28T06:59:01.860872Z","iopub.status.idle":"2025-06-28T06:59:02.360531Z","shell.execute_reply.started":"2025-06-28T06:59:01.860835Z","shell.execute_reply":"2025-06-28T06:59:02.359287Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# End","metadata":{}},{"cell_type":"code","source":"import cv2\nimport torch\nimport numpy as np\nimport os\n\n# Category mapping\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Function to preprocess the image\ndef preprocess_image(image_path):\n    # Load the image\n    image = cv2.imread(image_path)\n    if image is None:\n        raise ValueError(f\"Could not load image: {image_path}\")\n\n    # Convert to RGB\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n    # Resize\n    image = cv2.resize(image, (256, 256))\n\n    # CLAHE (as in training)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    r, g, b = cv2.split(image)\n    r = clahe.apply(r)\n    g = clahe.apply(g)\n    b = clahe.apply(b)\n    image = cv2.merge([r, g, b])\n\n    # Convert to tensor\n    image = image.transpose((2, 0, 1))  # HWC -> CHW\n    image = torch.tensor(image, dtype=torch.float32) / 255.0\n\n    # Apply same normalization\n    transform = T.Compose([\n        T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n    image = transform(image)\n    return image.unsqueeze(0)  # Add batch dimension\n\n# Load the model from checkpoint\ndef load_model(model_path, model_name=\"swinv2_small_window16_256\"):\n    model = timm.create_model(model_name, pretrained=False, num_classes=5)\n    model.load_state_dict(torch.load(model_path, map_location=device))\n    model.to(device)\n    model.eval()\n    return model\n\n# Main prediction loop\ndef predict_single_image():\n    image_path = input(\"🔍 Enter the path of the image to predict (e.g., /path/image.png): \").strip()\n    \n    if not os.path.exists(image_path):\n        print(\"❌ Image path does not exist. Please try again.\")\n        return\n\n    image = preprocess_image(image_path).to(device)\n\n    with torch.no_grad():\n        output = model(image)\n        pred_class = torch.argmax(output, dim=1).item()\n        print(f\"\\n✅ Predicted Severity Level: {pred_class} - {level_to_category[pred_class]}\\n\")\n\n# Load model only once\ncheckpoint_path = \"/kaggle/working/swinv2_small_window16_256_epoch_19.pt\"\nmodel = load_model(checkpoint_path)\n\n# Ask user to predict multiple images\nwhile True:\n    predict_single_image()\n    again = input(\"🔁 Do you want to predict another image? (y/n): \").strip().lower()\n    if again != 'y':\n        print(\"🛑 Prediction session ended.\")\n        break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-28T07:00:34.172494Z","iopub.execute_input":"2025-06-28T07:00:34.173343Z","iopub.status.idle":"2025-06-28T07:08:43.565956Z","shell.execute_reply.started":"2025-06-28T07:00:34.173307Z","shell.execute_reply":"2025-06-28T07:08:43.564752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}