{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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,"sourceType":"competition"},{"sourceId":10247719,"sourceType":"datasetVersion","datasetId":6338102}],"dockerImageVersionId":31011,"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=\"#DataPreprocssing\">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":"\n!pip install torch_optimizer","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:29:52.844481Z","iopub.execute_input":"2025-04-25T12:29:52.845017Z","iopub.status.idle":"2025-04-25T12:31:11.326750Z","shell.execute_reply.started":"2025-04-25T12:29:52.844995Z","shell.execute_reply":"2025-04-25T12:31:11.326040Z"}},"outputs":[],"execution_count":null},{"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 datasets, 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\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import accuracy_score\nimport torchvision.models as models\nfrom torch.utils.data import WeightedRandomSampler\nimport torch_optimizer as optim\n\n\n# Load a pre-trained MobileNetV2 model from torchvision\nprint(torch.__version__)","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:11.328202Z","iopub.execute_input":"2025-04-25T12:31:11.328460Z","iopub.status.idle":"2025-04-25T12:31:22.861580Z","shell.execute_reply.started":"2025-04-25T12:31:11.328438Z","shell.execute_reply":"2025-04-25T12:31:22.860729Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\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","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:22.862613Z","iopub.execute_input":"2025-04-25T12:31:22.863027Z","iopub.status.idle":"2025-04-25T12:31:22.867682Z","shell.execute_reply.started":"2025-04-25T12:31:22.863007Z","shell.execute_reply":"2025-04-25T12:31:22.866905Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(42)","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:33.386933Z","iopub.execute_input":"2025-04-25T12:31:33.387631Z","iopub.status.idle":"2025-04-25T12:31:33.395992Z","shell.execute_reply.started":"2025-04-25T12:31:33.387607Z","shell.execute_reply":"2025-04-25T12:31:33.395487Z"},"trusted":true},"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":"train= pd.read_csv('../input/aptos2019-blindness-detection/train.csv')\ntest= pd.read_csv('../input/aptos2019-blindness-detection/test.csv')","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:35.932766Z","iopub.execute_input":"2025-04-25T12:31:35.933038Z","iopub.status.idle":"2025-04-25T12:31:35.970763Z","shell.execute_reply.started":"2025-04-25T12:31:35.933017Z","shell.execute_reply":"2025-04-25T12:31:35.970204Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Number of train samples: ', train.shape[0])\nprint('Number of test samples: ', test.shape[0])\ndisplay(train.head())","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:37.716694Z","iopub.execute_input":"2025-04-25T12:31:37.716963Z","iopub.status.idle":"2025-04-25T12:31:37.739685Z","shell.execute_reply.started":"2025-04-25T12:31:37.716942Z","shell.execute_reply":"2025-04-25T12:31:37.739093Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"f, ax = plt.subplots(figsize=(14, 8.7))\nax = sns.countplot(x=\"diagnosis\", data=train, palette=\"GnBu_d\")\nsns.despine()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:38.855916Z","iopub.execute_input":"2025-04-25T12:31:38.856300Z","iopub.status.idle":"2025-04-25T12:31:39.070233Z","shell.execute_reply.started":"2025-04-25T12:31:38.856266Z","shell.execute_reply":"2025-04-25T12:31:39.069607Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\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 images along with their labels\ncount = 1\nplt.figure(figsize=[20, 20])\n\nfor img_name in train['id_code'][:15]:  # Assuming 'train' contains the dataset\n    img = cv2.imread(f\"../input/aptos2019-blindness-detection/train_images/{img_name}.png\")[..., [2, 1, 0]]  # Reading the image\n    \n    # Getting the label (class) for the image\n    label = train[train['id_code'] == img_name]['diagnosis'].values[0]  # Assuming 'diagnosis' is the label column\n    \n    # Setting up the subplot with image and label\n    plt.subplot(5, 5, count)\n    plt.imshow(img)\n    plt.title(f\"Image {count}: {level_to_category[label]}\")  # Display the class label\n    count += 1\n# Display the plot\nplt.savefig('/kaggle/working/imagebeforepreprecssing.png')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2025-04-25T12:31:40.240226Z","iopub.execute_input":"2025-04-25T12:31:40.240827Z","iopub.status.idle":"2025-04-25T12:31:58.806705Z","shell.execute_reply.started":"2025-04-25T12:31:40.240794Z","shell.execute_reply":"2025-04-25T12:31:58.805780Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"DataPreprocssing\"></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":{"execution":{"iopub.status.busy":"2025-04-25T12:32:08.826066Z","iopub.execute_input":"2025-04-25T12:32:08.827164Z","iopub.status.idle":"2025-04-25T12:32:08.834663Z","shell.execute_reply.started":"2025-04-25T12:32:08.827131Z","shell.execute_reply":"2025-04-25T12:32:08.833753Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set input and output directories\ninput_dir = '/kaggle/input/aptos2019-blindness-detection/train_images/'\noutput_dir = '/kaggle/working/processed_images/'\n\n# Ensure the output directory exists\nos.makedirs(output_dir, exist_ok=True)\n\n# Load the CSV containing image names and labels\ncsv_path = '/kaggle/input/aptos2019-blindness-detection/train.csv'\ndf = pd.read_csv(csv_path)\n\n# Create a CLAHE object\nclahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n\n# Function to process a single image\ndef process_image(row):\n    sample_image_id = row['id_code']  # Get the image ID\n    sample_image_file = sample_image_id + '.png'  # Assuming the image files have .png extension\n    sample_image_path = os.path.join(input_dir, sample_image_file)\n\n    if os.path.exists(sample_image_path):\n        # Load the image\n        image = cv2.imread(sample_image_path)\n        \n        # Crop the image based on gray threshold\n        image_cropped = crop_image_from_gray(image)\n        \n        # Resize the image to 256x256\n        image_resized = cv2.resize(image_cropped, (224, 224))\n        \n        # Split the image into its channels (BGR format)\n        blue, green, red = cv2.split(image_resized)\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        # Save the processed image to the output directory\n        output_path = os.path.join(output_dir, sample_image_file)\n        cv2.imwrite(output_path, result_image)\n\n# Using ThreadPoolExecutor to process images in parallel\nwith concurrent.futures.ThreadPoolExecutor(max_workers=8) as executor:\n    # Using tqdm to show progress bar while processing images\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 for all images.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:32:11.822205Z","iopub.execute_input":"2025-04-25T12:32:11.822519Z","iopub.status.idle":"2025-04-25T12:36:21.943539Z","shell.execute_reply.started":"2025-04-25T12:32:11.822499Z","shell.execute_reply":"2025-04-25T12:36:21.942823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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 images along with their labels\ncount = 1\nplt.figure(figsize=[20, 20])\n\nfor img_name in train['id_code'][:15]:  # Assuming 'train' contains the dataset\n    img = cv2.imread(f\"/kaggle/working/processed_images/{img_name}.png\")[..., [2, 1, 0]]  # Reading the image\n    \n    # Getting the label (class) for the image\n    label = train[train['id_code'] == img_name]['diagnosis'].values[0]  # Assuming 'diagnosis' is the label column\n    \n    # Setting up the subplot with image and label\n    plt.subplot(5, 5, count)\n    plt.imshow(img)\n    plt.title(f\"Image {count}: {level_to_category[label]}\")  # Display the class label\n    count += 1\n\n# Display the plot\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:36:24.114651Z","iopub.execute_input":"2025-04-25T12:36:24.115384Z","iopub.status.idle":"2025-04-25T12:36:27.207446Z","shell.execute_reply.started":"2025-04-25T12:36:24.115334Z","shell.execute_reply":"2025-04-25T12:36:27.206415Z"}},"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/input/aptos-224\"\nTRAIN_DIR = \"/kaggle/input/aptos-224/processed_images\"\nCSV_PATH = \"/kaggle/input/aptos2019-blindness-detection/train.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-04-25T12:36:46.547917Z","iopub.execute_input":"2025-04-25T12:36:46.548632Z","iopub.status.idle":"2025-04-25T12:36:46.552739Z","shell.execute_reply.started":"2025-04-25T12:36:46.548608Z","shell.execute_reply":"2025-04-25T12:36:46.552047Z"}},"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)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:36:47.795955Z","iopub.execute_input":"2025-04-25T12:36:47.796466Z","iopub.status.idle":"2025-04-25T12:36:47.801914Z","shell.execute_reply.started":"2025-04-25T12:36:47.796441Z","shell.execute_reply":"2025-04-25T12:36:47.801113Z"}},"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                         num_workers=NUM_WORKERS, drop_last=True, pin_memory=False)\n\nval_loader = DataLoader(val_dataset, batch_size=VALID_BATCH_SIZE, \n                        num_workers=NUM_WORKERS, drop_last=True, pin_memory=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:36:48.490710Z","iopub.execute_input":"2025-04-25T12:36:48.491476Z","iopub.status.idle":"2025-04-25T12:36:48.521842Z","shell.execute_reply.started":"2025-04-25T12:36:48.491450Z","shell.execute_reply":"2025-04-25T12:36:48.521107Z"}},"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:36:50.170793Z","iopub.execute_input":"2025-04-25T12:36:50.171553Z","iopub.status.idle":"2025-04-25T12:36:50.176835Z","shell.execute_reply.started":"2025-04-25T12:36:50.171521Z","shell.execute_reply":"2025-04-25T12:36:50.176046Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(torch.cuda.is_available())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:36:51.253686Z","iopub.execute_input":"2025-04-25T12:36:51.254365Z","iopub.status.idle":"2025-04-25T12:36:51.326309Z","shell.execute_reply.started":"2025-04-25T12:36:51.254318Z","shell.execute_reply":"2025-04-25T12:36:51.325592Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n<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-04-25T12:37:07.561349Z","iopub.execute_input":"2025-04-25T12:37:07.562058Z","iopub.status.idle":"2025-04-25T12:37:07.570025Z","shell.execute_reply.started":"2025-04-25T12:37:07.562036Z","shell.execute_reply":"2025-04-25T12:37:07.569279Z"}},"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-04-25T12:37:08.486874Z","iopub.execute_input":"2025-04-25T12:37:08.487467Z","iopub.status.idle":"2025-04-25T12:37:08.494656Z","shell.execute_reply.started":"2025-04-25T12:37:08.487444Z","shell.execute_reply":"2025-04-25T12:37:08.493936Z"}},"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":"# Define the models\nModel_P= \"pit_b_224\"\nMODEL_NAME='Mobile_Pit'\nset_debug_apis(False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:37:18.370441Z","iopub.execute_input":"2025-04-25T12:37:18.370962Z","iopub.status.idle":"2025-04-25T12:37:18.374539Z","shell.execute_reply.started":"2025-04-25T12:37:18.370941Z","shell.execute_reply":"2025-04-25T12:37:18.373810Z"}},"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-04-25T12:37:21.084822Z","iopub.execute_input":"2025-04-25T12:37:21.085387Z","iopub.status.idle":"2025-04-25T12:37:21.088830Z","shell.execute_reply.started":"2025-04-25T12:37:21.085364Z","shell.execute_reply":"2025-04-25T12:37:21.088127Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  (but you want MobileNet here)\nModel_Mobile = models.mobilenet_v2(pretrained=True)  # Use MobileNetV2 from torchvision\nnum_features_mobile = Model_Mobile.classifier[1].in_features  # Access the final layer's in_features\nModel_Mobile.classifier = nn.Identity()  # Replace the fully connected layer with Identity to get features\n\n# PiT (DeiT-style) Model Setup using timm\nModel_pit = timm.create_model(Model_P, pretrained=True, num_classes=5)\nnum_features_pit = Model_pit.head.in_features  # Access the head layer's in_features\nModel_pit.head = nn.Identity()  # Replace the head with Identity to get features\n\n# Combined Model Class\nclass CombinedModel(nn.Module):\n    def __init__(self, Model_Mobile, Model_pit, num_classes):\n        super(CombinedModel, self).__init__()\n        self.Model_Mobile = Model_Mobile\n        self.Model_pit = Model_pit\n        # Linear layer to classify after concatenating features\n        self.fc = nn.Linear(num_features_mobile + num_features_pit, num_classes)\n\n    def forward(self, x):\n        # Extract features from MobileNetV2\n        Model_Mobile_features = self.Model_Mobile(x)\n        # Extract features from PiT\n        Model_pit_features = self.Model_pit(x)\n\n        # Ensure outputs are 2D (batch_size, num_features)\n        if len(Model_Mobile_features.shape) > 2:\n            Model_Mobile_features = Model_Mobile_features.flatten(1)  # Flatten spatial dimensions\n        if len(Model_pit_features.shape) > 2:\n            Model_pit_features = Model_pit_features.flatten(1)\n\n        # Concatenate features along the feature axis (dim=1)\n        combined_features = torch.cat((Model_Mobile_features, Model_pit_features), dim=1)\n\n        # Pass through the final linear layer\n        output = self.fc(combined_features)\n        return output\n\n# Instantiate the combined model\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel_co = CombinedModel(Model_Mobile, Model_pit, num_classes=5).to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:37:22.275608Z","iopub.execute_input":"2025-04-25T12:37:22.275877Z","iopub.status.idle":"2025-04-25T12:37:26.525074Z","shell.execute_reply.started":"2025-04-25T12:37:22.275858Z","shell.execute_reply":"2025-04-25T12:37:26.524529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract the class labels (diagnosis column)\nclass_labels = train['diagnosis'].values\n\n# Compute class weights using sklearn\nunique_classes = np.unique(class_labels)\nclass_weights = compute_class_weight('balanced', classes=unique_classes, y=class_labels)\n\n# Create a dictionary that maps class labels to their corresponding weights\nclass_weight_dict = {class_label: weight for class_label, weight in zip(unique_classes, class_weights)}\n\nprint(\"Class Weights:\", class_weight_dict)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:37:28.116163Z","iopub.execute_input":"2025-04-25T12:37:28.116737Z","iopub.status.idle":"2025-04-25T12:37:28.123220Z","shell.execute_reply.started":"2025-04-25T12:37:28.116717Z","shell.execute_reply":"2025-04-25T12:37:28.122585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch_optimizer\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:37:33.489306Z","iopub.execute_input":"2025-04-25T12:37:33.489901Z","iopub.status.idle":"2025-04-25T12:37:36.779805Z","shell.execute_reply.started":"2025-04-25T12:37:33.489881Z","shell.execute_reply":"2025-04-25T12:37:36.779050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights_tensor = torch.tensor(list(class_weight_dict.values())).float()\ncriterion= nn.CrossEntropyLoss(weight=weights_tensor.to(device))\nimport torch.optim as optim  # Use torch.optim, not torch_op\n\nimport torch_optimizer as optim\noptimizer=  optim.Lamb(model_co.parameters(), lr=LEARNING_RATE, weight_decay=1e-4)\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-04-25T12:37:38.680231Z","iopub.execute_input":"2025-04-25T12:37:38.680765Z","iopub.status.idle":"2025-04-25T12:37:38.693931Z","shell.execute_reply.started":"2025-04-25T12:37:38.680728Z","shell.execute_reply":"2025-04-25T12:37:38.693084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf /kaggle/working/*\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:37:40.685600Z","iopub.execute_input":"2025-04-25T12:37:40.685868Z","iopub.status.idle":"2025-04-25T12:37:40.994851Z","shell.execute_reply.started":"2025-04-25T12:37:40.685848Z","shell.execute_reply":"2025-04-25T12:37:40.993827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\nimport time\n\n# Initialize loss and accuracy history before starting the loop\ntrain_loss = []\ntrain_accuracy = []\nval_loss = []\nval_accuracy = []\n\n# Start epoch loop with tqdm tracking total iterations per epoch\ntrain_batches = len(train_loader)\nval_batches = len(val_loader)\n\n# Start timer for overall training time\nstart_time = time.time()\n\nfor epoch in tqdm(range(EPOCHS), desc=\"Epochs\"):\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n\n    # Training loop\n    with tqdm(total=train_batches, desc=f\"Training Epoch {epoch+1}\", unit=\"batch\") as pbar_train:\n        train_metrics = train_step(model_co, train_loader, criterion, device, optimizer, scaler=scaler)\n        train_loss.append(train_metrics['loss'])\n        train_accuracy.append(train_metrics['accuracy'])\n        pbar_train.set_postfix({\"Loss\": f\"{train_metrics['loss']:.4f}\", \"Accuracy\": f\"{train_metrics['accuracy']:.2f}%\"})\n        pbar_train.update()\n\n    # Validation loop\n    with tqdm(total=val_batches, desc=f\"Validation Epoch {epoch+1}\", unit=\"batch\") as pbar_val:\n        val_metrics = val_step(model_co, val_loader, criterion, device)\n        val_loss.append(val_metrics['loss'])\n        val_accuracy.append(val_metrics['accuracy'])\n        pbar_val.set_postfix({\"Loss\": f\"{val_metrics['loss']:.4f}\", \"Accuracy\": f\"{val_metrics['accuracy']:.2f}%\"})\n        pbar_val.update()\n\n    # Print epoch summary\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 only the last checkpoint (overwrite previous one)\n    checkpoint_path = f\"{MODEL_NAME}_last_checkpoint.pt\"\n    torch.save({\n        \"epoch\": epoch + 1,\n        \"model_state_dict\": model_co.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"scaler_state_dict\": scaler.state_dict() if scaler else None,\n        \"train_loss\": train_loss,\n        \"train_accuracy\": train_accuracy,\n        \"val_loss\": val_loss,\n        \"val_accuracy\": val_accuracy,\n    }, checkpoint_path)\n    \n    print(f\"Checkpoint saved: {checkpoint_path}\")\n\n# Print overall training time\nend_time = time.time()\ntotal_time = end_time - start_time\nprint(f\"\\nTotal training time: {total_time:.2f} seconds\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T12:37:42.108044Z","iopub.execute_input":"2025-04-25T12:37:42.108396Z","iopub.status.idle":"2025-04-25T13:05:41.219951Z","shell.execute_reply.started":"2025-04-25T12:37:42.108355Z","shell.execute_reply":"2025-04-25T13:05:41.218934Z"}},"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":"total_time/60","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:06:13.665859Z","iopub.execute_input":"2025-04-25T13:06:13.666181Z","iopub.status.idle":"2025-04-25T13:06:13.671939Z","shell.execute_reply.started":"2025-04-25T13:06:13.666155Z","shell.execute_reply":"2025-04-25T13:06:13.671313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Calculate total number of parameters\ntotal_params = sum(p.numel() for p in model_co.parameters())\n\n# Print the result\nprint(f\"Total number of parameters in the model: {total_params}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:06:14.858853Z","iopub.execute_input":"2025-04-25T13:06:14.859138Z","iopub.status.idle":"2025-04-25T13:06:14.864987Z","shell.execute_reply.started":"2025-04-25T13:06:14.859105Z","shell.execute_reply":"2025-04-25T13:06:14.864370Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the checkpoint\ncheckpoint_path = f\"{MODEL_NAME}_last_checkpoint.pt\" # Replace with your latest checkpoint file\ncheckpoint = torch.load(checkpoint_path)\n\n# Extract metrics from the checkpoint\nepochs = list(range(1, checkpoint[\"epoch\"] + 1))\ntrain_loss = checkpoint[\"train_loss\"]\ntrain_accuracy = checkpoint[\"train_accuracy\"]\nval_loss = checkpoint[\"val_loss\"]\nval_accuracy = checkpoint[\"val_accuracy\"]\n\n# Set Seaborn style for better aesthetics\nsns.set_theme(style=\"ticks\")\n\n# Create a figure for visualizations\nplt.figure(figsize=(12, 8))\n\n# Plot training and validation loss\nplt.subplot(2, 1, 1)\nplt.plot(epochs, train_loss, label='Train Loss', color='blue')\nplt.plot(epochs, val_loss, label='Validation Loss', color='orange')\nplt.title(\"Training and Validation Loss per Epoch\", fontsize=14, fontweight='bold')\nplt.xlabel(\"Epoch\", fontsize=12)\nplt.ylabel(\"Loss\", fontsize=12)\nplt.legend()\n\n# Plot training and validation accuracy\nplt.subplot(2, 1, 2)\nplt.plot(epochs, train_accuracy, label='Train Accuracy', color='green')\nplt.plot(epochs, val_accuracy, label='Validation Accuracy', color='red')\nplt.title(\"Training and Validation Accuracy per Epoch\", fontsize=14, fontweight='bold')\nplt.xlabel(\"Epoch\", fontsize=12)\nplt.ylabel(\"Accuracy (%)\", fontsize=12)\nplt.legend()\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-04-25T13:06:16.320180Z","iopub.execute_input":"2025-04-25T13:06:16.320852Z","iopub.status.idle":"2025-04-25T13:06:18.268949Z","shell.execute_reply.started":"2025-04-25T13:06:16.320831Z","shell.execute_reply":"2025-04-25T13:06:18.268163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef load_model_from_checkpoint(checkpoint_path, model_class, device):\n    # Load the checkpoint\n    checkpoint = torch.load(checkpoint_path, map_location=device)\n    # Load model weights\n    model_co.load_state_dict(checkpoint[\"model_state_dict\"])\n    model_co.to(device)\n    model_co.eval()  # Set to evaluation mode\n    return model_co","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:06:18.269982Z","iopub.execute_input":"2025-04-25T13:06:18.270210Z","iopub.status.idle":"2025-04-25T13:06:18.274290Z","shell.execute_reply.started":"2025-04-25T13:06:18.270193Z","shell.execute_reply":"2025-04-25T13:06:18.273662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Predict on validation data\ndef predict(model_co, dataloader, device):\n    model_co.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 dataloader:  # 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-04-25T13:06:24.885761Z","iopub.execute_input":"2025-04-25T13:06:24.886454Z","iopub.status.idle":"2025-04-25T13:06:24.891055Z","shell.execute_reply.started":"2025-04-25T13:06:24.886432Z","shell.execute_reply":"2025-04-25T13:06:24.890380Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define paths and device\ncheckpoint_path = \"/kaggle/working/Mobile_Pit_last_checkpoint.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_train, y_pred_train = predict(model_co, train_loader, device)\n\ny_true_val, y_pred_val = predict(model_co, val_loader, device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:10:56.392893Z","iopub.execute_input":"2025-04-25T13:10:56.393481Z","iopub.status.idle":"2025-04-25T13:11:21.135600Z","shell.execute_reply.started":"2025-04-25T13:10:56.393456Z","shell.execute_reply":"2025-04-25T13:11:21.134755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_accuracy= accuracy_score(y_true_train, y_pred_train)\nprint(f\"Train Accuracy: {round(train_accuracy * 100, 2)}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:11:21.136974Z","iopub.execute_input":"2025-04-25T13:11:21.137264Z","iopub.status.idle":"2025-04-25T13:11:21.144222Z","shell.execute_reply.started":"2025-04-25T13:11:21.137241Z","shell.execute_reply":"2025-04-25T13:11:21.143628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Validation_accuracy = accuracy_score(y_true_val, y_pred_val)\nprint(f\"Validation Accuracy: {round(Validation_accuracy * 100, 2)}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:11:23.026209Z","iopub.execute_input":"2025-04-25T13:11:23.026978Z","iopub.status.idle":"2025-04-25T13:11:23.032297Z","shell.execute_reply.started":"2025-04-25T13:11:23.026952Z","shell.execute_reply":"2025-04-25T13:11:23.031620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"report = classification_report(y_true_val, y_pred_val, digits=2)\nprint(\"\\nClassification Report:\\n\", report)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:11:26.374158Z","iopub.execute_input":"2025-04-25T13:11:26.374476Z","iopub.status.idle":"2025-04-25T13:11:26.387315Z","shell.execute_reply.started":"2025-04-25T13:11:26.374453Z","shell.execute_reply":"2025-04-25T13:11:26.386771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport numpy as np\n\n# Assuming y_true and y_pred are your true and predicted labels, respectively\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Map multiclass labels to binary labels\nbinary_y_true = [0 if label == 0 else 1 for label in y_true_val]\nbinary_y_pred = [0 if label == 0 else 1 for label in y_pred_val]\n\n# Compute confusion matrix\ncm = confusion_matrix(binary_y_true, binary_y_pred)\nplt.figure(figsize=(6, 5))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=[\"No_DR\", \"DR\"], yticklabels=[\"No_DR\", \"DR\"])\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.title(\"Confusion Matrix (Binary Classification)\")\nplt.savefig('confusion matrix swin as binary calssification', dpi=300)\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:11:38.460627Z","iopub.execute_input":"2025-04-25T13:11:38.461297Z","iopub.status.idle":"2025-04-25T13:11:38.813026Z","shell.execute_reply.started":"2025-04-25T13:11:38.461272Z","shell.execute_reply":"2025-04-25T13:11:38.812371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Compute confusion matrix\ncm = confusion_matrix(y_true_val, y_pred_val)\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.savefig('confusion_matrixResnet50.png', bbox_inches='tight')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:12:42.560097Z","iopub.execute_input":"2025-04-25T13:12:42.560416Z","iopub.status.idle":"2025-04-25T13:12:42.910196Z","shell.execute_reply.started":"2025-04-25T13:12:42.560394Z","shell.execute_reply":"2025-04-25T13:12:42.909489Z"}},"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-04-25T13:11:47.926777Z","iopub.execute_input":"2025-04-25T13:11:47.927333Z","iopub.status.idle":"2025-04-25T13:11:47.932831Z","shell.execute_reply.started":"2025-04-25T13:11:47.927313Z","shell.execute_reply":"2025-04-25T13:11:47.932013Z"}},"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-04-25T13:11:49.386636Z","iopub.execute_input":"2025-04-25T13:11:49.387206Z","iopub.status.idle":"2025-04-25T13:11:49.391265Z","shell.execute_reply.started":"2025-04-25T13:11:49.387184Z","shell.execute_reply":"2025-04-25T13:11:49.390693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# رسم Barplot باستخدام Seaborn\nplt.figure(figsize=(10, 6))  # حجم الرسم البياني\nsns.barplot(x='Class', y='Error Rate', data=error_data, palette='viridis')\n\n# إضافة تسميات للمحاور\nplt.title('Error Rates per Class', fontsize=16)\nplt.xlabel('Class', fontsize=14)\nplt.ylabel('Error Rate', fontsize=14)\n\n# تحسين عرض الرسم البياني\nplt.xticks(rotation=45, ha='right')  # تدوير التسميات إذا كانت كثيفة\n\n# حفظ الصورة\nplt.tight_layout()\nplt.savefig('Error_Rates_Barplot.png')\n\n# عرض الرسم البياني\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:11:52.446417Z","iopub.execute_input":"2025-04-25T13:11:52.446981Z","iopub.status.idle":"2025-04-25T13:11:52.675572Z","shell.execute_reply.started":"2025-04-25T13:11:52.446962Z","shell.execute_reply":"2025-04-25T13:11:52.674802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the preprocessing pipeline (same as during training)\ntrain_transforms_DeiT_base_patch16 = T.Compose([\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\ndef process_external_image(image_path, transform=train_transforms_DeiT_base_patch16):\n    \"\"\"\n    Process an external image to match the model's expected input format.\n    \"\"\"\n    # Load the image\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  # Convert to RGB\n\n    # Crop, resize, and apply CLAHE preprocessing\n    image_cropped = crop_image_from_gray(image)\n    image_resized = cv2.resize(image_cropped, (256, 256))\n\n    # Apply CLAHE enhancement (you can reuse the previous CLAHE code for this)\n    blue, green, red = cv2.split(image_resized)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    blue_clahe = clahe.apply(blue)\n    green_clahe = clahe.apply(green)\n    red_clahe = clahe.apply(red)\n    result_image = cv2.merge([blue_clahe, green_clahe, red_clahe])\n\n    # Convert to PIL Image for transformation\n    pil_image = Image.fromarray(result_image)\n\n    # Apply the transform (normalization)\n    image_tensor = transform(pil_image)\n\n    # Add batch dimension\n    image_tensor = image_tensor.unsqueeze(0)\n    \n    return image_tensor\n\ndef predict_on_external_image(model, image_tensor, device):\n    \"\"\"\n    Predict the class label of an external image using the trained model.\n    \"\"\"\n    model.eval()  # Set model to evaluation mode\n    image_tensor = image_tensor.to(device)  # Move the image tensor to the device (GPU or CPU)\n\n    # Make prediction with the model\n    with torch.no_grad():\n        output = model(image_tensor)\n        _, predicted_label = torch.max(output, 1)  # Get the predicted label\n    \n    return predicted_label.item()\n\n# Example Usage:\nimage_path = \"/kaggle/input/diabetic-retinopathy-resized/resized_train/resized_train/10003_left.jpeg\"  # Replace with the path to the external image\n\n# Process the external image\nimage_tensor = process_external_image(image_path)\n\n# Predict the label\npredicted_label = predict_on_external_image(model, image_tensor, device)\n\n# Map the predicted label to category (based on your labels)\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Print the predicted category\nprint(f\"Predicted label: {level_to_category[predicted_label]}\")\n\n# Display the image\nimage = cv2.imread(image_path)\nimage_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\nplt.imshow(image_rgb)\nplt.title(f\"Predicted: {level_to_category[predicted_label]}\")\nplt.axis('off')  # Turn off axis\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-25T13:11:54.024110Z","iopub.execute_input":"2025-04-25T13:11:54.024509Z","iopub.status.idle":"2025-04-25T13:11:54.056008Z","shell.execute_reply.started":"2025-04-25T13:11:54.024482Z","shell.execute_reply":"2025-04-25T13:11:54.054931Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# End","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}