{"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":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\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\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\nfrom sklearn.preprocessing import label_binarize\nfrom sklearn.metrics import cohen_kappa_score\n\n# Set seeds for reproducibility\ndef seed_everything(seed=42):\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\n    torch.backends.cudnn.benchmark = False\n\n# Set seeds\nseed_everything()\n\n# Check if GPU is available\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\nimport seaborn as sns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T14:54:10.219363Z","iopub.execute_input":"2025-04-24T14:54:10.220125Z","iopub.status.idle":"2025-04-24T14:54:10.228658Z","shell.execute_reply.started":"2025-04-24T14:54:10.220098Z","shell.execute_reply":"2025-04-24T14:54:10.228002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set paths to datasets (reusing from baseline)\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}\")\n\n# Function to get file list from EyePACS folder structure (reused from baseline)\ndef get_eyepacs_files(root_path, subset='train'):\n    \"\"\"Get files from EyePACS folder structure where images are organized in class folders\"\"\"\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\"Path does not exist: {subset_path}\")\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    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 (reused from baseline)\ndef load_aptos_dataset(aptos_path):\n    \"\"\"Load APTOS dataset from CSV file\"\"\"\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    df_aptos = pd.read_csv(csv_path)\n    print(f\"Loaded APTOS 2019: {len(df_aptos)} images\")\n\n    # Add source and image path\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\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# Execute load_aptos_dataset\nprint(\"Executing load_aptos_dataset function...\")\naptos_df = load_aptos_dataset(aptos_path)\nprint(\"APTOS dataset loaded successfully.\")\n\n\ndef create_balanced_dataset(aptos_df, eyepacs_files, target_count_per_class=3000):\n    \"\"\"Create balanced dataset by combining APTOS and EyePACS images\"\"\"\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\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(\"\\nFinal class distribution:\")\n    print(balanced_df['class'].value_counts().sort_index())\n    \n    return balanced_df\n# Execute get_eyepacs_files\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# Execute 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-24T14:54:10.229922Z","iopub.execute_input":"2025-04-24T14:54:10.230111Z","iopub.status.idle":"2025-04-24T14:54:10.915008Z","shell.execute_reply.started":"2025-04-24T14:54:10.230097Z","shell.execute_reply":"2025-04-24T14:54:10.914335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Split dataset into train, validation, and test sets (reused from baseline)\ndef create_data_splits(balanced_df, test_size=0.15, val_size=0.15):\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']\n    )\n\n    # Recalculate validation size relative to remaining data\n    relative_val_size = val_size / (1 - test_size)\n\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']\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# Execute create_data_splits\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-24T14:54:10.915753Z","iopub.execute_input":"2025-04-24T14:54:10.915988Z","iopub.status.idle":"2025-04-24T14:54:10.936516Z","shell.execute_reply.started":"2025-04-24T14:54:10.915963Z","shell.execute_reply":"2025-04-24T14:54:10.935786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Data transformations for EfficientNet\n# EfficientNet-B0 uses 224x224 input size\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomAffine(\n        degrees=20,\n        translate=(0.1, 0.1),\n        scale=(0.8, 1.2),\n    ),\n    transforms.ColorJitter(brightness=(0.5, 1.0)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# Dataset class (reused from baseline)\nclass DiabeticRetinopathyDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        img_path = self.dataframe.iloc[idx]['image_path']\n        label = self.dataframe.iloc[idx]['class']\n\n        try:\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        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\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=32,  # Smaller batch size for EfficientNet\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=32,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=32,\n    shuffle=False,\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\")\n\n# class CustomEfficientNetB0(nn.Module):\n#     def __init__(self, num_classes=5):\n#         super(CustomEfficientNetB0, self).__init__()\n#         # Load pretrained EfficientNet-B0\n#         self.efficientnet = models.efficientnet_b0(pretrained=True)\n        \n#         # Get the number of features in the last layer\n#         in_features = self.efficientnet.classifier[1].in_features\n        \n#         # Replace classifier with a more compact head\n#         self.efficientnet.classifier = nn.Sequential(\n#             nn.Dropout(p=0.3),\n#             nn.Linear(in_features, 64),  # Reduced from original size\n#             nn.ReLU(),\n#             nn.Linear(64, num_classes)\n#         )\n    \n#     def forward(self, x):\n#         return self.efficientnet(x)\n\nclass CustomEfficientNetB0(nn.Module):\n    def __init__(self, num_classes=5):\n        super(CustomEfficientNetB0, self).__init__()\n        # Replacing EfficientNet with MobileNetV3Small\n        self.backbone = models.mobilenet_v3_small(pretrained=True)\n        # Replacing EfficientNet with shufflenet_v2_x0_5\n        # self.backbone = models.shufflenet_v2_x0_5(pretrained=True)\n        \n        # Get the number of features from the last layer\n        in_features = self.backbone.classifier[0].in_features #(changed to use sufflenet)\n        # in_features = self.backbone.fc.in_features\n        \n        # Create a more compact classifier\n        self.backbone.classifier = nn.Sequential(\n            nn.Linear(in_features, 32),  # Reduced from 64\n            nn.Hardswish(),  # More efficient than ReLU\n            nn.Dropout(p=0.2),\n            nn.Linear(32, num_classes)\n        )\n\n    #     self.backbone.fc = nn.Sequential( # classifier => fc\n    #         nn.Linear(in_features, 32),  # Reduced from 64\n    #         nn.Hardswish(),  # More efficient than ReLU\n    #         nn.Dropout(p=0.2),\n    #         nn.Linear(32, num_classes)\n    # )\n    \n    def forward(self, x):\n        return self.backbone(x)\n\n# Create teacher model (MobileNetV2 from baseline)\nclass TeacherModel(nn.Module):\n    def __init__(self, num_classes=5):\n        super(TeacherModel, self).__init__()\n        self.mobilenet = models.mobilenet_v2(pretrained=True)\n\n        self.mobilenet.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(self.mobilenet.last_channel, 32),\n            nn.ReLU(),\n            nn.Linear(32, 16),\n            nn.ReLU(),\n            nn.Linear(16, num_classes)\n        )\n\n    def forward(self, x):\n        return self.mobilenet(x)\n\n# Knowledge Distillation Loss Function\nclass DistillationLoss(nn.Module): # making alpha high and temperature\n    def __init__(self, alpha=0.7, temperature=4.0):\n        super(DistillationLoss, self).__init__()\n        self.alpha = alpha  # Weight for soft targets (teacher predictions)\n        self.temperature = temperature  # Temperature for softening probability distributions\n        self.ce_loss = nn.CrossEntropyLoss()\n        self.kl_loss = nn.KLDivLoss(reduction='batchmean')\n\n    def forward(self, student_outputs, teacher_outputs, targets):\n        # Hard target loss\n        hard_loss = self.ce_loss(student_outputs, targets)\n\n        # Soft target loss\n        soft_student = F.log_softmax(student_outputs / self.temperature, dim=1)\n        soft_teacher = F.softmax(teacher_outputs / self.temperature, dim=1)\n        soft_loss = self.kl_loss(soft_student, soft_teacher) * (self.temperature ** 2)\n\n        # Combined loss\n        loss = (1 - self.alpha) * hard_loss + self.alpha * soft_loss\n\n        return loss\n\n# Initialize models\nprint(\"\\nInitializing models...\")\n# Create the student model (EfficientNet-B0)\nstudent_model = CustomEfficientNetB0(num_classes=5).to(device)\n\n# Create the teacher model (from baseline)\nteacher_model = TeacherModel(num_classes=5).to(device)\n\n# Try to load the teacher model from saved weights\ntry:\n    teacher_model.load_state_dict(torch.load('best_dr_model.pth'))\n    print(\"Loaded teacher model weights successfully.\")\nexcept Exception as e:\n    print(f\"Error loading teacher model weights: {e}\")\n    print(\"Will use a pre-trained teacher model without fine-tuning.\")\n\n# Set teacher model to evaluation mode\nteacher_model.eval()\n\n# Define optimizer\noptimizer = optim.Adam(student_model.parameters(), lr=0.0005)  # Lower learning rate for EfficientNet\n\n# Learning rate scheduler\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode='min',\n    factor=0.5,\n    patience=5,\n    verbose=True\n)\n\n# Distillation loss\ndistillation_loss = DistillationLoss(alpha=0.7, temperature=4.0)  # Slightly higher alpha for EfficientNet\n\nprint(\"Models initialized successfully.\")\n\n\n# Training function with knowledge distillation\ndef train_model_with_distillation(student_model, teacher_model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=100, patience=10):\n    best_val_loss = float('inf')\n    counter = 0\n    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []}\n    start_time = time.time()\n    \n    best_model_path = 'best_efficientnet_dr_model.pth'\n    \n    for epoch in range(num_epochs):\n        epoch_start = time.time()\n        \n        # Training phase\n        student_model.train()\n        teacher_model.eval()  # Teacher is always in eval mode\n        \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            optimizer.zero_grad()\n            \n            # Get outputs from both models\n            student_outputs = student_model(inputs)\n            with torch.no_grad():\n                teacher_outputs = teacher_model(inputs)\n            \n            # Calculate distillation loss\n            loss = criterion(student_outputs, teacher_outputs, labels)\n            \n            loss.backward()\n            optimizer.step()\n            \n            running_loss += loss.item() * inputs.size(0)\n            _, predicted = torch.max(student_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            train_bar.set_postfix(loss=loss.item(), acc=correct/total if total > 0 else 0)\n        \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        # 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        # Validation phase\n        student_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        with torch.no_grad():\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                student_outputs = student_model(inputs)\n                \n                # For validation, we can use standard cross-entropy loss\n                loss = F.cross_entropy(student_outputs, labels)\n                \n                running_loss += loss.item() * inputs.size(0)\n                _, predicted = torch.max(student_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                    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                val_bar.set_postfix(loss=loss.item(), acc=correct/total if total > 0 else 0)\n        \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        # Print class-wise validation accuracy\n        print(f\"Val Loss: {epoch_val_loss:.4f}, Val Acc: {epoch_val_acc:.4f}\")\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        # Update learning rate based on validation loss\n        scheduler.step(epoch_val_loss)\n        \n        # Check if we should save the model\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(student_model.state_dict(), best_model_path)\n            counter = 0\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    student_model.load_state_dict(torch.load(best_model_path))\n    return student_model, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T14:54:10.938129Z","iopub.execute_input":"2025-04-24T14:54:10.938351Z","iopub.status.idle":"2025-04-24T14:54:11.148739Z","shell.execute_reply.started":"2025-04-24T14:54:10.938336Z","shell.execute_reply":"2025-04-24T14:54:11.148003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train the model\nprint(\"\\nStarting model training with knowledge distillation...\")\ntrained_student_model, history = train_model_with_distillation(\n    student_model,\n    teacher_model,\n    train_loader,\n    val_loader,\n    distillation_loss,\n    optimizer,\n    scheduler,\n    num_epochs=50,  # epochs 50\n    patience=20\n)\nprint(\"Model training with knowledge distillation completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T14:54:11.149577Z","iopub.execute_input":"2025-04-24T14:54:11.149853Z","iopub.status.idle":"2025-04-24T17:17:24.460755Z","shell.execute_reply.started":"2025-04-24T14:54:11.149827Z","shell.execute_reply":"2025-04-24T17:17:24.460015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot training history\nprint(\"\\nPlotting training history...\")\nplt.figure(figsize=(12, 5))\n\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\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('efficientnet_model_training_history.png')\nplt.show()\nprint(\"Training history plotted successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T17:17:24.461880Z","iopub.execute_input":"2025-04-24T17:17:24.462495Z","iopub.status.idle":"2025-04-24T17:17:24.989254Z","shell.execute_reply.started":"2025-04-24T17:17:24.462451Z","shell.execute_reply":"2025-04-24T17:17:24.988403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.utils.prune as prune  # for pruning\n\ndef prune_model(model, prune_amount=0.45):  # Increased from 0.35\n    \"\"\"Prune the model more aggressively with structured pruning\"\"\"\n    print(\"\\nPruning the model...\")\n    \n    # Step 1: Apply global unstructured pruning\n    parameters_to_prune = []\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d) or isinstance(module, nn.Linear):\n            parameters_to_prune.append((module, 'weight'))\n    \n    prune.global_unstructured(\n        parameters_to_prune,\n        pruning_method=prune.L1Unstructured,\n        amount=prune_amount,\n    )\n    \n    # Step 2: Apply structured pruning to Conv layers\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d) and module.out_channels > 8:\n            # Skip pruning the first layer and very small layers\n            if \"0.0\" not in name:\n                # Prune 30% of channels in convolutional layers\n                prune.ln_structured(module, name='weight', amount=0.3, n=2, dim=0)\n    \n    # Count the sparsity\n    zero_weights = 0\n    total_weights = 0\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d) or isinstance(module, nn.Linear):\n            zero_weights += torch.sum(module.weight == 0).item()\n            total_weights += module.weight.numel()\n    \n    sparsity = 100. * zero_weights / total_weights\n    print(f\"Model pruned with overall sparsity: {sparsity:.2f}%\")\n    \n    return model\n\n\n# Fine-tune after pruning\ndef fine_tune_pruned_model(model, train_loader, val_loader, epochs=10):  # Increased epochs\n    \"\"\"More extensive fine-tuning of the pruned model\"\"\"\n    print(\"\\nFine-tuning pruned model...\")\n\n    # Cosine annealing learning rate for better convergence\n    optimizer = optim.Adam(model.parameters(), lr=0.0002)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)\n    criterion = nn.CrossEntropyLoss()\n\n    best_val_loss = float('inf')\n    best_model_path = 'best_pruned_mobilenet_model.pth'\n\n    for epoch in range(epochs):\n        # Training phase\n        model.train()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n\n        train_bar = tqdm(train_loader, desc=f'Fine-tune Epoch {epoch+1}/{epochs} [Train]')\n        for inputs, labels in train_bar:\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            \n            # Gradient clipping to stabilize training\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            \n            optimizer.step()\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            train_bar.set_postfix(loss=loss.item(), acc=correct/total if total > 0 else 0)\n        \n        # Update learning rate\n        scheduler.step()\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        correct = 0\n        total = 0\n\n        with torch.no_grad():\n            for inputs, labels in val_loader:\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n                val_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        epoch_val_loss = val_loss / len(val_loader.dataset)\n        epoch_val_acc = correct / total\n\n        print(f\"Fine-tune Epoch {epoch+1}/{epochs}\")\n        print(f\"Val Loss: {epoch_val_loss:.4f}, Val Acc: {epoch_val_acc:.4f}\")\n\n        # Save best model\n        if epoch_val_loss < best_val_loss:\n            best_val_loss = epoch_val_loss\n            torch.save(model.state_dict(), best_model_path)\n\n    # Load best model\n    model.load_state_dict(torch.load(best_model_path))\n    print(\"Fine-tuning completed.\")\n\n    return model\n\n# Apply pruning and fine-tuning\nprint(\"\\nApplying pruning and fine-tuning...\")\npruned_model = prune_model(trained_student_model, prune_amount=0.25)  # Less aggressive pruning for EfficientNet\nfine_tuned_model = fine_tune_pruned_model(pruned_model, train_loader, val_loader, epochs=3)\nprint(\"Pruning and fine-tuning completed.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T17:17:24.990161Z","iopub.execute_input":"2025-04-24T17:17:24.990633Z","iopub.status.idle":"2025-04-24T17:26:00.530289Z","shell.execute_reply.started":"2025-04-24T17:17:24.990614Z","shell.execute_reply":"2025-04-24T17:26:00.529333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Make pruning permanent\ndef make_pruning_permanent(model):\n    \"\"\"Make the pruning permanent by removing the masks\"\"\"\n    print(\"\\nMaking pruning permanent...\")\n\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d) or isinstance(module, nn.Linear):\n            try:\n                torch.nn.utils.prune.remove(module, 'weight')\n            except:\n                print(f\"Could not remove pruning from {name}, skipping\")\n\n    print(\"Pruning made permanent.\")\n    return model\n\n# Make pruning permanent\nfinal_model = make_pruning_permanent(fine_tuned_model)\n\n# Apply quantization\ndef quantize_model(model):\n    \"\"\"Apply dynamic int8 quantization to the model\"\"\"\n    print(\"\\nQuantizing model with dynamic int8 quantization...\")\n    \n    # Move model to CPU for quantization\n    model = model.to('cpu').eval()\n    \n    # Apply dynamic quantization\n    quantized_model = torch.quantization.quantize_dynamic(\n        model,\n        {nn.Linear, nn.Conv2d, nn.BatchNorm2d},\n        dtype=torch.qint8\n    )\n    \n    print(\"Model quantized with int8 precision.\")\n    return quantized_model\n\n# Apply quantization\nprint(\"\\nApplying quantization...\")\nquantized_model = quantize_model(final_model)\nprint(\"Quantization applied.\")\n\n# Save the final model\ntorch.save(quantized_model.state_dict(), 'final_efficientnet_dr_model.pth')\nprint(\"Final EfficientNet model saved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T17:26:00.531461Z","iopub.execute_input":"2025-04-24T17:26:00.532244Z","iopub.status.idle":"2025-04-24T17:26:00.641329Z","shell.execute_reply.started":"2025-04-24T17:26:00.532212Z","shell.execute_reply":"2025-04-24T17:26:00.640675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Evaluation function\ndef evaluate_model(model, test_loader, criterion, class_names=None):\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    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 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('final_efficientnet_dr_model.pth') / (1024 * 1024)\n    asr = accuracy / model_size_mb\n\n    print(\"\\nEfficientNet Model 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    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    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T17:26:00.643075Z","iopub.execute_input":"2025-04-24T17:26:00.643264Z","iopub.status.idle":"2025-04-24T17:26:00.657204Z","shell.execute_reply.started":"2025-04-24T17:26:00.643250Z","shell.execute_reply":"2025-04-24T17:26:00.656377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save the final model after pruning and fine-tuning\ntorch.save(final_model.state_dict(), 'final_pruned_model.pth')\nprint(\"Final pruned model saved.\")\n\n# Apply quantization\nprint(\"\\nApplying quantization...\")\n# Move model to CPU for quantization\nquantized_model = final_model.to('cpu').eval()\n\n# Apply dynamic quantization to the model\nquantized_model = torch.quantization.quantize_dynamic(\n    quantized_model, \n    {nn.Linear, nn.Conv2d}, \n    dtype=torch.qint8\n)\n\n# Save the quantized model\ntorch.save(quantized_model.state_dict(), 'final_quantized_model.pth')\nprint(\"Quantized model saved.\")\n\n# IMPORTANT: Quantized models can only run on CPU, not on CUDA\nprint(\"\\nEvaluating the final model on CPU...\")\n# Define criterion for evaluation\ncriterion = nn.CrossEntropyLoss()\nclass_names = ['No DR', 'Mild DR', 'Moderate DR', 'Severe DR', 'Proliferative DR']\n\n\n\n\n# # For evaluation, we have two options:\n# # 1. Use the non-quantized model on GPU (faster but larger)\n# # 2. Use the quantized model on CPU (slower but smaller)\n# # Let's implement both to compare:\n\n# # Modified evaluation function to force CPU device\ndef evaluate_model_on_cpu(model, test_loader, criterion, class_names=None):\n    if class_names is None:\n        class_names = ['No DR', 'Mild DR', 'Moderate DR', 'Severe DR', 'Proliferative DR']\n    \n    model = model.to('cpu')  # Force CPU\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\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 on CPU'):\n            # Move tensors to CPU\n            inputs = inputs.to('cpu')\n            labels_cpu = labels.to('cpu')  # Store CPU version for evaluation\n            \n            outputs = model(inputs)\n            loss = criterion(outputs, labels_cpu)\n            \n            running_loss += loss.item() * inputs.size(0)\n            _, predicted = torch.max(outputs, 1)\n            total += labels_cpu.size(0)\n            correct += (predicted == labels_cpu).sum().item()\n            \n            all_labels.extend(labels_cpu.numpy())\n            all_predictions.extend(predicted.numpy())\n            all_outputs.extend(F.softmax(outputs, dim=1).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    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 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('final_quantized_model.pth') / (1024 * 1024)\n    asr = accuracy / model_size_mb\n    \n    print(\"\\nModel 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    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    }\n\n# Also evaluate the non-quantized model on GPU for comparison\ndef evaluate_model_on_gpu(model, test_loader, criterion, class_names=None):\n    if class_names is None:\n        class_names = ['No DR', 'Mild DR', 'Moderate DR', 'Severe DR', 'Proliferative DR']\n    \n    # Check if GPU is available\n    use_gpu = torch.cuda.is_available()\n    device = torch.device(\"cuda\" if use_gpu else \"cpu\")\n    print(f\"Evaluating non-quantized model on: {device}\")\n    \n    model = model.to(device)\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\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 on GPU'):\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            # Move to CPU for numpy conversion\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    # Follow the same metrics calculation as in the CPU function\n    cm = confusion_matrix(all_labels, all_predictions)\n    \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 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    # Calculate AUC-ROC\n    all_outputs = np.array(all_outputs)\n    all_labels = np.array(all_labels)\n    y_true_bin = label_binarize(all_labels, classes=range(5))\n    \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('final_pruned_model.pth') / (1024 * 1024)\n    asr = accuracy / model_size_mb\n    \n    class_recalls = recall_score(all_labels, all_predictions, average=None)\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    }\n\n# Run evaluation on both versions\nprint(\"\\nEvaluating quantized model on CPU...\")\ncpu_metrics = evaluate_model_on_cpu(quantized_model, test_loader, criterion, class_names)\n\nprint(\"\\nEvaluating original (non-quantized) model...\")\ngpu_metrics = evaluate_model_on_gpu(final_model, test_loader, criterion, class_names)\n\n# Measure inference time correctly for both models\ndef measure_inference_time(quant_model, non_quant_model, batch_size=1, image_size=(224, 224), n_runs=100):\n    # Create dummy inputs\n    dummy_input_cpu = torch.randn(batch_size, 3, *image_size)\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    dummy_input_gpu = torch.randn(batch_size, 3, *image_size).to(device)\n    \n    # Ensure models are on correct devices\n    quant_model = quant_model.to('cpu')\n    non_quant_model = non_quant_model.to(device)\n    \n    # 1. Measure quantized model on CPU\n    # Warm-up runs\n    for _ in range(10):\n        with torch.no_grad():\n            _ = quant_model(dummy_input_cpu)\n    \n    # Measure inference time\n    start = time.time()\n    for _ in range(n_runs):\n        with torch.no_grad():\n            _ = quant_model(dummy_input_cpu)\n    quant_cpu_time = (time.time() - start) / n_runs\n    print(f\"Quantized model CPU inference time: {quant_cpu_time*1000:.2f} ms per image\")\n    \n    # 2. Measure non-quantized model on GPU (if available)\n    if device.type == 'cuda':\n        # Warm-up runs\n        for _ in range(10):\n            with torch.no_grad():\n                _ = non_quant_model(dummy_input_gpu)\n        \n        torch.cuda.synchronize()\n        start = time.time()\n        for _ in range(n_runs):\n            with torch.no_grad():\n                _ = non_quant_model(dummy_input_gpu)\n            torch.cuda.synchronize()\n        non_quant_gpu_time = (time.time() - start) / n_runs\n        print(f\"Non-quantized model GPU inference time: {non_quant_gpu_time*1000:.2f} ms per image\")\n    else:\n        non_quant_gpu_time = float('nan')\n    \n    # 3. Measure non-quantized model on CPU\n    non_quant_model = non_quant_model.to('cpu')\n    \n    # Warm-up runs\n    for _ in range(10):\n        with torch.no_grad():\n            _ = non_quant_model(dummy_input_cpu)\n    \n    start = time.time()\n    for _ in range(n_runs):\n        with torch.no_grad():\n            _ = non_quant_model(dummy_input_cpu)\n    non_quant_cpu_time = (time.time() - start) / n_runs\n    print(f\"Non-quantized model CPU inference time: {non_quant_cpu_time*1000:.2f} ms per image\")\n    \n    return {\n        'quant_cpu_time': quant_cpu_time,\n        'non_quant_gpu_time': non_quant_gpu_time if device.type == 'cuda' else None,\n        'non_quant_cpu_time': non_quant_cpu_time,\n        'cpu_speedup': non_quant_cpu_time / quant_cpu_time\n    }\n\n# Measure inference times for both models\nprint(\"\\nMeasuring inference times...\")\ninference_times = measure_inference_time(quantized_model, final_model)\n\n# Print comprehensive summary\nprint(\"\\n\" + \"=\"*70)\nprint(\"COMPLETE MODEL COMPARISON\")\nprint(\"=\"*70)\nprint(f\"{'Metric':<25} | {'Non-Quantized':<15} | {'Quantized':<15}\")\nprint(\"-\"*70)\nprint(f\"{'Model Size (MB)':<25} | {gpu_metrics['model_size_mb']:<15.2f} | {cpu_metrics['model_size_mb']:<15.2f}\")\nprint(f\"{'Accuracy':<25} | {gpu_metrics['accuracy']:<15.4f} | {cpu_metrics['accuracy']:<15.4f}\")\nprint(f\"{'Sensitivity':<25} | {gpu_metrics['sensitivity']:<15.4f} | {cpu_metrics['sensitivity']:<15.4f}\")\nprint(f\"{'Specificity':<25} | {gpu_metrics['specificity']:<15.4f} | {cpu_metrics['specificity']:<15.4f}\")\nprint(f\"{'F1-Score':<25} | {gpu_metrics['f1']:<15.4f} | {cpu_metrics['f1']:<15.4f}\")\nprint(f\"{'ASR (Accuracy/Size)':<25} | {gpu_metrics['asr']:<15.6f} | {cpu_metrics['asr']:<15.6f}\")\n\n# Print inference times\nprint(\"-\"*70)\nprint(f\"{'CPU Inference (ms)':<25} | {inference_times['non_quant_cpu_time']*1000:<15.2f} | {inference_times['quant_cpu_time']*1000:<15.2f}\")\nif inference_times['non_quant_gpu_time'] is not None:\n    print(f\"{'GPU Inference (ms)':<25} | {inference_times['non_quant_gpu_time']*1000:<15.2f} | {'N/A':<15}\")\n# print(f\"{'CPU Speedup':<25} | {'1.00x':<15} | {inference_times['cpu_speedup']:<15.2f}x\")\nprint(\"=\"*70)\n\n# Visualize confusion matrix for the quantized model\ntry:\n    import seaborn as sns\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cpu_metrics['confusion_matrix'], annot=True, fmt='d', cmap='Blues', \n                xticklabels=class_names, yticklabels=class_names)\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.title('Confusion Matrix - Quantized Model')\n    plt.tight_layout()\n    plt.show()\nexcept ImportError:\n    print(\"Seaborn not available for confusion matrix visualization\")\n\nprint(\"\\nModel evaluation and analysis completed successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T17:26:00.658084Z","iopub.execute_input":"2025-04-24T17:26:00.658315Z","iopub.status.idle":"2025-04-24T17:27:07.174171Z","shell.execute_reply.started":"2025-04-24T17:26:00.658299Z","shell.execute_reply":"2025-04-24T17:27:07.173359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}