{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install timm\n!pip install git+https://github.com/jacobgil/pytorch-grad-cam.git\n\nimport pandas as pd\nimport os\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import transforms, datasets\nfrom torch.utils.data import DataLoader, random_split\nimport timm\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\nfrom sklearn.metrics import confusion_matrix, classification_report\nimport seaborn as sns\n\nprint(\"Setup Complete. All libraries are installed and imported.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:32:38.678072Z","iopub.execute_input":"2025-07-26T14:32:38.678347Z","iopub.status.idle":"2025-07-26T14:34:16.243456Z","shell.execute_reply.started":"2025-07-26T14:32:38.678323Z","shell.execute_reply":"2025-07-26T14:34:16.242616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"KAGGLE_INPUT_DIR = '/kaggle/input/aptos2019-blindness-detection'\nCSV_PATH = os.path.join(KAGGLE_INPUT_DIR, 'train.csv')\nSOURCE_IMAGES_FOLDER = os.path.join(KAGGLE_INPUT_DIR, 'train_images')\n\nOUTPUT_FOLDER = '/kaggle/working/processed_npdr_data'\n\nprint(\"Step 1: Reading and filtering the CSV file...\")\ndf = pd.read_csv(CSV_PATH)\n\ndf_filtered = df[df['diagnosis'].isin([1, 2, 3])].copy()\nprint(f\"Found {len(df_filtered)} images for NPDR stages 1, 2, and 3.\")\n\nprint(\"\\nStep 2: Creating new directories for each class...\")\nos.makedirs(OUTPUT_FOLDER, exist_ok=True)\nfor label in ['1', '2', '3']:\n    class_dir = os.path.join(OUTPUT_FOLDER, str(label))\n    os.makedirs(class_dir, exist_ok=True)\n    print(f\"- Directory created: {class_dir}\")\n\nprint(\"\\nStep 3: Resizing and saving images to their new folders...\")\nfor index, row in tqdm(df_filtered.iterrows(), total=df_filtered.shape[0], desc=\"Processing Images\"):\n    image_name = f\"{row['id_code']}.png\"\n    label = str(row['diagnosis'])\n    source_path = os.path.join(SOURCE_IMAGES_FOLDER, image_name)\n    destination_path = os.path.join(OUTPUT_FOLDER, label, image_name)\n    \n    if os.path.exists(source_path):\n        try:\n            with Image.open(source_path) as img:\n                # [cite_start]Resize the image to 224x224 pixels [cite: 59]\n                img_resized = img.resize((224, 224))\n                img_resized.save(destination_path)\n        except Exception as e:\n            print(f\"\\nCould not process {image_name}. Error: {e}\")\n\nprint(\"\\nData preparation is complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:34:16.244891Z","iopub.execute_input":"2025-07-26T14:34:16.245387Z","iopub.status.idle":"2025-07-26T14:40:01.692090Z","shell.execute_reply.started":"2025-07-26T14:34:16.245361Z","shell.execute_reply":"2025-07-26T14:40:01.691296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_transforms = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nprint(\"1. Loading dataset from the processed folder...\")\nfull_dataset = datasets.ImageFolder(OUTPUT_FOLDER, transform=data_transforms)\nprint(f\"Dataset loaded. Found {len(full_dataset)} images in {len(full_dataset.classes)} classes.\")\n\ntrain_size = int(0.8 * len(full_dataset))\nval_size = len(full_dataset) - train_size\ntrain_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)\nprint(f\"\\n2. Data split into {len(train_dataset)} training and {len(val_dataset)} validation images.\")\n\nprint(\"\\n3. Loading pre-trained Vision Transformer model (vit_base_patch16_224)...\")\nmodel = timm.create_model('vit_base_patch16_224', pretrained=True)\nnum_classes = 3\nmodel.head = nn.Linear(model.head.in_features, num_classes)\nprint(\"Model modified for 3 classes.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:40:01.692916Z","iopub.execute_input":"2025-07-26T14:40:01.693675Z","iopub.status.idle":"2025-07-26T14:40:04.833333Z","shell.execute_reply.started":"2025-07-26T14:40:01.693653Z","shell.execute_reply":"2025-07-26T14:40:04.832695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_EPOCHS = 15\nLEARNING_RATE = 3e-5\nMODEL_SAVE_PATH = '/kaggle/working/aptos_vit_model.pth'\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\nmodel.to(device)\n\n# Optimizer and scheduler (fine-tuning ViT works better with lower LR)\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=NUM_EPOCHS)\n\n# Handle class imbalance\nclass_counts = torch.bincount(torch.tensor(full_dataset.targets))\nclass_weights = 1. / class_counts.float()\nclass_weights = class_weights / class_weights.sum()\nclass_weights = class_weights.to(device)\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\nprint(\"Using weighted loss to handle class imbalance.\")\nprint(f\"Class Weights: {class_weights.cpu().numpy()}\")\n\nbest_val_acc = 0.0\nprint(\"\\nStarting training...\\n\")\n\nfor epoch in range(NUM_EPOCHS):\n    # --- Training ---\n    model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n\n    for inputs, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{NUM_EPOCHS} [Training]\"):\n        inputs, labels = inputs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * inputs.size(0)\n        _, preds = torch.max(outputs, 1)\n        correct_train += (preds == labels).sum().item()\n        total_train += labels.size(0)\n\n    train_loss = running_loss / len(train_loader.dataset)\n    train_acc = 100 * correct_train / total_train\n\n    # --- Validation ---\n    model.eval()\n    running_vloss = 0.0\n    correct_val = 0\n    total_val = 0\n    with torch.no_grad():\n        for inputs, labels in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{NUM_EPOCHS} [Validation]\"):\n            inputs, labels = inputs.to(device), labels.to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            running_vloss += loss.item() * inputs.size(0)\n            _, predicted = torch.max(outputs, 1)\n            correct_val += (predicted == labels).sum().item()\n            total_val += labels.size(0)\n\n    val_loss = running_vloss / len(val_loader.dataset)\n    val_acc = 100 * correct_val / total_val\n\n    scheduler.step()  # 🔁 Adjust learning rate\n\n    print(f\"Epoch {epoch+1}/{NUM_EPOCHS} | Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% | Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), MODEL_SAVE_PATH)\n        print(f\"✅ Model saved to {MODEL_SAVE_PATH} (Val Acc: {val_acc:.2f}%)\")\n\nprint(f\"\\n🎉 Training complete! Best validation accuracy: {best_val_acc:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:40:04.835003Z","iopub.execute_input":"2025-07-26T14:40:04.835309Z","iopub.status.idle":"2025-07-26T14:46:42.099472Z","shell.execute_reply.started":"2025-07-26T14:40:04.835289Z","shell.execute_reply":"2025-07-26T14:46:42.098655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 1. Load the Best Model ---\nmodel.load_state_dict(torch.load(MODEL_SAVE_PATH))\nmodel.to(device)\nmodel.eval()\n\n# --- 2. Get Predictions ---\nall_preds = []\nall_labels = []\nwith torch.no_grad():\n    for inputs, labels in val_loader:\n        inputs, labels = inputs.to(device), labels.to(device)\n        outputs = model(inputs)\n        _, predicted = torch.max(outputs, 1)\n        all_preds.extend(predicted.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n# --- 3. Classification Report ---\nclass_names = ['Mild', 'Moderate', 'Severe'] # Assuming 0->Mild, 1->Moderate, 2->Severe\n# The class indices might be different, let's get them from the dataset object\nclass_names_from_dataset = [full_dataset.classes[i] for i in sorted(full_dataset.class_to_idx.values())]\n# Assuming the folder names '1', '2', '3' map to Mild, Moderate, Severe\nclass_map = {'1': 'Mild', '2': 'Moderate', '3': 'Severe'}\nclass_names = [class_map[c] for c in class_names_from_dataset]\n\n\nprint(\"Classification Report:\\n\")\nprint(classification_report(all_labels, all_preds, target_names=class_names))\n\n# --- 4. Confusion Matrix ---\ncm = confusion_matrix(all_labels, all_preds)\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names)\nplt.xlabel('Predicted Label')\nplt.ylabel('True Label')\nplt.title('Confusion Matrix')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:46:42.100516Z","iopub.execute_input":"2025-07-26T14:46:42.100908Z","iopub.status.idle":"2025-07-26T14:46:44.831221Z","shell.execute_reply.started":"2025-07-26T14:46:42.100877Z","shell.execute_reply":"2025-07-26T14:46:44.830446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 1. Define a Reshape Transform for Vision Transformers ---\ndef reshape_transform(tensor, height=14, width=14):\n    # Discard the [CLS] token\n    result = tensor[:, 1:, :].reshape(tensor.size(0), height, width, tensor.size(2))\n    # Bring the channels to the front\n    result = result.permute(0, 3, 1, 2)\n    return result\n\n# --- 2. Define Grad-CAM Function with the Transform ---\ntarget_layers = [model.blocks[-1].norm1]\ncam = GradCAM(model=model, target_layers=target_layers, reshape_transform=reshape_transform)\n\ndef generate_heatmap(image_path):\n    pil_img = Image.open(image_path).convert('RGB')\n    torch_img = data_transforms(pil_img).unsqueeze(0).to(device)\n    rgb_img = np.array(pil_img.resize((224, 224))) / 255.0\n    \n    grayscale_cam = cam(input_tensor=torch_img, targets=None)[0, :]\n    visualization = show_cam_on_image(rgb_img, grayscale_cam, use_rgb=True)\n    return visualization, pil_img\n\n# --- 3. Visualize Results on Sample Images ---\nsample_image_paths = [val_dataset.dataset.samples[i][0] for i in val_dataset.indices[:5]]\n\nfor img_path in sample_image_paths:\n    \n    with torch.no_grad():\n        pil_img = Image.open(img_path).convert('RGB')\n        input_tensor = data_transforms(pil_img).unsqueeze(0).to(device)\n        output = model(input_tensor)\n        _, prediction_idx = torch.max(output, 1)\n        predicted_class_name = class_names[prediction_idx.item()]\n\n   \n    heatmap, original_image = generate_heatmap(img_path)\n    \n   \n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 5))\n    \n    ax1.imshow(original_image)\n    ax1.set_title(f'Original Image\\nTrue Label: {class_map[os.path.basename(os.path.dirname(img_path))]}')\n    ax1.axis('off')\n    \n    ax2.imshow(heatmap)\n    ax2.set_title(f'Grad-CAM Heatmap\\nPrediction: {predicted_class_name}')\n    ax2.axis('off')\n    \n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:46:44.832406Z","iopub.execute_input":"2025-07-26T14:46:44.832943Z","iopub.status.idle":"2025-07-26T14:46:46.388301Z","shell.execute_reply.started":"2025-07-26T14:46:44.832916Z","shell.execute_reply":"2025-07-26T14:46:46.387619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\nimport seaborn as sns\n\n# --- 1. Reshape Transform for Vision Transformers ---\ndef reshape_transform(tensor, height=14, width=14):\n    result = tensor[:, 1:, :].reshape(tensor.size(0), height, width, tensor.size(2))\n    result = result.permute(0, 3, 1, 2)\n    return result\n\n# --- 2. Grad-CAM Setup ---\ntarget_layers = [model.blocks[-1].norm1]\ncam = GradCAM(model=model, target_layers=target_layers, reshape_transform=reshape_transform)\n\ndef generate_heatmap(image_path):\n    pil_img = Image.open(image_path).convert('RGB')\n    torch_img = data_transforms(pil_img).unsqueeze(0).to(device)\n    rgb_img = np.array(pil_img.resize((224, 224))) / 255.0\n    \n    grayscale_cam = cam(input_tensor=torch_img, targets=None)[0, :]\n    visualization = show_cam_on_image(rgb_img, grayscale_cam, use_rgb=True)\n    return visualization, pil_img\n\n# --- 3. Function to Get Top-K Class Probabilities ---\ndef get_topk_predictions(image, model, class_names, k=5):\n    input_tensor = data_transforms(image).unsqueeze(0).to(device)\n    with torch.no_grad():\n        output = model(input_tensor)\n        probs = F.softmax(output, dim=1).squeeze().cpu().numpy()\n    \n    topk_indices = probs.argsort()[-k:][::-1]\n    topk_probs = probs[topk_indices]\n    topk_labels = [class_names[i] for i in topk_indices]\n    return topk_labels, topk_probs\n\n# --- 4. Visualize Results ---\nsample_image_paths = [val_dataset.dataset.samples[i][0] for i in val_dataset.indices[:5]]\n\nfor img_path in sample_image_paths:\n    pil_img = Image.open(img_path).convert('RGB')\n    \n    # Get model prediction and class probabilities\n    input_tensor = data_transforms(pil_img).unsqueeze(0).to(device)\n    with torch.no_grad():\n        output = model(input_tensor)\n        _, prediction_idx = torch.max(output, 1)\n        predicted_class_name = class_names[prediction_idx.item()]\n    \n    topk_labels, topk_probs = get_topk_predictions(pil_img, model, class_names, k=5)\n    \n    # Grad-CAM Heatmap\n    heatmap, original_image = generate_heatmap(img_path)\n    \n    # --- Plot All Together ---\n    fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(18, 5))\n    \n    # Original Image\n    ax1.imshow(original_image)\n    ax1.set_title(f'Original Image\\nTrue Label: {class_map[os.path.basename(os.path.dirname(img_path))]}')\n    ax1.axis('off')\n    \n    # Grad-CAM Heatmap\n    ax2.imshow(heatmap)\n    ax2.set_title(f'Grad-CAM Heatmap\\nPrediction: {predicted_class_name}')\n    ax2.axis('off')\n    \n    # Confidence Bar Plot\n    sns.barplot(x=topk_probs, y=topk_labels, palette='mako', ax=ax3)\n    ax3.set_title('Top-5 Class Probabilities')\n    ax3.set_xlabel('Confidence')\n    ax3.set_xlim(0, 1)\n    for i, v in enumerate(topk_probs):\n        ax3.text(v + 0.01, i, f\"{v*100:.1f}%\", va='center', fontsize=9)\n    \n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:46:46.389449Z","iopub.execute_input":"2025-07-26T14:46:46.389815Z","iopub.status.idle":"2025-07-26T14:46:49.094737Z","shell.execute_reply.started":"2025-07-26T14:46:46.389794Z","shell.execute_reply":"2025-07-26T14:46:49.094017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Labels for radar chart\nlabels = ['Explainability', 'Progression Tracking', 'Accuracy Focus', 'Monitoring Use', 'Modern Architecture']\nnum_vars = len(labels)\n\n# Data for each model\nvit_gradcam = [5, 5, 4, 5, 5]\ncnn_resnet = [2, 1, 5, 1, 3]\nquellec2017 = [3, 1, 4, 1, 3]\nlam2018 = [2, 1, 4, 1, 4]\n\n# Model names and colors\nmodels = {\n    \"ViT + Grad-CAM\": vit_gradcam,\n    \"CNN (ResNet)\": cnn_resnet,\n    \"Quellec et al.\": quellec2017,\n    \"Lam et al.\": lam2018\n}\ncolors = ['#636EFA', '#EF553B', '#00CC96', '#AB63FA']\n\n# Radar chart setup\nangles = np.linspace(0, 2 * np.pi, num_vars, endpoint=False).tolist()\nangles += angles[:1]  # complete the loop\n\nfig, ax = plt.subplots(figsize=(7, 7), subplot_kw=dict(polar=True))\nfor (name, values), color in zip(models.items(), colors):\n    values += values[:1]\n    ax.plot(angles, values, color=color, linewidth=2, label=name)\n    ax.fill(angles, values, color=color, alpha=0.25)\n\nax.set_theta_offset(np.pi / 2)\nax.set_theta_direction(-1)\nax.set_thetagrids(np.degrees(angles[:-1]), labels)\nax.set_title('Model Comparison Radar Chart', size=16, pad=20)\nax.set_rlim(0, 5)\nax.legend(loc='upper right', bbox_to_anchor=(1.3, 1.1))\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T14:46:49.095760Z","iopub.execute_input":"2025-07-26T14:46:49.096686Z","iopub.status.idle":"2025-07-26T14:46:49.398876Z","shell.execute_reply.started":"2025-07-26T14:46:49.096664Z","shell.execute_reply":"2025-07-26T14:46:49.398070Z"}},"outputs":[],"execution_count":null}]}