{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":7352272,"sourceType":"datasetVersion","datasetId":4244554}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### 1: SETUP AND CONFIGURATION","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport seaborn as sns \nfrom tqdm.notebook import tqdm\nimport time\nimport hashlib\nfrom collections import defaultdict\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nimport torchvision.transforms.functional as TF\n\nfrom sklearn.metrics import confusion_matrix, classification_report,cohen_kappa_score\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\nfrom sklearn.preprocessing import label_binarize\n\n# ====================================================================================================\n# SECTION 1: SETUP AND CONFIGURATION\n# ====================================================================================================\n\n# Set seeds for reproducibility to ensure consistent results across runs\ndef seed_everything(seed=42):\n    \"\"\"Set seeds for all random number generators to ensure reproducibility.\"\"\"\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True  # Ensures that CUDA selects deterministic algorithms\n    torch.backends.cudnn.benchmark = False     # Disables CUDA benchmark to ensure reproducibility\n\n# Initialize all random seeds\nseed_everything()\n\n# Availability of CUDA\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Define paths to the datasets\naptos_path = \"/kaggle/input/aptos2019-blindness-detection\"\neyepacs_path = \"/kaggle/input/eyepacs-aptos-messidor-diabetic-retinopathy\"\nprint(f\"APTOS 2019 path: {aptos_path}\")\nprint(f\"EyePACS-APTOS-Messidor path: {eyepacs_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:13:50.334155Z","iopub.execute_input":"2025-04-23T22:13:50.334423Z","iopub.status.idle":"2025-04-23T22:13:58.914111Z","shell.execute_reply.started":"2025-04-23T22:13:50.334401Z","shell.execute_reply":"2025-04-23T22:13:58.912824Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###  2: DATA LOADING AND PREPROCESSING","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 2: DATA LOADING AND PREPROCESSING\n# ====================================================================================================\n\n# Function to get file list from EyePACS dataset folder structure\ndef get_eyepacs_files(root_path, subset='train'):\n    \"\"\"\n    Get files from EyePACS folder structure where images are organized in class folders.\n    \n    Args:\n        root_path (str): Path to the EyePACS dataset\n        subset (str): Dataset subset to use ('train' or 'test')\n        \n    Returns:\n        dict: Dictionary with class numbers as keys and lists of file info as values\n    \"\"\"\n    files_dict = defaultdict(list)\n    subset_path = os.path.join(root_path, 'augmented_resized_V2', subset)\n    \n    # Check if path exists\n    if not os.path.exists(subset_path):\n        print(f\"Error: Path {subset_path} does not exist\")\n        return files_dict\n    \n    # Iterate through class folders (0, 1, 2, 3, 4)\n    for class_folder in os.listdir(subset_path):\n        if class_folder.isdigit():  # Only process folders that are class numbers\n            class_num = int(class_folder)\n            class_path = os.path.join(subset_path, class_folder)\n            \n            if os.path.isdir(class_path):\n                # Get all files in this class folder\n                for file_name in os.listdir(class_path):\n                    if file_name.endswith(('.jpg', '.jpeg', '.png')):\n                        file_path = os.path.join(class_path, file_name)\n                        files_dict[class_num].append({\n                            'image_id': file_name,\n                            'class': class_num,\n                            'image_path': file_path,\n                            'source': 'eyepacs'\n                        })\n    \n    # Print summary information\n    total_files = sum(len(files) for files in files_dict.values())\n    print(f\"Found {total_files} images in EyePACS {subset} folder\")\n    \n    # Print class distribution\n    for cls in sorted(files_dict.keys()):\n        print(f\"  Class {cls}: {len(files_dict[cls])} images\")\n    \n    return files_dict\n\n# Function to load APTOS dataset from CSV file\ndef load_aptos_dataset(aptos_path):\n    \"\"\"\n    Load APTOS dataset from CSV file and preprocess for consistency.\n    \n    Args:\n        aptos_path (str): Path to the APTOS dataset\n        \n    Returns:\n        DataFrame: DataFrame containing processed APTOS dataset\n    \"\"\"\n    csv_path = os.path.join(aptos_path, 'train.csv')\n    if not os.path.exists(csv_path):\n        print(f\"Error: APTOS CSV file not found at {csv_path}\")\n        return pd.DataFrame()\n    \n    # Load CSV data\n    df_aptos = pd.read_csv(csv_path)\n    print(f\"Loaded APTOS 2019: {len(df_aptos)} images\")\n    \n    # Add source and image path columns for consistency\n    df_aptos['source'] = 'aptos'\n    df_aptos['image_path'] = df_aptos['id_code'].apply(\n        lambda x: os.path.join(aptos_path, 'train_images', f\"{x}.png\")\n    )\n    \n    # Rename columns for consistency with EyePACS data format\n    df_aptos = df_aptos.rename(columns={'id_code': 'image_id', 'diagnosis': 'class'})\n    \n    # Print class distribution\n    print(\"APTOS class distribution:\")\n    print(df_aptos['class'].value_counts().sort_index())\n    \n    return df_aptos\n\n# Load APTOS dataset\nprint(\"Executing load_aptos_dataset function...\")\naptos_df = load_aptos_dataset(aptos_path)\nprint(\"APTOS dataset loaded successfully.\")\n\n# Function to create a balanced dataset by combining APTOS and EyePACS\ndef create_balanced_dataset(aptos_df, eyepacs_files, target_count_per_class=3000):\n    \"\"\"\n    Create balanced dataset by combining APTOS and EyePACS images with equal samples per class.\n    \n    Args:\n        aptos_df (DataFrame): DataFrame containing APTOS dataset\n        eyepacs_files (dict): Dictionary with EyePACS files by class\n        target_count_per_class (int): Target number of images per class\n        \n    Returns:\n        DataFrame: Balanced dataset with combined images\n    \"\"\"\n    balanced_rows = []\n    \n    # For each class (0-4), take samples up to target count\n    for cls in range(5):\n        # Get APTOS images for this class\n        aptos_cls = aptos_df[aptos_df['class'] == cls]\n        aptos_count = len(aptos_cls)\n        \n        # Get all EyePACS images for this class\n        eyepacs_cls = eyepacs_files.get(cls, [])\n        eyepacs_count = len(eyepacs_cls)\n        \n        print(f\"Class {cls}: {aptos_count} from APTOS, {eyepacs_count} from EyePACS\")\n        \n        # Determine how many images we need from each source\n        if aptos_count >= target_count_per_class:\n            # If APTOS has enough, just use those\n            balanced_rows.extend(aptos_cls.sample(target_count_per_class, random_state=42).to_dict('records'))\n            print(f\"  Using {target_count_per_class} images from APTOS for class {cls}\")\n        else:\n            # Use all APTOS images\n            balanced_rows.extend(aptos_cls.to_dict('records'))\n            print(f\"  Using all {aptos_count} images from APTOS for class {cls}\")\n            \n            # Calculate how many more we need\n            needed_from_eyepacs = target_count_per_class - aptos_count\n            \n            if eyepacs_count >= needed_from_eyepacs:\n                # Randomly sample from EyePACS\n                selected_eyepacs = random.sample(eyepacs_cls, needed_from_eyepacs)\n                balanced_rows.extend(selected_eyepacs)\n                print(f\"  Adding {needed_from_eyepacs} images from EyePACS for class {cls}\")\n            else:\n                # Use all available from EyePACS\n                balanced_rows.extend(eyepacs_cls)\n                print(f\"  Adding all {eyepacs_count} images from EyePACS for class {cls}\")\n                print(f\"  Warning: Could only reach {aptos_count + eyepacs_count}/{target_count_per_class} for class {cls}\")\n    \n    # Convert to DataFrame\n    balanced_df = pd.DataFrame(balanced_rows)\n    \n    # Verify image paths exist (sample check)\n    sample_size = min(10, len(balanced_df))\n    for idx, row in balanced_df.sample(sample_size).iterrows():\n        print(f\"Checking {row['image_path']}: {os.path.exists(row['image_path'])}\")\n    \n    # Print final distribution\n    print(\"\\nFinal class distribution:\")\n    print(balanced_df['class'].value_counts().sort_index())\n    \n    return balanced_df\n\n# Load EyePACS data\nprint(\"\\nExecuting get_eyepacs_files function...\")\n# Load from train folder\neyepacs_train_files = get_eyepacs_files(eyepacs_path, subset='train')\n# Also load from test folder for more diversity\neyepacs_test_files = get_eyepacs_files(eyepacs_path, subset='test')\n\n# Combine train and test files\neyepacs_files = defaultdict(list)\nfor cls in range(5):\n    eyepacs_files[cls].extend(eyepacs_train_files.get(cls, []))\n    eyepacs_files[cls].extend(eyepacs_test_files.get(cls, []))\nprint(\"EyePACS files loaded successfully.\")\n\n# Create balanced dataset\nprint(\"\\nExecuting create_balanced_dataset function...\")\nbalanced_df = create_balanced_dataset(aptos_df, eyepacs_files, target_count_per_class=3000)\nprint(\"Balanced dataset created successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:14:03.727380Z","iopub.execute_input":"2025-04-23T22:14:03.728172Z","iopub.status.idle":"2025-04-23T22:14:05.003671Z","shell.execute_reply.started":"2025-04-23T22:14:03.728142Z","shell.execute_reply":"2025-04-23T22:14:05.002948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to split dataset into train, validation, and test sets\ndef create_data_splits(balanced_df, test_size=0.15, val_size=0.15):\n    \"\"\"\n    Split dataset into train, validation, and test sets with stratification.\n    \n    Args:\n        balanced_df (DataFrame): Balanced dataset to split\n        test_size (float): Proportion of data for testing\n        val_size (float): Proportion of data for validation\n        \n    Returns:\n        tuple: (train_df, val_df, test_df) - DataFrames for each split\n    \"\"\"\n    # First split out test set\n    train_val_df, test_df = train_test_split(\n        balanced_df, \n        test_size=test_size,\n        random_state=42,\n        stratify=balanced_df['class']  # Maintain class proportions\n    )\n    \n    # Recalculate validation size relative to remaining data\n    relative_val_size = val_size / (1 - test_size)\n    \n    # Split remaining data into train and validation\n    train_df, val_df = train_test_split(\n        train_val_df,\n        test_size=relative_val_size,\n        random_state=42,\n        stratify=train_val_df['class']  # Maintain class proportions\n    )\n    \n    print(f\"Data split sizes: Train={len(train_df)}, Validation={len(val_df)}, Test={len(test_df)}\")\n    \n    # Check class distribution in each split\n    print(\"\\nTrain class distribution:\")\n    print(train_df['class'].value_counts().sort_index())\n    \n    print(\"\\nValidation class distribution:\")\n    print(val_df['class'].value_counts().sort_index())\n    \n    print(\"\\nTest class distribution:\")\n    print(test_df['class'].value_counts().sort_index())\n    \n    return train_df, val_df, test_df\n\n# Split data into train, validation, and test sets\nprint(\"\\nExecuting create_data_splits function...\")\ntrain_df, val_df, test_df = create_data_splits(balanced_df)\nprint(\"Data splits created successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:14:10.302226Z","iopub.execute_input":"2025-04-23T22:14:10.302895Z","iopub.status.idle":"2025-04-23T22:14:10.327967Z","shell.execute_reply.started":"2025-04-23T22:14:10.302869Z","shell.execute_reply":"2025-04-23T22:14:10.327162Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3: DATA AUGMENTATION AND DATASET PREPARATION","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 3: DATA AUGMENTATION AND DATASET PREPARATION (according to paper-baseline model)\n# ====================================================================================================\n\n# Define data transformations for training (with augmentation)\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224)),              # Resize to standard input size\n    transforms.RandomAffine(                    # Apply random affine transformations\n        degrees=20,                             # Rotation range\n        translate=(0.1, 0.1),                   # Translation range\n        scale=(0.8, 1.2),                       # Scale range\n    ),\n    transforms.ColorJitter(brightness=(0.5, 1.0)),  # Randomly adjust brightness\n    transforms.ToTensor(),                      # Convert to tensor\n    transforms.Normalize(                       # Normalize with ImageNet mean and std\n        mean=[0.485, 0.456, 0.406], \n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# Define data transformations for validation and testing (no augmentation)\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),              # Resize to standard input size\n    transforms.ToTensor(),                      # Convert to tensor\n    transforms.Normalize(                       # Normalize with ImageNet mean and std\n        mean=[0.485, 0.456, 0.406], \n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# Custom dataset class for Diabetic Retinopathy images\nclass DiabeticRetinopathyDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset class for Diabetic Retinopathy images.\n    \n    Attributes:\n        dataframe (DataFrame): DataFrame containing image paths and labels\n        transform (callable): Transformation to apply to images\n    \"\"\"\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n        \n    def __len__(self):\n        \"\"\"Return the number of samples in the dataset.\"\"\"\n        return len(self.dataframe)\n    \n    def __getitem__(self, idx):\n        \"\"\"\n        Get an item from the dataset by index.\n        \n        Args:\n            idx (int): Index of the item to get\n            \n        Returns:\n            tuple: (image, label) where image is the transformed image and label is the class\n        \"\"\"\n        img_path = self.dataframe.iloc[idx]['image_path']\n        label = self.dataframe.iloc[idx]['class']\n        \n        try:\n            # Load and convert image to RGB\n            image = Image.open(img_path).convert('RGB')\n        except Exception as e:\n            print(f\"Error loading image {img_path}: {e}\")\n            # Return a blank image if loading fails\n            image = Image.new('RGB', (224, 224), color='black')\n        \n        # Apply transformations if specified\n        if self.transform:\n            image = self.transform(image)\n            \n        return image, label\n\n# Create datasets and dataloaders\nprint(\"\\nCreating datasets and dataloaders...\")\n# Create datasets\ntrain_dataset = DiabeticRetinopathyDataset(train_df, transform=train_transform)\nval_dataset = DiabeticRetinopathyDataset(val_df, transform=val_transform)\ntest_dataset = DiabeticRetinopathyDataset(test_df, transform=val_transform)\n\n# Create data loaders with specified batch size and workers\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=64,             # Number of samples per batch\n    shuffle=True,              # Shuffle data for training\n    num_workers=4,             # Number of parallel workers for data loading\n    pin_memory=True            # Pin memory for faster GPU transfer\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=64,\n    shuffle=False,             # No need to shuffle validation data\n    num_workers=4,\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=64,\n    shuffle=False,             # No need to shuffle test data\n    num_workers=4,\n    pin_memory=True\n)\n\nprint(f\"Created dataloaders with {len(train_dataset)} training, {len(val_dataset)} validation, and {len(test_dataset)} test images\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:14:17.363151Z","iopub.execute_input":"2025-04-23T22:14:17.363921Z","iopub.status.idle":"2025-04-23T22:14:17.375352Z","shell.execute_reply.started":"2025-04-23T22:14:17.363892Z","shell.execute_reply":"2025-04-23T22:14:17.374518Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4: MODEL ARCHITECTURE","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 4: MODEL ARCHITECTURE\n# ====================================================================================================\n\n# Define custom MobileNetV2 model with additional fully connected layers\nclass CustomMobileNetV2(nn.Module):\n    \"\"\"\n    Custom MobileNetV2 model with modified classifier for Diabetic Retinopathy classification.\n    \n    Attributes:\n        mobilenet (nn.Module): MobileNetV2 backbone with pretrained weights\n    \"\"\"\n    def __init__(self, num_classes=5):\n        \"\"\"\n        Initialize the model with pretrained MobileNetV2 backbone and custom classifier.\n        \n        Args:\n            num_classes (int): Number of output classes (5 for DR grades)\n        \"\"\"\n        super(CustomMobileNetV2, self).__init__()\n        # Load pretrained MobileNetV2 model\n        self.mobilenet = models.mobilenet_v2(pretrained=True)\n        \n        # Replace classifier with custom layers\n        self.mobilenet.classifier = nn.Sequential(\n            nn.Dropout(0.2),                          # Dropout for regularization\n            nn.Linear(self.mobilenet.last_channel, 32),  # First FC layer\n            nn.ReLU(),                                # Activation function\n            nn.Linear(32, 16),                        # Second FC layer\n            nn.ReLU(),                                # Activation function\n            nn.Linear(16, num_classes)                # Output layer\n        )\n    \n    def forward(self, x):\n        \"\"\"Forward pass through the network.\"\"\"\n        return self.mobilenet(x)\n\n# Initialize model\nprint(\"\\nInitializing model...\")\nmodel = CustomMobileNetV2(num_classes=5).to(device)\n\n# Define loss function - standard cross entropy as per paper\ncriterion = nn.CrossEntropyLoss()\n\n# Define optimizer - Adam with learning rate 0.0001 as per paper\noptimizer = optim.Adam(model.parameters(), lr=0.0001)\n\n# Learning rate scheduler to reduce LR when validation loss plateaus\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, \n    mode='min',               # Monitor minimum validation loss\n    factor=0.5,               # Multiply LR by this factor when reducing\n    patience=5,               # Wait for 5 epochs before reducing LR\n    verbose=True              # Print message when reducing LR\n)\nprint(\"Model initialized successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:14:33.342152Z","iopub.execute_input":"2025-04-23T22:14:33.342432Z","iopub.status.idle":"2025-04-23T22:14:33.911807Z","shell.execute_reply.started":"2025-04-23T22:14:33.342411Z","shell.execute_reply":"2025-04-23T22:14:33.910926Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5: TRAINING FUNCTION","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 5: TRAINING FUNCTION\n# ====================================================================================================\n\n# Function to train the model\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=100, patience=7):\n    \"\"\"\n    Train the model with early stopping and learning rate scheduling.\n    \n    Args:\n        model (nn.Module): Model to train\n        train_loader (DataLoader): Training data loader\n        val_loader (DataLoader): Validation data loader\n        criterion (nn.Module): Loss function\n        optimizer (optim.Optimizer): Optimizer\n        scheduler: Learning rate scheduler\n        num_epochs (int): Maximum number of epochs to train\n        patience (int): Early stopping patience (epochs without improvement)\n        \n    Returns:\n        tuple: (trained_model, history) - Trained model and training history\n    \"\"\"\n    best_val_loss = float('inf')\n    counter = 0  # Counter for early stopping\n    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []}\n    start_time = time.time()\n    \n    best_model_path = 'baseline_model.pth'\n    \n    for epoch in range(num_epochs):\n        epoch_start = time.time()\n        \n        # -----------------\n        # Training phase\n        # -----------------\n        model.train()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n        \n        # Track class-wise accuracy\n        class_correct = [0] * 5\n        class_total = [0] * 5\n        \n        train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Train]')\n        for inputs, labels in train_bar:\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n            \n            # Zero gradients\n            optimizer.zero_grad()\n            \n            # Forward pass\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            \n            # Backward pass and optimize\n            loss.backward()\n            optimizer.step()\n            \n            # Track statistics\n            running_loss += loss.item() * inputs.size(0)\n            _, predicted = torch.max(outputs, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n            \n            # Class-wise accuracy\n            for i in range(5):\n                label_mask = (labels == i)\n                class_total[i] += label_mask.sum().item()\n                if label_mask.sum() > 0:\n                    class_correct[i] += (predicted[label_mask] == i).sum().item()\n            \n            # Update progress bar\n            train_bar.set_postfix(loss=loss.item(), acc=correct/total if total > 0 else 0)\n        \n        # Calculate epoch metrics\n        epoch_train_loss = running_loss / len(train_loader.dataset) if len(train_loader.dataset) > 0 else 0\n        epoch_train_acc = correct / total if total > 0 else 0\n        history['train_loss'].append(epoch_train_loss)\n        history['train_acc'].append(epoch_train_acc)\n        \n        # -----------------\n        # Validation phase\n        # -----------------\n        model.eval()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n        \n        # Class-wise validation accuracy\n        val_class_correct = [0] * 5\n        val_class_total = [0] * 5\n        \n        # Confusion matrix for per-class metrics\n        confusion_mat = torch.zeros(5, 5)\n        \n        with torch.no_grad():  # No gradients needed for validation\n            val_bar = tqdm(val_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Val]')\n            for inputs, labels in val_bar:\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n                \n                # Forward pass\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n                \n                # Track statistics\n                running_loss += loss.item() * inputs.size(0)\n                _, predicted = torch.max(outputs, 1)\n                total += labels.size(0)\n                correct += (predicted == labels).sum().item()\n                \n                # Update confusion matrix\n                for t, p in zip(labels.view(-1), predicted.view(-1)):\n                    confusion_mat[t.long(), p.long()] += 1\n                \n                # Class-wise accuracy\n                for i in range(5):\n                    label_mask = (labels == i)\n                    val_class_total[i] += label_mask.sum().item()\n                    if label_mask.sum() > 0:\n                        val_class_correct[i] += (predicted[label_mask] == i).sum().item()\n                \n                # Update progress bar\n                val_bar.set_postfix(loss=loss.item(), acc=correct/total if total > 0 else 0)\n        \n        # Calculate epoch metrics\n        epoch_val_loss = running_loss / len(val_loader.dataset) if len(val_loader.dataset) > 0 else 0\n        epoch_val_acc = correct / total if total > 0 else 0\n        history['val_loss'].append(epoch_val_loss)\n        history['val_acc'].append(epoch_val_acc)\n        \n        # Update learning rate based on validation loss\n        scheduler.step(epoch_val_loss)\n        \n        # Print metrics\n        epoch_end = time.time()\n        epoch_time = (epoch_end - epoch_start) / 60  # in minutes\n        \n        print(f\"Epoch {epoch+1}/{num_epochs} completed in {epoch_time:.2f} minutes\")\n        print(f\"Train Loss: {epoch_train_loss:.4f}, Train Acc: {epoch_train_acc:.4f}\")\n        \n        # Print class-wise training accuracy\n        for i in range(5):\n            if class_total[i] > 0:\n                print(f\"Training Accuracy of class {i}: {100 * class_correct[i] / class_total[i]:.2f}%\")\n            else:\n                print(f\"Training Accuracy of class {i}: N/A (no training examples)\")\n        \n        print(f\"Val Loss: {epoch_val_loss:.4f}, Val Acc: {epoch_val_acc:.4f}\")\n        \n        # Print class-wise validation accuracy\n        for i in range(5):\n            if val_class_total[i] > 0:\n                print(f\"Validation Accuracy of class {i}: {100 * val_class_correct[i] / val_class_total[i]:.2f}%\")\n            else:\n                print(f\"Validation Accuracy of class {i}: N/A (no validation examples)\")\n        \n        # Calculate and print per-class metrics\n        for i in range(5):\n            tp = confusion_mat[i, i]\n            fp = confusion_mat[:, i].sum() - tp\n            fn = confusion_mat[i, :].sum() - tp\n            tn = confusion_mat.sum() - (tp + fp + fn)\n            \n            # Calculate performance metrics\n            sensitivity = tp / (tp + fn) if tp + fn > 0 else 0\n            specificity = tn / (tn + fp) if tn + fp > 0 else 0\n            precision = tp / (tp + fp) if tp + fp > 0 else 0\n            \n            print(f\"Class {i} - Sensitivity: {sensitivity:.4f}, Specificity: {specificity:.4f}, Precision: {precision:.4f}\")\n        \n        # Check if we should save the model (when validation loss improves)\n        if epoch_val_loss < best_val_loss:\n            print(f\"Validation loss decreased ({best_val_loss:.6f} --> {epoch_val_loss:.6f}). Saving model...\")\n            best_val_loss = epoch_val_loss\n            torch.save(model.state_dict(), best_model_path)\n            counter = 0  # Reset early stopping counter\n        else:\n            counter += 1\n            print(f\"Early stopping counter: {counter} out of {patience}\")\n            if counter >= patience:\n                print(\"Early stopping triggered\")\n                break\n    \n    total_time = (time.time() - start_time) / 60  # in minutes\n    print(f\"Training completed in {total_time:.2f} minutes\")\n    \n    # Load the best model\n    model.load_state_dict(torch.load(best_model_path))\n    return model, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:14:39.743955Z","iopub.execute_input":"2025-04-23T22:14:39.744553Z","iopub.status.idle":"2025-04-23T22:14:39.762239Z","shell.execute_reply.started":"2025-04-23T22:14:39.744514Z","shell.execute_reply":"2025-04-23T22:14:39.761452Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"###  6: MODEL TRAINING","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 6: MODEL TRAINING\n# ====================================================================================================\n\n# Train the model\nprint(\"\\nStarting model training...\")\ntrained_model, history = train_model(\n    model, \n    train_loader, \n    val_loader, \n    criterion, \n    optimizer,\n    scheduler,\n    num_epochs=50,  # Epochs 50 for faster testing\n    patience=20\n)\nprint(\"Model training completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:14:49.342447Z","iopub.execute_input":"2025-04-23T22:14:49.342774Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 7: VISUALIZATION AND EVALUATION","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 7: VISUALIZATION AND EVALUATION\n# ====================================================================================================\n\n# Plot training history\nprint(\"\\nPlotting training history...\")\nplt.figure(figsize=(12, 5))\n\n# Plot loss curves\nplt.subplot(1, 2, 1)\nplt.plot(history['train_loss'], label='Train Loss')\nplt.plot(history['val_loss'], label='Validation Loss')\nplt.title('Model Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\n\n# Plot accuracy curves\nplt.subplot(1, 2, 2)\nplt.plot(history['train_acc'], label='Train Accuracy')\nplt.plot(history['val_acc'], label='Validation Accuracy')\nplt.title('Model Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.legend()\n\nplt.tight_layout()\nplt.savefig('training_history.png')\nplt.show()\nprint(\"Training history plotted successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:13:35.924630Z","iopub.status.idle":"2025-04-23T22:13:35.924940Z","shell.execute_reply.started":"2025-04-23T22:13:35.924781Z","shell.execute_reply":"2025-04-23T22:13:35.924796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#evaluate_model function\ndef evaluate_model(model, test_loader, criterion, class_names=None):\n    \"\"\"\n    Evaluate the model on test data and report metrics including Cohen's Kappa.\n    \n    Args:\n        model (nn.Module): Model to evaluate\n        test_loader (DataLoader): Test data loader\n        criterion (nn.Module): Loss function\n        class_names (list): Names of classes\n        \n    Returns:\n        dict: Dictionary containing evaluation metrics\n    \"\"\"\n    if class_names is None:\n        class_names = ['No DR', 'Mild DR', 'Moderate DR', 'Severe DR', 'Proliferative DR']\n    \n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    all_labels = []\n    all_predictions = []\n    all_outputs = []\n    \n    with torch.no_grad():\n        for inputs, labels in tqdm(test_loader, desc='Testing'):\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n            \n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item() * inputs.size(0)\n            _, predicted = torch.max(outputs, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n            \n            all_labels.extend(labels.cpu().numpy())\n            all_predictions.extend(predicted.cpu().numpy())\n            all_outputs.extend(F.softmax(outputs, dim=1).cpu().numpy())\n    \n    test_loss = running_loss / len(test_loader.dataset)\n    test_acc = correct / total\n    \n    print(f\"Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}\")\n    \n    # Calculate and display confusion matrix and per-class metrics\n    cm = confusion_matrix(all_labels, all_predictions)\n    print(\"Confusion Matrix:\")\n    print(cm)\n    \n    # Create and display confusion matrix heatmap\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                xticklabels=class_names, \n                yticklabels=class_names)\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.title('Confusion Matrix')\n    plt.tight_layout()\n    plt.savefig('confusion_matrix.png')\n    plt.show()\n    \n    # Print classification report\n    print(\"\\nClassification Report:\")\n    print(classification_report(all_labels, all_predictions, target_names=class_names))\n    \n    # Calculate metrics\n    accuracy = accuracy_score(all_labels, all_predictions)\n    precision = precision_score(all_labels, all_predictions, average='weighted')\n    sensitivity = recall_score(all_labels, all_predictions, average='weighted')\n    \n    # Calculate Cohen's Kappa\n    kappa = cohen_kappa_score(all_labels, all_predictions)\n    weighted_kappa = cohen_kappa_score(all_labels, all_predictions, weights='quadratic')\n    \n    print(f\"\\nCohen's Kappa: {kappa:.4f}\")\n    print(f\"Quadratic Weighted Kappa: {weighted_kappa:.4f}\")\n    \n    # Interpret Kappa value\n    if kappa < 0:\n        interpretation = \"Poor agreement (worse than random)\"\n    elif kappa < 0.2:\n        interpretation = \"Slight agreement\"\n    elif kappa < 0.4:\n        interpretation = \"Fair agreement\"\n    elif kappa < 0.6:\n        interpretation = \"Moderate agreement\"\n    elif kappa < 0.8:\n        interpretation = \"Substantial agreement\"\n    else:\n        interpretation = \"Almost perfect agreement\"\n    \n    print(f\"Interpretation: {interpretation}\")\n    \n    # Calculate specificity\n    specificity_list = []\n    for i in range(5):\n        true_negatives = np.sum(np.delete(np.delete(cm, i, 0), i, 1))\n        false_positives = np.sum(np.delete(cm[:, i], i))\n        specificity = true_negatives / (true_negatives + false_positives) if (true_negatives + false_positives) > 0 else 0\n        specificity_list.append(specificity)\n    \n    specificity = np.mean(specificity_list)\n    f1 = f1_score(all_labels, all_predictions, average='weighted')\n    \n    # For AUC-ROC\n    all_outputs = np.array(all_outputs)\n    all_labels = np.array(all_labels)\n    \n    # One-hot encode true labels for multiclass ROC AUC\n    y_true_bin = label_binarize(all_labels, classes=range(5))\n    \n    # Calculate AUC for each class\n    auc_scores = []\n    for i in range(5):\n        if len(np.unique(y_true_bin[:, i])) > 1:\n            auc_scores.append(roc_auc_score(y_true_bin[:, i], all_outputs[:, i]))\n    \n    auc_roc = np.mean(auc_scores) if auc_scores else 0\n    \n    # Calculate model size and ASR\n    model_size_mb = os.path.getsize('baseline_model.pth') / (1024 * 1024)\n    asr = accuracy / model_size_mb\n    \n    print(\"\\nBaseline Paper Metrics:\")\n    print(f\"Accuracy: {accuracy:.4f}\")\n    print(f\"Sensitivity: {sensitivity:.4f}\")\n    print(f\"Specificity: {specificity:.4f}\")\n    print(f\"Precision: {precision:.4f}\")\n    print(f\"F1-Score: {f1:.4f}\")\n    print(f\"AUC-ROC: {auc_roc:.4f}\")\n    print(f\"Model Size: {model_size_mb:.2f} MB\")\n    print(f\"Accuracy-to-Size Ratio (ASR): {asr:.6f}\")\n    \n    # Print per-class sensitivities\n    class_recalls = recall_score(all_labels, all_predictions, average=None)\n    print(\"\\nPer-class Sensitivity:\")\n    for i, recall in enumerate(class_recalls):\n        print(f\"Class {i} ({class_names[i]}): {recall:.4f}\")\n    \n    # Generate a detailed per-class confusion matrix visualization\n    plt.figure(figsize=(12, 10))\n    cm_normalized = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]\n    \n    # Plot normalized confusion matrix\n    sns.heatmap(cm_normalized, annot=True, fmt='.2f', cmap='Blues',\n                xticklabels=class_names, \n                yticklabels=class_names)\n    plt.xlabel('Predicted Label', fontsize=12)\n    plt.ylabel('True Label', fontsize=12)\n    plt.title('Normalized Confusion Matrix', fontsize=14)\n    plt.tight_layout()\n    plt.savefig('normalized_confusion_matrix.png')\n    plt.show()\n    \n    return {\n        'accuracy': accuracy,\n        'sensitivity': sensitivity,\n        'specificity': specificity,\n        'precision': precision,\n        'f1': f1,\n        'auc_roc': auc_roc,\n        'model_size_mb': model_size_mb,\n        'asr': asr,\n        'class_recalls': class_recalls,\n        'confusion_matrix': cm,\n        'kappa': kappa,\n        'weighted_kappa': weighted_kappa\n    }\n\n# Evaluate the model\nprint(\"\\nEvaluating model...\")\nclass_names = ['No DR', 'Mild DR', 'Moderate DR', 'Severe DR', 'Proliferative DR']\ntest_metrics = evaluate_model(trained_model, test_loader, criterion, class_names)\nprint(\"Model evaluation completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:13:35.929142Z","iopub.status.idle":"2025-04-23T22:13:35.929450Z","shell.execute_reply.started":"2025-04-23T22:13:35.929290Z","shell.execute_reply":"2025-04-23T22:13:35.929318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to plot class-wise performance metrics as bar charts\ndef plot_class_performance(test_metrics, class_names):\n    \"\"\"\n    Plot per-class performance metrics as bar charts.\n    \n    Args:\n        test_metrics (dict): Dictionary containing evaluation metrics\n        class_names (list): Names of classes\n    \"\"\"\n    # Extract class-wise metrics from confusion matrix\n    cm = test_metrics['confusion_matrix']\n    class_metrics = []\n    \n    for i in range(len(class_names)):\n        tp = cm[i, i]\n        fp = np.sum(cm[:, i]) - tp\n        fn = np.sum(cm[i, :]) - tp\n        tn = np.sum(cm) - (tp + fp + fn)\n        \n        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n        recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n        specificity = tn / (tn + fp) if (tn + fp) > 0 else 0\n        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n        \n        class_metrics.append({\n            'class': class_names[i],\n            'precision': precision,\n            'recall': recall,\n            'specificity': specificity,\n            'f1': f1\n        })\n    \n    # Create a figure with multiple subplots\n    fig, axs = plt.subplots(2, 2, figsize=(14, 10))\n    fig.suptitle('Class-wise Performance Metrics', fontsize=16)\n    \n    # Plot precision\n    x = range(len(class_names))\n    axs[0, 0].bar(x, [m['precision'] for m in class_metrics], color='skyblue')\n    axs[0, 0].set_title('Precision')\n    axs[0, 0].set_xticks(x)\n    axs[0, 0].set_xticklabels(class_names, rotation=45)\n    axs[0, 0].set_ylim(0, 1)\n    \n    # Plot recall/sensitivity\n    axs[0, 1].bar(x, [m['recall'] for m in class_metrics], color='lightgreen')\n    axs[0, 1].set_title('Recall/Sensitivity')\n    axs[0, 1].set_xticks(x)\n    axs[0, 1].set_xticklabels(class_names, rotation=45)\n    axs[0, 1].set_ylim(0, 1)\n    \n    # Plot specificity\n    axs[1, 0].bar(x, [m['specificity'] for m in class_metrics], color='salmon')\n    axs[1, 0].set_title('Specificity')\n    axs[1, 0].set_xticks(x)\n    axs[1, 0].set_xticklabels(class_names, rotation=45)\n    axs[1, 0].set_ylim(0, 1)\n    \n    # Plot F1 score\n    axs[1, 1].bar(x, [m['f1'] for m in class_metrics], color='mediumpurple')\n    axs[1, 1].set_title('F1 Score')\n    axs[1, 1].set_xticks(x)\n    axs[1, 1].set_xticklabels(class_names, rotation=45)\n    axs[1, 1].set_ylim(0, 1)\n    \n    plt.tight_layout()\n    plt.subplots_adjust(top=0.9)\n    plt.savefig('class_performance_metrics.png')\n    plt.show()\n\n# Call the function to plot class-wise performance\nplot_class_performance(test_metrics, class_names)\n\n# Inference time measurement function\ndef measure_inference_time(model, device, batch_size=1, image_size=(224, 224), n_runs=100):\n    \"\"\"\n    Measure model inference time on both GPU and CPU.\n    \n    Args:\n        model (nn.Module): Model to evaluate\n        device (torch.device): Current device\n        batch_size (int): Batch size for inference\n        image_size (tuple): Input image dimensions\n        n_runs (int): Number of runs to average\n        \n    Returns:\n        dict: Dictionary containing GPU and CPU inference times\n    \"\"\"\n    # Create a dummy input\n    dummy_input = torch.randn(batch_size, 3, *image_size).to(device)\n    \n    # Warm-up runs to ensure fair measurement\n    for _ in range(10):\n        _ = model(dummy_input)\n    \n    # Measure GPU inference time\n    if device.type == 'cuda':\n        torch.cuda.synchronize()  # Wait for all CUDA operations to finish\n        start = time.time()\n        for _ in range(n_runs):\n            _ = model(dummy_input)\n            torch.cuda.synchronize()  # Wait for each forward pass to complete\n        gpu_time = (time.time() - start) / n_runs\n        print(f\"GPU inference time: {gpu_time*1000:.2f} ms per image\")\n    else:\n        gpu_time = float('nan')\n        \n    # Move model to CPU for CPU inference time\n    model_cpu = model.to('cpu')\n    dummy_input_cpu = torch.randn(batch_size, 3, *image_size)\n    \n    # Warm-up runs for CPU\n    for _ in range(10):\n        _ = model_cpu(dummy_input_cpu)\n    \n    # Measure CPU inference time\n    start = time.time()\n    for _ in range(n_runs):\n        _ = model_cpu(dummy_input_cpu)\n    cpu_time = (time.time() - start) / n_runs\n    print(f\"CPU inference time: {cpu_time*1000:.2f} ms per image\")\n    \n    # Move model back to original device\n    model.to(device)\n    \n    return {\n        'gpu_time': gpu_time if device.type == 'cuda' else None,\n        'cpu_time': cpu_time\n    }\n\n# Measure inference time\nprint(\"\\nMeasuring inference time...\")\ninference_times = measure_inference_time(trained_model, device)\nprint(\"Inference time measurement completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:13:35.994965Z","iopub.execute_input":"2025-04-23T22:13:35.995236Z","iopub.status.idle":"2025-04-23T22:13:36.016273Z","shell.execute_reply.started":"2025-04-23T22:13:35.995212Z","shell.execute_reply":"2025-04-23T22:13:36.014909Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 8: VISUALIZATIONS FOR CONFUSION MATRIX","metadata":{}},{"cell_type":"code","source":"# ====================================================================================================\n# SECTION 8: VISUALIZATIONS FOR CONFUSION MATRIX\n# ====================================================================================================\n\ndef plot_confusion_matrix_heatmap(cm, class_names):\n    \"\"\"\n    Create a more detailed and customized confusion matrix heatmap.\n    \n    Args:\n        cm (numpy.ndarray): Confusion matrix\n        class_names (list): Names of the classes\n    \"\"\"\n    # Calculate metrics for each class from confusion matrix\n    n_classes = len(class_names)\n    class_accuracy = np.zeros(n_classes)\n    class_precision = np.zeros(n_classes)\n    class_recall = np.zeros(n_classes)\n    class_f1 = np.zeros(n_classes)\n    \n    for i in range(n_classes):\n        tp = cm[i, i]\n        fp = np.sum(cm[:, i]) - tp\n        fn = np.sum(cm[i, :]) - tp\n        tn = np.sum(cm) - (tp + fp + fn)\n        \n        class_accuracy[i] = (tp + tn) / np.sum(cm)\n        class_precision[i] = tp / (tp + fp) if (tp + fp) > 0 else 0\n        class_recall[i] = tp / (tp + fn) if (tp + fn) > 0 else 0\n        class_f1[i] = 2 * class_precision[i] * class_recall[i] / (class_precision[i] + class_recall[i]) if (class_precision[i] + class_recall[i]) > 0 else 0\n    \n    # Set up the figure with multiple subplots\n    fig = plt.figure(figsize=(20, 15))\n    \n    # 1. Raw counts confusion matrix\n    ax1 = plt.subplot2grid((2, 2), (0, 0))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=ax1,\n                xticklabels=class_names, yticklabels=class_names)\n    ax1.set_xlabel('Predicted Label')\n    ax1.set_ylabel('True Label')\n    ax1.set_title('Confusion Matrix (Raw Counts)')\n    \n    # 2. Row-normalized confusion matrix (recall/sensitivity for each class)\n    ax2 = plt.subplot2grid((2, 2), (0, 1))\n    row_sums = cm.sum(axis=1)\n    cm_row_norm = cm / row_sums[:, np.newaxis]\n    sns.heatmap(cm_row_norm, annot=True, fmt='.2f', cmap='Greens', ax=ax2,\n                xticklabels=class_names, yticklabels=class_names)\n    ax2.set_xlabel('Predicted Label')\n    ax2.set_ylabel('True Label')\n    ax2.set_title('Row-Normalized Confusion Matrix (Recall/Sensitivity)')\n    \n    # 3. Column-normalized confusion matrix (precision for each class)\n    ax3 = plt.subplot2grid((2, 2), (1, 0))\n    col_sums = cm.sum(axis=0)\n    cm_col_norm = cm / col_sums[np.newaxis, :]\n    sns.heatmap(cm_col_norm, annot=True, fmt='.2f', cmap='Oranges', ax=ax3,\n                xticklabels=class_names, yticklabels=class_names)\n    ax3.set_xlabel('Predicted Label')\n    ax3.set_ylabel('True Label')\n    ax3.set_title('Column-Normalized Confusion Matrix (Precision)')\n    \n    # 4. Per-class metrics as a bar chart\n    ax4 = plt.subplot2grid((2, 2), (1, 1))\n    metrics_df = pd.DataFrame({\n        'Accuracy': class_accuracy,\n        'Precision': class_precision,\n        'Recall': class_recall,\n        'F1-Score': class_f1\n    }, index=class_names)\n    metrics_df.plot(kind='bar', ax=ax4, rot=45)\n    ax4.set_ylim(0, 1)\n    ax4.set_title('Per-Class Performance Metrics')\n    ax4.set_ylabel('Score')\n    ax4.legend(loc='upper center', bbox_to_anchor=(0.5, -0.15), ncol=4)\n    \n    plt.tight_layout()\n    plt.savefig('detailed_confusion_matrix.png')\n    plt.show()\n\n# Create the detailed confusion matrix visualization\nprint(\"\\nGenerating detailed confusion matrix visualization...\")\nplot_confusion_matrix_heatmap(test_metrics['confusion_matrix'], class_names)\nprint(\"Detailed confusion matrix visualization completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:13:36.016717Z","iopub.status.idle":"2025-04-23T22:13:36.016957Z","shell.execute_reply.started":"2025-04-23T22:13:36.016846Z","shell.execute_reply":"2025-04-23T22:13:36.016859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_cohen_kappa(all_labels, all_predictions):\n    \"\"\"\n    Calculate Cohen's Kappa coefficient, which measures inter-rater agreement.\n    \n    Args:\n        all_labels (array-like): True labels\n        all_predictions (array-like): Predicted labels\n        \n    Returns:\n        float: Cohen's Kappa score\n    \"\"\"\n    # Calculate Cohen's Kappa\n    kappa = cohen_kappa_score(all_labels, all_predictions)\n    \n    # Interpret Kappa value\n    if kappa < 0:\n        interpretation = \"Poor agreement (worse than random)\"\n    elif kappa < 0.2:\n        interpretation = \"Slight agreement\"\n    elif kappa < 0.4:\n        interpretation = \"Fair agreement\"\n    elif kappa < 0.6:\n        interpretation = \"Moderate agreement\"\n    elif kappa < 0.8:\n        interpretation = \"Substantial agreement\"\n    else:\n        interpretation = \"Almost perfect agreement\"\n    \n    print(f\"\\nCohen's Kappa: {kappa:.4f}\")\n    print(f\"Interpretation: {interpretation}\")\n    \n    # For quadratic weighted kappa (commonly used in DR grading)\n    weighted_kappa = cohen_kappa_score(all_labels, all_predictions, weights='quadratic')\n    print(f\"Quadratic Weighted Kappa: {weighted_kappa:.4f}\")\n    \n    return kappa, weighted_kappa\n\n# Calculate Cohen's Kappa from test_metrics\nall_labels = np.array([])\nall_predictions = np.array([])\n\n# Extract actual and predicted labels from confusion matrix\ncm = test_metrics['confusion_matrix']\nfor i in range(len(class_names)):\n    for j in range(len(class_names)):\n        count = cm[i, j]\n        all_labels = np.append(all_labels, np.full(int(count), i))\n        all_predictions = np.append(all_predictions, np.full(int(count), j))\n\n# Calculate and display Cohen's Kappa\nkappa, weighted_kappa = calculate_cohen_kappa(all_labels, all_predictions)\n\n# ====================================================================================================\n# SECTION 10: ENHANCED SUMMARY\n# ====================================================================================================\n\n# Print enhanced summary of results\nprint(\"\\nEnhanced Summary of Results:\")\nprint(\"=\" * 70)\nprint(f\"Model: MobileNetV2 with dataset fusion\")\nprint(f\"Model Size: {test_metrics['model_size_mb']:.2f} MB\")\nprint(f\"Test Accuracy: {test_metrics['accuracy']:.4f}\")\nprint(f\"Cohen's Kappa: {kappa:.4f}\")\nprint(f\"Quadratic Weighted Kappa: {weighted_kappa:.4f}\")\nprint(f\"Test Sensitivity (Recall): {test_metrics['sensitivity']:.4f}\")\nprint(f\"Test Specificity: {test_metrics['specificity']:.4f}\")\nprint(f\"Test Precision: {test_metrics['precision']:.4f}\")\nprint(f\"Test F1-Score: {test_metrics['f1']:.4f}\")\nprint(f\"AUC-ROC: {test_metrics['auc_roc']:.4f}\")\nprint(f\"Accuracy-to-Size Ratio (ASR): {test_metrics['asr']:.6f}\")\nprint(\"\\nPer-class Sensitivities:\")\nfor i, recall in enumerate(test_metrics['class_recalls']):\n    print(f\"  - {class_names[i]}: {recall:.4f}\")\nprint(f\"\\nGPU Inference Time: {inference_times['gpu_time']*1000 if inference_times['gpu_time'] else 'N/A':.2f} ms per image\")\nprint(f\"CPU Inference Time: {inference_times['cpu_time']*1000:.2f} ms per image\")\nprint(\"=\" * 70)\n\n# Visualize a comparison of all metrics\ndef plot_metrics_comparison():\n    \"\"\"\n    Create a visual comparison of all evaluation metrics in a single radar chart.\n    \"\"\"\n    metrics = {\n        'Accuracy': test_metrics['accuracy'],\n        'Sensitivity': test_metrics['sensitivity'],\n        'Specificity': test_metrics['specificity'],\n        'Precision': test_metrics['precision'],\n        'F1-Score': test_metrics['f1'],\n        'AUC-ROC': test_metrics['auc_roc'],\n        'Cohen\\'s Kappa': kappa,\n        'Weighted Kappa': weighted_kappa\n    }\n    \n    # Create radar chart\n    categories = list(metrics.keys())\n    values = list(metrics.values())\n    \n    # Calculate angle for each category\n    angles = np.linspace(0, 2*np.pi, len(categories), endpoint=False).tolist()\n    \n    # Complete the loop for the radar chart by appending the first value at the end\n    values.append(values[0])\n    angles.append(angles[0])\n    categories.append(categories[0])\n    \n    # Create radar chart\n    fig, ax = plt.subplots(figsize=(10, 10), subplot_kw=dict(polar=True))\n    \n    # Draw the chart\n    ax.plot(angles, values, 'o-', linewidth=2, label='Metrics')\n    ax.fill(angles, values, alpha=0.25)\n    \n    # Set category labels\n    ax.set_thetagrids(np.degrees(angles[:-1]), categories[:-1])\n    \n    # Set radial limits\n    ax.set_ylim(0, 1)\n    \n    # Add grid and labels\n    ax.grid(True)\n    plt.title('Model Performance Metrics', size=15)\n    \n    plt.tight_layout()\n    plt.savefig('metrics_radar_chart.png')\n    plt.show()\n\n# Generate metrics comparison visualization\nprint(\"\\nGenerating metrics comparison visualization...\")\nplot_metrics_comparison()\nprint(\"Metrics comparison visualization completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-23T22:13:36.018926Z","iopub.status.idle":"2025-04-23T22:13:36.019319Z","shell.execute_reply.started":"2025-04-23T22:13:36.019081Z","shell.execute_reply":"2025-04-23T22:13:36.019096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}