{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594},{"sourceType":"datasetVersion","sourceId":6717213,"datasetId":1317048,"databundleVersionId":6801677},{"sourceType":"datasetVersion","sourceId":11028374,"datasetId":6868079,"databundleVersionId":11411117},{"sourceType":"datasetVersion","sourceId":11026852,"datasetId":6866921,"databundleVersionId":11409423},{"sourceType":"datasetVersion","sourceId":10269174,"datasetId":6353528,"databundleVersionId":10568641},{"sourceType":"datasetVersion","sourceId":9864328,"datasetId":6054580,"databundleVersionId":10116567},{"sourceType":"datasetVersion","sourceId":11202373,"datasetId":6994345,"databundleVersionId":11607168}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Robust Multi-Label Chest X-Ray Diagnosis Using a Comprehensive Multi-Source Dataset\n\n## Overview\n\nThe dataset used in this project is a **comprehensive multi-source chest X-ray dataset**, curated from multiple publicly available sources. After extensive cleaning and filtering, the final dataset comprises **317,871 images** covering a wide range of thoracic conditions, including **pneumonia, tuberculosis, cardiomegaly, COVID-19, and other pulmonary abnormalities**. The dataset has been **carefully curated, cleaned, and preprocessed** to ensure high reliability and consistency, making it suitable for deep learning applications in multi-label classification.\n\nGiven computational resource constraints, a **reduced subset** of these 374,718 images was selected for training while maintaining the **original class distribution** to ensure model generalizability.\n\n## Project Goals and Objectives\n\nThe goal of this project is to develop a robust AI-driven diagnostic tool for chest X-ray interpretation. Specifically, the project aims to:\n\n- **Automate the diagnosis** of multiple thoracic conditions using a multi-label classification approach.\n- **Enhance early detection** and treatment planning for conditions like pneumonia and tuberculosis, while also accurately identifying other significant abnormalities such as cardiomegaly and COVID-19.\n- **Leverage expert-annotated data** to build a reliable model that generalizes well across diverse patient populations and imaging conditions.\n- **Ensure model interpretability** through explainability techniques (e.g., SHAP, Grad-CAM) and robustness via out-of-distribution (OOD) detection.\n\n## Data Sources and Their Contributions\n\nThe dataset integrates images from multiple high-quality sources, ensuring diverse imaging conditions, patient demographics, and pathological variations. Below is a summary of the data sources and their label contributions:\n\n| **Source**        | **Images** | **Key Labels**                                                                                                                          |\n| ----------------- | ---------- | --------------------------------------------------------------------------------------------------------------------------------------- |\n| **CheXpert**      | 154,054    | Normal, Pneumonia, Lung Opacity, Cardiomegaly, Edema, Consolidation, Atelectasis, Pneumothorax, Pleural Effusion, Support Devices, etc. |\n| **VIN Big Data**  | 67,914     | Normal, Cardiomegaly, Lung Opacity, Pleural Effusion, Aortic Enlargement, Pulmonary Fibrosis, Nodule/Mass, Other Lesions, etc.          |\n| **BIMCV**         | 49,687     | COVID-19                                                                                                                                |\n| **RSNA**          | 18,406     | Normal, Pneumonia, Lung Opacity                                                                                                         |\n| **Stony Brook**   | 13,636     | COVID-19                                                                                                                                |\n| **OCT**           | 5,856      | Normal, Pneumonia                                                                                                                       |\n| **NIAID**         | 3,499      | Tuberculosis                                                                                                                            |\n| **Pavan**         | 1,357      | Non-XRay                                                                                                                                       |\n| **RICORD**        | 1,096      | COVID-19                                                                                                                                |\n| **SIRM**          | 943        | COVID-19                                                                                                                                |\n| **Shenzhen**      | 662        | Normal, Tuberculosis                                                                                                                    |\n| **Belarus**       | 304        | Tuberculosis                                                                                                                            |\n| **Cohen**         | 270        | COVID-19                                                                                                                                |\n| **MontgomerySet** | 138        | Normal, Tuberculosis                                                                                                                    |\n| **ACTmed**        | 25         | COVID-19                                                                                                                                |\n| **FIG1**          | 24         | COVID-19                                                                                                                                |\n                                                                                                                  \n## Detailed Source Descriptions\n\n- **CheXpert**:  \n  Originating from Stanford Hospital, CheXpert is one of the largest chest X-ray datasets available. It covers 14 common thoracic conditions and uses a combination of automated NLP-based annotation and manual validation by board-certified radiologists, ensuring a high level of reliability.\n\n- **VIN Big Data**:  \n  Sourced from Vietnamese hospitals, this dataset focuses on a diverse range of pulmonary and cardiovascular abnormalities. **Annotations were manually performed by radiologists**, which significantly boosts the credibility of the labels and ensures precise identification of subtle abnormalities.\n\n- **BIMCV**:  \n  Collected from Spanish hospitals, BIMCV primarily contributes COVID-19 cases, offering a rich set of examples for pandemic-related diagnosis along with detailed metadata.\n\n- **RSNA**:  \n  Developed as part of the RSNA Pneumonia Detection Challenge, this dataset is meticulously annotated by expert radiologists, focusing on pneumonia and lung opacity. Its ground-truth labels are highly reliable for clinical applications.\n\n- **Stony Brook**:  \n  A dataset from Stony Brook University Hospital that adds valuable COVID-19 cases, enhancing the dataset with geographically diverse imaging.\n\n- **OCT**:  \n  Although primarily known for ophthalmic imaging, this dataset contributes a subset of chest X-rays (normal and pneumonia cases), further diversifying the overall image distribution.\n\n- **NIAID, Belarus, Shenzhen, MontgomerySet**:  \n  These datasets focus on tuberculosis detection and are sourced from various public health research initiatives and hospitals. They provide high diagnostic accuracy for TB through rigorous annotation protocols, making them vital for global TB screening research.\n\n- **RICORD, SIRM, Cohen, ACTmed, FIG1**:  \n  These sources contribute COVID-19 cases, collected from multiple hospitals and research centers worldwide, ensuring that the dataset captures a wide variety of COVID-19 manifestations.\n\n- **Pavan**:  \n  This source includes non-X-ray images such as humans, horses, flowers, and other miscellaneous objects, making it distinct from the other medical imaging datasets.\n\n## Annotation and Reliability\n\nThe dataset sources include images labeled through various annotation methods, ensuring overall reliability:\n\n- **CheXpert**:  \n  Labels were generated using an automated NLP-based approach, followed by manual expert validation by board-certified radiologists, enhancing accuracy and consistency.\n\n- **RSNA**:  \n  Annotations were performed directly by expert radiologists, providing highly reliable pneumonia detection labels.\n\n- **VIN Big Data**:  \n  **Annotations were manually performed by radiologists**, which significantly contributes to its reliability and accuracy.\n\n- **COVID-19 Cases (BIMCV, RICORD, SIRM, Stony Brook, Cohen, ACTmed, FIG1)**:  \n  Collected from multiple reputable centers, these datasets offer broad representation and high reliability of COVID-19 manifestations.\n\n- **Tuberculosis Cases (NIAID, Belarus, Shenzhen, MontgomerySet)**:  \n  Sourced from public health research initiatives, these datasets provide high diagnostic accuracy for TB through stringent annotation protocols.\n\n## Disease Prevalence and Descriptions\n\nBelow is the list of diseases with their prevalence (image counts) in the dataset, along with a brief description and typical radiographic features:\n\n| **Disease**                    | **Count** | **Description & Radiographic Features**                                                                                                               |\n| ------------------------------ | --------- | -------------------------------------------------------------------------------------------------------------------------------------------------------- |\n| **Support_Devices**            | 95,328    | Medical implants such as pacemakers, tubes, and other devices visible on the X-ray; essential for assessing patient management.                         |\n| **Lung_Opacity**               | 89,578    | Areas of increased lung density; can indicate infection, edema, or inflammation. Often seen as diffuse or localized opacities.                          |\n| **Pleural_Effusion**           | 73,378    | Fluid accumulation in the pleural space; appears as blunting of costophrenic angles and a meniscus sign.                                                |\n| **Covid**                      | 65,681    | Chest X-rays in COVID-19 typically show bilateral, peripheral ground-glass opacities, sometimes with consolidation in severe cases.                      |\n| **Normal**                     | 59,575    | X-rays without any significant pathological findings; serve as controls for comparison.                                                                 |\n| **Edema**                      | 47,033    | Fluid accumulation in lung tissue; appears as diffuse opacities, often bilateral, due to congestive heart failure or other causes.                      |\n| **Atelectasis**                | 28,762    | Partial or complete collapse of lung tissue; presents as areas of increased opacity with volume loss and possible mediastinal shift.                     |\n| **Cardiomegaly**               | 27,058    | Enlarged cardiac silhouette; often quantified by an increased cardiothoracic ratio (greater than 0.5).                                                  |\n| **Pneumonia**                  | 18,191    | Infection of lung tissue; can present as focal or diffuse consolidation with or without air bronchograms.                                                |\n| **Pneumothorax**               | 17,233    | Presence of air in the pleural space causing lung collapse; appears as a distinct line with absence of lung markings beyond it.                          |\n| **Consolidation**              | 12,379    | Dense lung opacification due to the presence of inflammatory exudates; commonly seen in bacterial pneumonia.                                            |\n| **Enlarged_Cardiomediastinum** | 8,504     | Widening of the mediastinum; may indicate a variety of conditions including vascular anomalies or neoplasms.                                             |\n| **Aortic_Enlargement**         | 7,162     | Dilation of the aorta; may be indicative of aneurysm or other vascular conditions.                                                                       |\n| **Pleural_Other**              | 7,060     | Other pleural abnormalities not categorized under effusion; may include thickening or scarring.                                                          |\n| **Fracture**                   | 6,789     | Breaks or cracks in the rib cage or other bony structures; critical for trauma assessment.                                                               |\n| **Lung_Lesion**                | 6,345     | Focal abnormality within the lung parenchyma; could be benign or malignant, requiring further investigation.                                             |\n| **Pulmonary_Fibrosis**         | 4,655     | Chronic scarring of lung tissue; appears as reticular opacities with volume loss, often in a basal distribution.                                         |\n| **Tuberculosis**               | 4,197     | Infectious disease caused by *Mycobacterium tuberculosis*; typically shows upper lobe consolidation, cavitary lesions, and calcified nodules.            |\n| **Nodule_Mass**                | 2,580     | Small, rounded opacities that could indicate benign nodules or malignancies; size and shape often determine further investigation.                       |\n| **Other_Lesion**               | 2,203     | Lesions not otherwise classified; may include atypical or mixed-pattern abnormalities requiring additional imaging or clinical correlation.               |\n| **Non_XRay**                   | 1,357     | Images that do not belong to the X-ray category, including photographs of humans, animals, and objects.                                                  |\n| **Infiltration**               | 1,247     | Diffuse or patchy increases in lung density, often indicative of infection or inflammatory processes.                                                   |\n| **ILD (Interstitial Lung Disease)** | 1,000 | A group of lung diseases causing progressive scarring of lung tissue; presents with reticular opacities and ground-glass changes.                        |\n| **Calcification**              | 960       | Deposition of calcium in lung tissues or lymph nodes; may be seen in healed infections or granulomatous diseases.                                        |\n\n> **Note**: The disease prevalence in the dataset reflects how the data was collected from various sources. This distribution may not be representative of real-world case prevalence but is rather driven by the characteristics of the datasets included.\n## Labeling and Observations\n\nEach image is annotated using a **multi-label format**, where:  \n- **1** indicates the presence of a condition.  \n- **0** indicates the absence of a condition.  \n- **0.5** denotes uncertainty in the diagnosis.  \n\nA single image may have multiple associated conditions. The dataset covers the full spectrum of diseases listed above, ensuring comprehensive coverage of thoracic abnormalities.\n\n## Preprocessing and Standardization\n\nTo ensure uniformity and high-quality data, several preprocessing steps were applied:\n\n1. **View Selection**:  \n   Only **Frontal PA and AP** views were retained, as these are the most clinically relevant. Currently, the model is being trained using only **Frontal cases** (PA and AP views).\n\n2. **Missing Value Handling**:  \n   All `NaN` values were converted to `0` (absence of findings).\n\n3. **Filtering Criteria**:  \n   - Images with **all `0` labels**, except for the **Normal** category, were removed to ensure clinical significance.  \n   - Images with **all `-1` labels** (fully uncertain annotations) were removed.  \n\n4. **Uncertain Label Handling**:  \n   In the **CheXpert dataset**, labels marked as `-1` (uncertain) were converted to `0.5`. This preserves ambiguity rather than assuming presence or absence.\n\n## Dataset Characteristics\n\n- **Multi-label Classification**:  \n  The dataset supports images having multiple conditions simultaneously.\n\n- **Diverse Patient Demographics**:  \n  Aggregating data from multiple institutions allows better generalization across different populations.\n\n- **One-Hot Encoding**:  \n  The labeling structure ensures that deep learning models can handle multi-label outputs effectively.\n\n- **Clinical Relevance**:  \n  The included conditions reflect real-world diagnostic challenges, making the dataset highly applicable for medical AI applications.\n\n- **Resource-Aware Training Subset**:  \n  A subset of the dataset was used for model training due to hardware constraints, while maintaining representative label distributions.\n\n## Intended Use and Applications\n\nThis dataset is designed to support AI-driven diagnostic applications and research in chest X-ray interpretation. Key applications include:\n\n- **Deep Learning-Based Disease Detection**:  \n  Training **CNNs and Transformer models** (e.g., MAX-ViT, EfficientNetV2-M) for thoracic disease classification.\n\n- **Medical Image Explainability**:  \n  Using **Grad-CAM ++, Bounding Boxes and Out Of Order Distribution** to improve model interpretability.\n\n- **Out-of-Distribution (OOD) Detection**:  \n  Enhancing model robustness by identifying abnormal cases not present in the training distribution.\n\n## Conclusion\n\nThis dataset serves as a **high-quality, multi-source collection** for advancing automated chest X-ray analysis. With **expert-annotated images, extensive preprocessing, and multi-label classification capabilities**, it provides a strong foundation for developing AI models in healthcare. The **diverse data sources and clinically relevant labeling**—including **manually annotated VIN Big Data and tuberculosis datasets, as well as radiologist-verified cases in CheXpert and RSNA**—ensure robust and reliable AI-driven diagnostic tools for medical imaging.\n","metadata":{"_kg_hide-output":true}},{"cell_type":"markdown","source":"# Data Loading & Pre-Processing","metadata":{}},{"cell_type":"markdown","source":"### Importing Basic Libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing","metadata":{"_uuid":"af1bcddd-213b-4816-ae2d-085ad4baf60b","_cell_guid":"d6dda341-bbaa-4d54-b572-4ba5ba27e22f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-04-15T03:52:59.346351Z","iopub.execute_input":"2025-04-15T03:52:59.346693Z","iopub.status.idle":"2025-04-15T03:52:59.350542Z","shell.execute_reply.started":"2025-04-15T03:52:59.346669Z","shell.execute_reply":"2025-04-15T03:52:59.349508Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Reading Metadata for Images","metadata":{}},{"cell_type":"code","source":"data = pd.read_csv('/kaggle/input/cxr-metadata-with-uncertain-labels/CXR Metadata With Uncertain Labels.csv')\ndata.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:01.259972Z","iopub.execute_input":"2025-04-15T03:53:01.260264Z","iopub.status.idle":"2025-04-15T03:53:02.223073Z","shell.execute_reply.started":"2025-04-15T03:53:01.260241Z","shell.execute_reply":"2025-04-15T03:53:02.222245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:02.224106Z","iopub.execute_input":"2025-04-15T03:53:02.22441Z","iopub.status.idle":"2025-04-15T03:53:02.228931Z","shell.execute_reply.started":"2025-04-15T03:53:02.224388Z","shell.execute_reply":"2025-04-15T03:53:02.228125Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Initially we have 3,17,871 Images For Our Use Case","metadata":{}},{"cell_type":"code","source":"data.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:02.230245Z","iopub.execute_input":"2025-04-15T03:53:02.230557Z","iopub.status.idle":"2025-04-15T03:53:02.331058Z","shell.execute_reply.started":"2025-04-15T03:53:02.230534Z","shell.execute_reply":"2025-04-15T03:53:02.330064Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This are the attributes of the metadata of the Images","metadata":{}},{"cell_type":"markdown","source":"# Data Visualizations ","metadata":{}},{"cell_type":"markdown","source":"### Bar Plot showing Frequency-Distribution of X-Rays Based on Source","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Select disease columns\ndisease_cols = ['Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum',\n                'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation', \n                'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other', 'Fracture', \n                'Support_Devices', 'Aortic_Enlargement', 'Calcification', 'ILD', 'Infiltration', \n                'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay']\n\n# Print source counts\nsource_counts = data['source'].value_counts()\n\n# Filter out sources with value 0.5\nfiltered_counts = source_counts[source_counts != 0.5]\n\nprint(filtered_counts)\n\n# Plot bar chart for filtered source counts\nplt.figure(figsize=(10, 6))\nfiltered_counts.plot(kind='bar', color='lightcoral', edgecolor='black')\nplt.title(\"Number of Images per Source\")\nplt.xlabel(\"Source\")\nplt.ylabel(\"Count\")\nplt.xticks(rotation=45, ha='right')\nplt.grid(axis='y', linestyle='--', alpha=0.7)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:02.595521Z","iopub.execute_input":"2025-04-15T03:53:02.595957Z","iopub.status.idle":"2025-04-15T03:53:02.913334Z","shell.execute_reply.started":"2025-04-15T03:53:02.595922Z","shell.execute_reply":"2025-04-15T03:53:02.912271Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Bar Plot showing Data Distribution Based on Modality and Position of X-Rays","metadata":{}},{"cell_type":"code","source":"# Create a new column that combines 'modality' and 'position'\ndata['modality_position'] = data['modality'] + \" + \" + data['position']\n\n# Count occurrences of each unique combination\nmodality_position_counts = data['modality_position'].value_counts()\n\n# Print the counts\nprint(\"Modality + Position Combinations:\")\nprint(modality_position_counts)\n\n# Plot the counts\nplt.figure(figsize=(12, 6))\nmodality_position_counts.plot(kind='bar', color='teal', edgecolor='black')\nplt.title(\"Counts of Modality + Position Combinations\")\nplt.xlabel(\"Modality + Position\")\nplt.ylabel(\"Count\")\nplt.xticks(rotation=45, ha='right')\nplt.grid(axis='y', linestyle='--', alpha=0.7)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:02.933313Z","iopub.execute_input":"2025-04-15T03:53:02.933636Z","iopub.status.idle":"2025-04-15T03:53:03.23272Z","shell.execute_reply.started":"2025-04-15T03:53:02.93361Z","shell.execute_reply":"2025-04-15T03:53:03.231718Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Bar Plots showing Disease-Frequency Distribtuion From Each Source","metadata":{}},{"cell_type":"code","source":"import math\nimport matplotlib.pyplot as plt\n\n# Group by 'source' and sum disease counts\nsource_disease_counts = data.groupby('source')[disease_cols].sum()\n\n# Remove values of exactly 0.5 using .where()\nsource_disease_counts = source_disease_counts.where(source_disease_counts != 0.5, 0)\n\n# Print only non-zero disease counts for each source\nfor source, counts in source_disease_counts.iterrows():\n    non_zero_counts = counts[counts > 0]  # Filter out zero counts\n    if not non_zero_counts.empty:\n        print(f\"Source: {source}\")\n        for disease, count in non_zero_counts.items():\n            print(f\"  {disease}: {count}\")\n        print(\"-\" * 50)  # Separator for readability\n\n# Determine grid layout: 3 columns, adjust rows dynamically\nnum_sources = len(source_disease_counts)\nnum_cols = 3  # Number of columns in the grid\nnum_rows = math.ceil(num_sources / num_cols)  # Calculate required rows\n\n# Create subplots with a compact figure size\nfig, axes = plt.subplots(nrows=num_rows, ncols=num_cols, figsize=(12, 3 * num_rows))\n\n# Flatten axes array if needed (when more than 1 row/column)\naxes = axes.flatten() if num_sources > 1 else [axes]\n\n# Plot bar charts for each source with all diseases (including zero counts)\nfor ax, (source, counts) in zip(axes, source_disease_counts.iterrows()):\n    counts.plot(kind='bar', ax=ax, color='skyblue', edgecolor='black')\n    ax.set_title(f\"{source}\", fontsize=10)  # Smaller title\n    ax.set_xlabel(\"\", fontsize=8)  # Remove xlabel\n    ax.set_ylabel(\"Count\", fontsize=8)  # Smaller ylabel\n    ax.set_xticklabels(counts.index, rotation=45, ha='right', fontsize=7)  # Smaller x-tick labels\n    ax.tick_params(axis='y', labelsize=7)  # Smaller y-tick labels\n    ax.grid(axis='y', linestyle='--', alpha=0.7)\n\n# Hide empty subplots if any\nfor i in range(len(source_disease_counts), len(axes)):\n    fig.delaxes(axes[i])\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:03.262989Z","iopub.execute_input":"2025-04-15T03:53:03.26342Z","iopub.status.idle":"2025-04-15T03:53:07.627059Z","shell.execute_reply.started":"2025-04-15T03:53:03.263384Z","shell.execute_reply":"2025-04-15T03:53:07.626135Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Bar Plots Showing Disease Frequency Distribution Of the Data","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Ensure all selected columns exist in data\ndisease_cols = [col for col in disease_cols if col in data.columns]\n\n# Replace values of exactly 0.5 with 0 before summing\nfiltered_data = data[disease_cols].where(data[disease_cols] != 0.5, 0)\n\n# Compute disease counts and convert to integers\ndisease_counts = filtered_data.sum().astype(int).sort_values(ascending=False)\n\n# Remove diseases with zero cases\ndisease_counts = disease_counts[disease_counts > 0]\n\n# Print disease counts as integers\nprint(disease_counts)\n\n# Plot\nplt.figure(figsize=(12, 6))\nsns.barplot(x=disease_counts.index, y=disease_counts.values, palette=\"viridis\")\nplt.xticks(rotation=90)\nplt.xlabel(\"Disease\")\nplt.ylabel(\"Count\")\nplt.title(\"Disease Counts in Dataset\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:07.628234Z","iopub.execute_input":"2025-04-15T03:53:07.628552Z","iopub.status.idle":"2025-04-15T03:53:08.220655Z","shell.execute_reply.started":"2025-04-15T03:53:07.628526Z","shell.execute_reply":"2025-04-15T03:53:08.21975Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Co-Occurence Matrix to Represent Distribution of Diseases Which Occur Together ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Drop 'Support_Devices' and 'Non_XRay' columns from dataset (if they exist)\ndata_filtered = data.drop(columns=[\"Support_Devices\", \"Non_XRay\"], errors=\"ignore\")\n\n# Remove \"Support_Devices\" and \"Non_XRay\" from disease_cols if present\ndisease_cols = [col for col in disease_cols if col not in [\"Support_Devices\", \"Non_XRay\"]]\n\n# Replace values of exactly 0.5 with 0, then convert NaN (or pd.NA) to 0 before computing co-occurrence\ndata_filtered[disease_cols] = data_filtered[disease_cols].where(data_filtered[disease_cols] != 0.5, 0).fillna(0).astype(int)\n\n# Compute co-occurrence matrix\nco_occurrence_matrix = data_filtered[disease_cols].T.dot(data_filtered[disease_cols])\n\n# Fill diagonal with zeros (to ignore self-co-occurrence)\nnp.fill_diagonal(co_occurrence_matrix.values, 0)\n\n# Plot heatmap with adjustments\nplt.figure(figsize=(18, 16))  # Increase figure size for better visibility\nsns.heatmap(co_occurrence_matrix, cmap=\"Blues\", annot=True, fmt=\".0f\", square=True, \n            annot_kws={\"size\": 8})  # Smaller font for annotations\nplt.xticks(rotation=45, fontsize=10)  # Adjust x-axis labels\nplt.yticks(fontsize=10)  # Adjust y-axis labels\nplt.title(\"Disease Co-Occurrence Matrix\", fontsize=16)\nplt.xlabel(\"Diseases\", fontsize=12)\nplt.ylabel(\"Diseases\", fontsize=12)\nplt.tight_layout()  # Ensure labels and titles don't overlap\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:08.222022Z","iopub.execute_input":"2025-04-15T03:53:08.22227Z","iopub.status.idle":"2025-04-15T03:53:10.024377Z","shell.execute_reply.started":"2025-04-15T03:53:08.222249Z","shell.execute_reply":"2025-04-15T03:53:10.023374Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Combination Pairs of Diseases Which Most Freqeuntly.\n\nSupport Devices and Non X-Ray Has Been Removed For Consideration.","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Drop 'Support_Devices' and 'Non_XRay' columns from dataset (if they exist)\ndata_filtered = data.drop(columns=[\"Support_Devices\", \"Non_XRay\"], errors=\"ignore\")\n\n# Remove \"Support_Devices\" and \"Non_XRay\" from disease_cols if present\ndisease_cols = [col for col in disease_cols if col not in [\"Support_Devices\", \"Non_XRay\"]]\n\n# Replace values of exactly 0.5 with 0, then convert NaN (or pd.NA) to 0\ndata_filtered[disease_cols] = data_filtered[disease_cols].where(data_filtered[disease_cols] != 0.5, 0).fillna(0).astype(int)\n\n# Convert each row of disease labels into a frozenset (only keeping labels with value 1)\ndata_filtered[\"Combination\"] = data_filtered[disease_cols].apply(\n    lambda row: frozenset(row[row == 1].index), axis=1\n)\n\n# Exclude single-disease cases\ndata_filtered = data_filtered[data_filtered[\"Combination\"].apply(lambda x: len(x) > 1)]\n\n# Count occurrences of each unique multi-disease combination\ncombination_counts = data_filtered[\"Combination\"].value_counts().reset_index()\ncombination_counts.columns = [\"Combination\", \"Count\"]\n\n# Keep only the top 50 most frequent multi-disease combinations\ntop_50_combinations = combination_counts.head(50).copy()\n\n# Convert sets to readable strings\ntop_50_combinations[\"Combination\"] = top_50_combinations[\"Combination\"].apply(lambda x: \" & \".join(sorted(x)))\n\n# Plot the top 50 most frequent multi-disease combinations\nplt.figure(figsize=(14, 12))\nsns.barplot(y=\"Combination\", x=\"Count\", data=top_50_combinations, palette=\"magma\")\nplt.xlabel(\"Frequency\", fontsize=14)\nplt.ylabel(\"Disease Combination\", fontsize=14)\nplt.title(\"Top 50 Most Frequent Multi-Disease Combinations\", fontsize=16)\nplt.gca().invert_yaxis()  # Invert for better readability\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:15.708841Z","iopub.execute_input":"2025-04-15T03:53:15.709137Z","iopub.status.idle":"2025-04-15T03:53:54.610593Z","shell.execute_reply.started":"2025-04-15T03:53:15.709114Z","shell.execute_reply":"2025-04-15T03:53:54.609652Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### List of Top 15 Co-occurring Disease Combinations","metadata":{}},{"cell_type":"code","source":"for _, row in top_50_combinations.head(15).iterrows():\n    print(f\"{row['Combination']}: {row['Count']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:54.61171Z","iopub.execute_input":"2025-04-15T03:53:54.611994Z","iopub.status.idle":"2025-04-15T03:53:54.621943Z","shell.execute_reply.started":"2025-04-15T03:53:54.611968Z","shell.execute_reply":"2025-04-15T03:53:54.620868Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Below are Few Sample Images From the Dataset With their Metadata","metadata":{}},{"cell_type":"code","source":"import os\nimport pydicom\nimport cv2\nimport math\nimport textwrap  # For wrapping long text\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Define the disease label columns\nlabel_columns = [\n    'Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum', \n    'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation', \n    'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other', \n    'Fracture', 'Support_Devices', 'Aortic_Enlargement', 'Calcification', \n    'ILD', 'Infiltration', 'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay'\n]\n\n# Load dataset\ndf = data.copy()\n\n# Replace 0.5 with 0 in disease columns\ndf[label_columns] = df[label_columns].where(df[label_columns] != 0.5, 0).astype(int)\n\n# Get unique sources\nsources = df[\"source\"].unique()\nnum_sources = len(sources)\n\n# Define grid layout (4 columns per row)\ncols = 4\nrows = math.ceil(num_sources / cols)\nfig, axes = plt.subplots(rows, cols, figsize=(5 * cols, 6 * rows))  # Adjust figure size\n\n# Flatten axes for easier iteration\nif num_sources == 1:\n    axes = [axes]\nelse:\n    axes = axes.flatten()\n\nfor i, source in enumerate(sources):\n    df_source = df[df[\"source\"] == source]  # Filter data for this source\n\n    # Pick one random sample safely\n    if not df_source.empty:\n        sample = df_source.sample(1, random_state=108)\n    else:\n        sample = pd.DataFrame()  # Empty DataFrame\n\n    if not sample.empty:\n        img_path = sample[\"img_path\"].values[0]\n\n        # Extract metadata\n        modality = sample[\"modality\"].values[0]\n        position = sample[\"position\"].values[0]\n\n        # Extract diseases (where value = 1)\n        present_diseases = [col for col in label_columns if sample.iloc[0][col] == 1]\n        disease_label = \", \".join(present_diseases) if present_diseases else \"No Disease\"\n\n        # Wrap text for better readability\n        title_text = f\"Source: {source}\\nModality: {modality} | Position: {position}\\nDiseases: {disease_label}\"\n        wrapped_text = \"\\n\".join(textwrap.wrap(title_text, width=50))\n\n        # Read Image (DICOM for VIN Big Data, JPG/PNG/BMP for others)\n        img = None\n        if os.path.exists(img_path):\n            if img_path.lower().endswith((\".dcm\", \".dicom\")):  # Handle DICOM formats\n                try:\n                    dicom_data = pydicom.dcmread(img_path)\n                    img = dicom_data.pixel_array  \n                    \n                    # Normalize safely\n                    min_val, max_val = np.min(img), np.max(img)\n                    if min_val != max_val:  \n                        img = ((img - min_val) / (max_val - min_val) * 255).astype(np.uint8)\n                    \n                except Exception as e:\n                    print(f\"Error reading DICOM: {img_path} - {e}\")\n            else:\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n\n        # Display Image\n        if img is not None:\n            axes[i].imshow(img, cmap=\"gray\")\n            axes[i].set_title(wrapped_text, fontsize=12, fontweight=\"bold\", color=\"blue\", loc=\"left\")\n            axes[i].axis(\"off\")\n        else:\n            # Display \"Image Not Found\" if img_path is missing\n            missing_text = f\"Image Not Found\\n(Source: {source})\"\n            axes[i].text(0.5, 0.5, missing_text, fontsize=14, ha=\"center\", va=\"center\", color=\"darkred\")\n            axes[i].axis(\"off\")\n\n# Hide extra unused axes\nfor j in range(i + 1, len(axes)):\n    axes[j].axis(\"off\")\n\n# Add a main title\nplt.suptitle(\"Sample X-ray Images with Metadata\\n\\n\", fontsize=18, fontweight=\"bold\", color=\"black\")\n\n# Adjust layout for better spacing\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:53:54.623735Z","iopub.execute_input":"2025-04-15T03:53:54.624078Z","iopub.status.idle":"2025-04-15T03:54:04.120236Z","shell.execute_reply.started":"2025-04-15T03:53:54.624042Z","shell.execute_reply":"2025-04-15T03:54:04.119303Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Splitting the Data into Train, Validation and Test Sets\n\n# Code Logic Explanation\n\nThis code is designed to create a balanced, multi-label dataset from a larger data collection and then split it into training, validation, and test sets using stratified sampling. The key steps are outlined below:\n\n## 1. Define Label Columns\n\n- **Label Columns**:  \n  A list named `label_columns` is defined that includes all disease labels (e.g., Normal, Pneumonia, Covid, etc.). These columns are used throughout the script to assess the presence (1), absence (0), or uncertainty (0.5, later handled as 0.5) of each condition in an image.\n\n## 2. Sampling Common Cases for Each Disease\n\n- **Function `sample_common_cases_for_disease`**:  \n  This function is responsible for sampling a fixed number (`samples_per_class`) of cases for a specific common disease while ensuring that no single data source contributes more than a defined fraction (`max_source_fraction`) of the total samples.  \n  - **Step-by-step process**:\n    - **Subset Data**: It filters the dataset to include only rows where the target disease is present (i.e., label equals 1).\n    - **Determine Per-Source Cap**: It computes the maximum allowed number of samples from any one source as a fraction of `samples_per_class`.\n    - **Group Sampling**: The data is grouped by the \"source\" column, and for each source, it randomly samples up to the allowed number of cases.\n    - **Fill Remaining Slots**: If the combined samples across sources are fewer than `samples_per_class`, additional cases are randomly sampled from the remaining pool (ignoring source limits).\n    - **Trim Excess**: In case the sampling overshoots the target, it randomly trims the combined samples to exactly `samples_per_class`.\n\n## 3. Creating the Balanced Dataset\n\n- **Function `create_balanced_dataset`**:  \n  This function builds a balanced dataset by handling rare and common diseases differently.\n  - **Disease Frequency Calculation**:  \n    It computes the overall frequency (i.e., count of positive cases) for each disease label.\n  - **Identifying Rare and Common Diseases**:  \n    Diseases with counts below a specified threshold (`rare_threshold`) are marked as rare. Those meeting or exceeding the threshold are considered common.\n  - **Retaining Rare Cases**:  \n    All images that have any rare disease (i.e., at least one positive label for any rare condition) are fully retained in the final dataset.\n  - **Sampling Common Cases**:  \n    For each common disease, the function calls `sample_common_cases_for_disease` on the remaining data (i.e., after excluding rare cases) to sample a fixed number of cases per common disease while enforcing source-based limits.\n  - **Merging Datasets**:  \n    Finally, the rare cases and the sampled common cases are merged. Note that an image might be multi-labeled and appear in multiple disease samples.\n\n## 4. Stratified Data Splitting\n\n- **Function `stratified_split`**:  \n  This function splits the balanced dataset into training, validation, and test sets using multi-label stratified sampling. The `MultilabelStratifiedShuffleSplit` class from the `iterstrat.ml_stratifiers` module is used to ensure that the distribution of disease labels is similar across all splits.\n  - **Process**:\n    - **Test Split**: The data is first split into train/validation and test sets.\n    - **Validation Split**: The remaining train/validation set is further split into separate training and validation sets, ensuring that the stratification accounts for the overall multi-label distribution.\n    \n## 5. Usage Example and Output\n\n- **Creating the Balanced Dataset**:  \n  The function `create_balanced_dataset` is called with the provided data and parameters (`samples_per_class=5000`, `rare_threshold=5000`, `max_source_fraction=0.2`). This produces a balanced dataset combining both rare and common disease cases.\n  \n- **Splitting the Data**:  \n  The balanced dataset is then split into train, validation, and test sets using the `stratified_split` function.\n  \n- **Reporting Distributions**:  \n  Finally, the code prints out the number of images in each split along with the disease distributions in each set. This helps verify that the stratification and balancing steps have preserved the multi-label nature of the data.\n\nOverall, this code ensures that the final dataset has a balanced representation of common diseases (with controlled source contribution) and retains all rare disease cases, while also maintaining a similar label distribution across different splits for reliable model training and evaluation.\n","metadata":{}},{"cell_type":"code","source":"pip install iterative-stratification","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:54:04.12173Z","iopub.execute_input":"2025-04-15T03:54:04.122172Z","iopub.status.idle":"2025-04-15T03:54:09.023526Z","shell.execute_reply.started":"2025-04-15T03:54:04.122115Z","shell.execute_reply":"2025-04-15T03:54:09.022424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\n\n# Define the disease label columns\nlabel_columns = [\n    'Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum', \n    'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation', \n    'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other', \n    'Fracture', 'Support_Devices', 'Aortic_Enlargement', 'Calcification', \n    'ILD', 'Infiltration', 'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay'\n]\n\ndef sample_common_cases_for_disease(disease, df, samples_per_class, max_source_fraction, random_state=42):\n    \"\"\"\n    For a given common disease, sample exactly samples_per_class rows from df (which should already\n    exclude rare cases) ensuring that no source contributes more than a fraction (max_source_fraction)\n    of the samples.\n    \"\"\"\n    # Subset rows positive for the disease\n    disease_df = df[df[disease] == 1].copy()\n    if disease_df.empty:\n        return pd.DataFrame()\n    \n    # Determine allowed samples per source for this disease.\n    allowed_per_source = int(max_source_fraction * samples_per_class)\n    \n    # Group by source and sample up to allowed_per_source from each group.\n    samples_list = []\n    for source, group in disease_df.groupby(\"source\"):\n        n_sample = min(len(group), allowed_per_source)\n        samples_list.append(group.sample(n=n_sample, random_state=random_state))\n    \n    # Combine samples from all sources.\n    sampled_disease = pd.concat(samples_list)\n    \n    # If we still need more cases to reach samples_per_class, sample additional rows\n    # from disease_df (excluding already chosen ones), even if they come from sources that\n    # have already reached the cap.\n    remaining_needed = samples_per_class - len(sampled_disease)\n    if remaining_needed > 0:\n        remaining_pool = disease_df.loc[~disease_df.index.isin(sampled_disease.index)]\n        if not remaining_pool.empty:\n            additional = remaining_pool.sample(n=min(remaining_needed, len(remaining_pool)), random_state=random_state)\n            sampled_disease = pd.concat([sampled_disease, additional])\n    \n    # In case we oversampled, randomly trim to exactly samples_per_class.\n    if len(sampled_disease) > samples_per_class:\n        sampled_disease = sampled_disease.sample(n=samples_per_class, random_state=random_state)\n    \n    return sampled_disease\n\ndef create_balanced_dataset(data, samples_per_class=5000, rare_threshold=5000, max_source_fraction=0.2, random_state=42):\n    \"\"\"\n    1. Compute disease frequencies and identify rare and common diseases.\n    2. Retain all rows that contain any rare disease.\n    3. For each common disease, sample exactly `samples_per_class` images from the remaining data,\n       enforcing a maximum per-source contribution.\n    4. Merge the rare cases and the common disease samples.\n    \"\"\"\n    print(\"\\nOriginal dataset disease frequency:\")\n    print((data[label_columns] == 1).sum())\n    \n    # Calculate overall counts per disease.\n    overall_counts = (data[label_columns] == 1).sum()\n    rare_diseases = overall_counts[overall_counts < rare_threshold].index.tolist()\n    common_diseases = overall_counts[overall_counts >= rare_threshold].index.tolist()\n    \n    print(\"\\nRare diseases:\", rare_diseases)\n    print(\"Common diseases:\", common_diseases)\n    \n    # Step 1: Extract all rows that have any rare disease (fully retained).\n    rare_set = data[data[rare_diseases].sum(axis=1) > 0]\n    \n    # Step 2: Exclude rare cases from common disease sampling.\n    remaining_set = data.loc[~data.index.isin(rare_set.index)]\n    \n    # Step 3: For each common disease, sample with source balancing.\n    common_samples = []\n    for disease in common_diseases:\n        print(f\"\\nSampling for common disease: {disease}\")\n        sampled_disease = sample_common_cases_for_disease(\n            disease,\n            remaining_set,\n            samples_per_class,\n            max_source_fraction,\n            random_state\n        )\n        print(f\"Selected {len(sampled_disease)} samples for {disease}\")\n        common_samples.append(sampled_disease)\n    \n    # Combine all common disease samples. Note: duplicates may occur if a row is positive for multiple common diseases.\n    common_set = pd.concat(common_samples)\n    \n    # Step 4: Merge rare and common cases.\n    final_data = pd.concat([rare_set, common_set])\n    \n    print(\"\\nFinal dataset disease frequency (counts may exceed quota if images are multi-labeled):\")\n    print((final_data[label_columns] == 1).sum())\n    \n    return final_data\n\ndef stratified_split(data, test_size=0.15, val_size=0.15, random_state=42):\n    \"\"\"\n    Split the data into train, validation, and test sets using multilabel stratification.\n    \"\"\"\n    msss = MultilabelStratifiedShuffleSplit(n_splits=1, test_size=test_size, random_state=random_state)\n    X = data.index.values.reshape(-1, 1)\n    y = (data[label_columns] == 1).astype(int).values\n\n    for train_val_idx, test_idx in msss.split(X, y):\n        train_val = data.iloc[train_val_idx]\n        test = data.iloc[test_idx]\n\n    msss_val = MultilabelStratifiedShuffleSplit(n_splits=1, test_size=val_size / (1 - test_size), random_state=random_state)\n    X_val = train_val.index.values.reshape(-1, 1)\n    y_val = (train_val[label_columns] == 1).astype(int).values\n\n    for train_idx, val_idx in msss_val.split(X_val, y_val):\n        train = train_val.iloc[train_idx]\n        val = train_val.iloc[val_idx]\n\n    return train, val, test\n\n# Usage example:\n# Step 1-4: Create balanced dataset with exactly samples_per_class images per common disease.\nfinal_data = create_balanced_dataset(\n    data,\n    samples_per_class=5000,\n    rare_threshold=5000,         # rare_threshold now matches the intended value\n    max_source_fraction=0.2,\n    random_state=42\n)\n\n# Step 5: Perform stratified splits\ntrain_data, val_data, test_data = stratified_split(final_data)\n\n# Step 6: Print dataset distributions\nprint(\"\\nFinal Train size:\", len(train_data))\nprint(\"Final Validation size:\", len(val_data))\nprint(\"Final Test size:\", len(test_data))\n\nprint(\"\\nTrain set disease distribution:\")\nprint((train_data[label_columns] == 1).sum())\n\nprint(\"\\nValidation set disease distribution:\")\nprint((val_data[label_columns] == 1).sum())\n\nprint(\"\\nTest set disease distribution:\")\nprint((test_data[label_columns] == 1).sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:54:09.024764Z","iopub.execute_input":"2025-04-15T03:54:09.025018Z","iopub.status.idle":"2025-04-15T03:54:12.963721Z","shell.execute_reply.started":"2025-04-15T03:54:09.024997Z","shell.execute_reply":"2025-04-15T03:54:12.962812Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Bar Plots Showing Sources Distribution & Disease Distribution in Training, Validation & Testing Sets\n\nNote: We are using a reduced subset of the Training Set","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\n\n# Assuming you have already loaded your data and defined the `label_columns` and the downsampling logic above...\n\n# === Diagnostics: Check distributions and plot results ===\n\n# Plot source distributions for train, validation, and test sets\nfig, axs = plt.subplots(1, 3, figsize=(18, 6))\n\n# Calculate value counts for the 'source' column in each split\ntrain_source_counts = train_data['source'].value_counts()\nvalid_source_counts = val_data['source'].value_counts()\ntest_source_counts  = test_data['source'].value_counts()\n\n# Plotting for Train Set\naxs[0].bar(train_source_counts.index, train_source_counts.values, color='skyblue')\naxs[0].set_title(\"Train Set Source Distribution\")\naxs[0].set_xlabel(\"Source\")\naxs[0].set_ylabel(\"Count\")\naxs[0].tick_params(axis='x', rotation=45)\n\n# Plotting for Validation Set\naxs[1].bar(valid_source_counts.index, valid_source_counts.values, color='lightgreen')\naxs[1].set_title(\"Validation Set Source Distribution\")\naxs[1].set_xlabel(\"Source\")\naxs[1].set_ylabel(\"Count\")\naxs[1].tick_params(axis='x', rotation=45)\n\n# Plotting for Test Set\naxs[2].bar(test_source_counts.index, test_source_counts.values, color='salmon')\naxs[2].set_title(\"Test Set Source Distribution\")\naxs[2].set_xlabel(\"Source\")\naxs[2].set_ylabel(\"Count\")\naxs[2].tick_params(axis='x', rotation=45)\n\nplt.tight_layout()\nplt.show()\n\n# Plot disease distributions for train, validation, and test sets\n# Sort by counts to see the distribution clearly\ntrain_disease_counts = train_data[label_columns].sum().sort_values(ascending=False)\nvalid_disease_counts = val_data[label_columns].sum().sort_values(ascending=False)\ntest_disease_counts  = test_data[label_columns].sum().sort_values(ascending=False)\n\n# Set up the subplots for disease distribution plots\nfig, axs = plt.subplots(3, 1, figsize=(12, 18))\n\n# Plotting for Train Set Diseases\naxs[0].bar(train_disease_counts.index, train_disease_counts.values, color='skyblue')\naxs[0].set_title(\"Train Set Disease Distribution\")\naxs[0].set_xlabel(\"Disease\")\naxs[0].set_ylabel(\"Count\")\naxs[0].tick_params(axis='x', rotation=90)\n\n# Plotting for Validation Set Diseases\naxs[1].bar(valid_disease_counts.index, valid_disease_counts.values, color='lightgreen')\naxs[1].set_title(\"Validation Set Disease Distribution\")\naxs[1].set_xlabel(\"Disease\")\naxs[1].set_ylabel(\"Count\")\naxs[1].tick_params(axis='x', rotation=90)\n\n# Plotting for Test Set Diseases\naxs[2].bar(test_disease_counts.index, test_disease_counts.values, color='salmon')\naxs[2].set_title(\"Test Set Disease Distribution\")\naxs[2].set_xlabel(\"Disease\")\naxs[2].set_ylabel(\"Count\")\naxs[2].tick_params(axis='x', rotation=90)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:54:12.964599Z","iopub.execute_input":"2025-04-15T03:54:12.964837Z","iopub.status.idle":"2025-04-15T03:54:14.532173Z","shell.execute_reply.started":"2025-04-15T03:54:12.964807Z","shell.execute_reply":"2025-04-15T03:54:14.531219Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Custom ImageDataset Class To Read Dicom and other Formats Images & Apply Augmentations, CLAHE, Noise to the Image.","metadata":{}},{"cell_type":"code","source":"pip install -U albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:54:14.533121Z","iopub.execute_input":"2025-04-15T03:54:14.533408Z","iopub.status.idle":"2025-04-15T03:54:19.990503Z","shell.execute_reply.started":"2025-04-15T03:54:14.533384Z","shell.execute_reply":"2025-04-15T03:54:19.989238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport numpy as np\nimport pydicom\nfrom torch.utils.data import Dataset\nfrom albumentations import (Compose, Resize, RandomResizedCrop, Normalize, \n                            Rotate, HorizontalFlip, RandomBrightnessContrast)\nfrom albumentations.pytorch import ToTensorV2\n\n\ndef load_dicom_image(img_path):\n    \"\"\"Loads and normalizes a DICOM image safely.\"\"\"\n    ds = pydicom.dcmread(img_path)\n    img = ds.pixel_array.astype(np.float32)\n\n    # Normalize to [0, 255] if min/max are different\n    min_val, max_val = np.min(img), np.max(img)\n    if max_val > min_val:\n        img = ((img - min_val) / (max_val - min_val) * 255.0).astype(np.uint8)\n\n    # Convert grayscale to RGB\n    return cv2.cvtColor(img, cv2.COLOR_GRAY2RGB) if len(img.shape) == 2 else img\n\n\ndef apply_gaussian_filter(img, ksize=(3, 3), sigmaX=0.5):\n    \"\"\"Applies Gaussian blur to reduce artifacts.\"\"\"\n    return cv2.GaussianBlur(img, ksize, sigmaX)\n\n\ndef apply_clahe(img):\n    \"\"\"Applies CLAHE for low and medium contrast images.\"\"\"\n    lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n    l, a, b = cv2.split(lab)\n\n    # Apply CLAHE if contrast is low or medium\n    if np.std(l) < 70:  \n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        l = clahe.apply(l)\n\n    enhanced_lab = cv2.merge((l, a, b))\n    return cv2.cvtColor(enhanced_lab, cv2.COLOR_LAB2RGB)\n\n\nclass ImageDataset(Dataset):\n    def __init__(self, df, data_dir=None, augs=None, is_test=False, is_valid=False, \n                 use_clahe=False, use_gaussian=True, use_rgb_conversion=True):\n        \n        self.df = df\n        self.data_dir = data_dir\n        self.is_test = is_test\n        self.is_valid = is_valid\n        self.use_clahe = use_clahe\n        self.use_gaussian = use_gaussian\n        self.use_rgb_conversion = use_rgb_conversion\n        self.augs = augs if augs else self.default_augs(is_test, is_valid)\n\n        self.label_columns = [\n            'Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum',\n            'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation',\n            'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other',\n            'Fracture', 'Support_Devices', 'Aortic_Enlargement', 'Calcification',\n            'ILD', 'Infiltration', 'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay'\n        ]\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row['img_path']\n\n        # Load Image\n        if img_path.lower().endswith(('.dicom', '.dcm')):\n            img = load_dicom_image(img_path)\n        elif img_path.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp')):\n            img = cv2.imread(img_path)\n            if img is None:\n                raise FileNotFoundError(f\"Image not found at {img_path}.\")\n            if self.use_rgb_conversion:\n                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        else:\n            raise ValueError(f\"Unsupported image format: {img_path}\")\n\n        # Apply CLAHE if enabled\n        if self.use_clahe:\n            img = apply_clahe(img)\n\n        # Apply Gaussian filtering before resizing\n        if self.use_gaussian:\n            img = apply_gaussian_filter(img)\n\n        # Get labels\n        label = row[self.label_columns].values.astype(np.float32)\n\n        # Apply augmentations (including resizing only for test/valid)\n        img = self.augs(image=img)[\"image\"]\n\n        return img, torch.tensor(label, dtype=torch.float32)\n\n    def default_augs(self, is_test, is_valid):\n        \"\"\"Defines default augmentations with high-quality interpolation.\"\"\"\n        if is_test or is_valid:\n            return Compose([\n                Resize(224, 224, interpolation=cv2.INTER_LANCZOS4),  # High-quality resize for test/valid\n                Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n                ToTensorV2()\n            ])\n        return Compose([\n            RandomResizedCrop(size=(224, 224), scale=(0.8, 1.0), interpolation=cv2.INTER_CUBIC),  # FIXED TUPLE ISSUE\n            Rotate(limit=12, p=0.5),\n            HorizontalFlip(p=0.1),\n            RandomBrightnessContrast(p=0.5),\n            Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ])\n\n    def __repr__(self):\n        return (f\"ImageDataset(size={len(self.df)}, CLAHE={self.use_clahe}, \"\n                f\"Gaussian={self.use_gaussian}, RGB_Conversion={self.use_rgb_conversion})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:54:19.992903Z","iopub.execute_input":"2025-04-15T03:54:19.993138Z","iopub.status.idle":"2025-04-15T03:54:24.444855Z","shell.execute_reply.started":"2025-04-15T03:54:19.993118Z","shell.execute_reply":"2025-04-15T03:54:24.444123Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training, Validation & Test Examples Counts","metadata":{}},{"cell_type":"code","source":"# Create datasets with both CLAHE and Super-Resolution enabled\ntrainset = ImageDataset(train_data, data_dir=None, is_test=False, use_clahe=True)\nvalidset = ImageDataset(val_data, data_dir=None, is_valid=True, use_clahe=True)\ntestset = ImageDataset(test_data, data_dir=None, is_test=True, use_clahe=True)  \n\n# Optional: Verify dataset sizes\nprint(f\"Training set size: {len(trainset)}\")\nprint(f\"Validation set size: {len(validset)}\")\nprint(f\"Test set size: {len(testset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T03:58:47.805427Z","iopub.execute_input":"2025-04-15T03:58:47.805742Z","iopub.status.idle":"2025-04-15T03:58:47.82006Z","shell.execute_reply.started":"2025-04-15T03:58:47.805719Z","shell.execute_reply":"2025-04-15T03:58:47.819331Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Some Sample Images from Trainset, Validset and Testset after apply necessary Augmentations, Noise & CLAHE","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Function to display an image and its multi-labels\ndef show_image(img_tensor, labels, label_columns, title=\"Sample Image\"):\n    img = img_tensor.permute(1, 2, 0).numpy()  # Convert (C, H, W) → (H, W, C)\n    \n    # Denormalize the image\n    img = img * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])  \n    img = np.clip(img, 0, 1)  # Clip values to [0, 1]\n     \n    # Ensure the range is correctly adjusted (for multi-label classification)\n    active_labels = [label_columns[i] for i in range(len(labels)) if labels[i] == 1]\n    \n    label_text = \", \".join(active_labels) if active_labels else \"No labels\"\n\n    plt.imshow(img)\n    plt.title(f\"{title}: {label_text}\")  # Display active labels\n    plt.axis(\"off\")\n    plt.show()\n\n# Display a sample image from each dataset (adjusting for correct indexing of labels)\nsample_train_img, sample_train_labels = trainset[13415]\nsample_valid_img, sample_valid_labels = validset[6379]\nsample_test_img, sample_test_labels = testset[1011]\n\n# Use the same label columns used for multi-label classification\nlabel_columns = [\n    'Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum', \n    'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation', \n    'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other', \n    'Fracture', 'Support_Devices', 'Aortic_Enlargement', 'Calcification', \n    'ILD', 'Infiltration', 'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay'\n]\n\n# Display images\nshow_image(sample_train_img, sample_train_labels, label_columns, \"Train Sample\")\nshow_image(sample_valid_img, sample_valid_labels, label_columns, \"Validation Sample\")\nshow_image(sample_test_img, sample_test_labels, label_columns, \"Test Sample\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T04:00:28.756151Z","iopub.execute_input":"2025-04-15T04:00:28.75652Z","iopub.status.idle":"2025-04-15T04:00:29.430021Z","shell.execute_reply.started":"2025-04-15T04:00:28.756492Z","shell.execute_reply":"2025-04-15T04:00:29.429161Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Creating Batches for Trainset, Validset & Testset","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\n# Use CUDA if available\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Create DataLoader instances\ntrain_loader = DataLoader(trainset, batch_size=16, shuffle=True, \n                          num_workers=4, pin_memory=True if torch.cuda.is_available() else False)\n\nval_loader = DataLoader(validset, batch_size=16, shuffle=False, \n                        num_workers=4, pin_memory=True if torch.cuda.is_available() else False)\n\ntest_loader = DataLoader(testset, batch_size=16, shuffle=False, \n                         num_workers=4, pin_memory=True if torch.cuda.is_available() else False)\n\n# Print number of batches\nprint(f\"No. of batches in train_loader : {len(train_loader)}\")\nprint(f\"No. of batches in val_loader   : {len(val_loader)}\")\nprint(f\"No. of batches in test_loader  : {len(test_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T04:00:29.431316Z","iopub.execute_input":"2025-04-15T04:00:29.431633Z","iopub.status.idle":"2025-04-15T04:00:29.501818Z","shell.execute_reply.started":"2025-04-15T04:00:29.431602Z","shell.execute_reply":"2025-04-15T04:00:29.500905Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training and Validation Phase - Model: maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k","metadata":{}},{"cell_type":"code","source":"!pip install timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T04:00:31.793904Z","iopub.execute_input":"2025-04-15T04:00:31.79419Z","iopub.status.idle":"2025-04-15T04:00:35.127321Z","shell.execute_reply.started":"2025-04-15T04:00:31.794169Z","shell.execute_reply":"2025-04-15T04:00:35.126434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom torch.utils.tensorboard import SummaryWriter\nimport os\nfrom tqdm import tqdm\nimport timm\nfrom torch.cuda.amp import GradScaler, autocast\nfrom sklearn.metrics import roc_auc_score, precision_score, recall_score, f1_score\n\n# Device and Hyperparameter Setup\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nNUM_LABELS = 24  # Number of multi-label classes\n\n# --- Positive Weight Calculation for BCE Loss ---\n# (Assuming that `train_loader` is already defined)\npositive_counts = np.array([\n    3747, 5239, 2938, 3500, 5303, 7828, 22921, 4737, 11701, 5831,\n    9010, 6290, 18379, 4050, 4789, 23132, 3500, 672,\n    700, 873, 1806, 1542, 3259, 950\n], dtype=np.float32)\ntotal_samples = len(train_loader.dataset)\nnegative_counts = total_samples - positive_counts\n\n# Compute raw positive weights and normalize so that max weight is 20.\nraw_pos_weights = negative_counts / (positive_counts + 1e-6)\nmax_weight = np.max(raw_pos_weights)\nnormalized_pos_weights = (raw_pos_weights / max_weight) * 20\npos_weights_tensor = torch.tensor(normalized_pos_weights, dtype=torch.float32, device=DEVICE)\n\n# Use in BCEWithLogitsLoss\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights_tensor, reduction='none')\nprint(\"Positive weights:\", pos_weights_tensor)\n\n# --- FastConfidenceScheduler Definition ---\nclass FastConfidenceScheduler:\n    \"\"\"\n    Quickly adapts the confidence threshold within a few epochs.\n    Supports linear or cosine adjustment.\n    \"\"\"\n    def __init__(self, start=0.65, end=0.85, warmup_steps=2000, mode=\"linear\"):\n        self.start = start\n        self.end = end\n        self.warmup_steps = warmup_steps  # Increase confidence threshold within this many steps\n        self.mode = mode\n\n    def get_threshold(self, step):\n        \"\"\"Computes the confidence threshold for a given step.\"\"\"\n        progress = min(step / self.warmup_steps, 1.0)  # Normalize between 0 and 1\n        if self.mode == \"linear\":\n            return self.start + (self.end - self.start) * progress\n        elif self.mode == \"cosine\":\n            return self.start + (self.end - self.start) * (1 - 0.5 * (1 + torch.cos(torch.tensor(progress * 3.1416, device=DEVICE))))\n        else:\n            raise ValueError(\"Unsupported mode. Choose 'linear' or 'cosine'.\")\n\n# Initialize confidence scheduler (for global dynamic threshold in soft smoothing)\nconfidence_scheduler = FastConfidenceScheduler(start=0.65, end=0.85, warmup_steps=2000, mode=\"linear\")\n\n# --- Loss Function with Adaptive Label Smoothing ---\ndef soft_target_loss(outputs, labels, step, conf_scheduler):\n    \"\"\"\n    Adaptive label smoothing for uncertain labels (0.5) with a fast-changing confidence threshold.\n    Uses the provided confidence scheduler to get the threshold for the current step.\n    \"\"\"\n    confidence_threshold = conf_scheduler.get_threshold(step)\n    probs = torch.sigmoid(outputs).clamp(1e-6, 1 - 1e-6)\n    # For uncertain labels (0.5), if the model is confident then use model's prediction; else keep 0.5.\n    smoothed_labels = torch.where(\n        (labels == 0.5) & ((probs > confidence_threshold) | (probs < (1 - confidence_threshold))),\n        probs,\n        torch.full_like(labels, 0.5)\n    )\n    weights = torch.where(labels == 0.5, 0.2, 1.0)\n    loss = criterion(outputs, smoothed_labels) * weights\n    return loss.mean()\n\n# --- Differentiable Soft F1 Score (for potential logging) ---\ndef soft_f1_score(y_true, y_pred, mask, epsilon=1e-7, return_classwise=False):\n    y_pred = torch.sigmoid(y_pred)\n    tp = torch.sum(y_pred * y_true * mask, dim=0)\n    fp = torch.sum(y_pred * (1 - y_true) * mask, dim=0)\n    fn = torch.sum((1 - y_pred) * y_true * mask, dim=0)\n    f1 = (2 * tp + epsilon) / (2 * tp + fp + fn + epsilon)\n    mask_sum = mask.sum(dim=0).clamp(min=1)\n    if return_classwise:\n        return f1\n    return (f1 * mask).sum() / mask_sum\n\n# --- Metric Functions ---\ndef multi_label_accuracy(y_true, y_pred, thresholds):\n    \"\"\"\n    Computes the multi-label accuracy using per-class thresholds.\n    thresholds: A tensor of thresholds with shape (NUM_LABELS,).\n    \"\"\"\n    y_pred_bin = (y_pred >= thresholds).float()\n    correct = (y_pred_bin == y_true).float()\n    per_sample_acc = correct.sum(dim=1) / float(y_true.shape[1])\n    return per_sample_acc.mean().item()\n\ndef compute_extra_metrics(all_labels, all_preds):\n    \"\"\"\n    Computes additional metrics (ROC-AUC, Precision, Recall) for multi-label classification.\n    \"\"\"\n    all_labels = np.concatenate(all_labels, axis=0)\n    all_preds = np.concatenate(all_preds, axis=0)\n    \n    # Default threshold of 0.5 is used for binarization here.\n    all_preds_bin = (all_preds >= 0.5).astype(int)\n    \n    aucs, precisions, recalls = [], [], []\n    for i in range(all_preds.shape[1]):\n        true_label = all_labels[:, i]\n        pred_prob = all_preds[:, i]\n        pred_bin = all_preds_bin[:, i]\n        if len(np.unique(true_label)) < 2:\n            continue\n        auc = roc_auc_score(true_label, pred_prob)\n        precision = precision_score(true_label, pred_bin, zero_division=0)\n        recall = recall_score(true_label, pred_bin, zero_division=0)\n        aucs.append(auc)\n        precisions.append(precision)\n        recalls.append(recall)\n    \n    return np.nanmean(aucs), np.nanmean(precisions), np.nanmean(recalls)\n\ndef find_optimal_thresholds(labels, preds, masks=None):\n    \"\"\"\n    Computes the best threshold for each class by maximizing F1-score.\n    If masks is provided, only considers confident labels.\n    \"\"\"\n    num_classes = labels.shape[1]\n    thresholds = np.zeros(num_classes)\n    for i in range(num_classes):\n        best_f1 = 0.0\n        best_thresh = 0.5  # Default\n        for t in np.linspace(0.1, 0.9, 17):\n            pred_bin = (preds[:, i] >= t).astype(int)\n            f1 = f1_score(labels[:, i], pred_bin, zero_division=0)\n            if f1 > best_f1:\n                best_f1 = f1\n                best_thresh = t\n        thresholds[i] = best_thresh\n    return thresholds\n\ndef evaluate_model(model, loader, mode=\"Validation\", compute_thresholds=False):\n    \"\"\"\n    Evaluates the model using per-class thresholds computed from the validation set.\n    Uses the stable state (i.e. thresholds computed over the entire validation set).\n    \"\"\"\n    model.eval()\n    total_loss = 0.0\n    valid_batches = 0\n    all_labels, all_preds = [], []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=f\"[{mode}]\", unit=\"batch\", leave=True):\n            images = images.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True).float()\n            outputs = model(images)\n            loss = criterion(outputs, labels).mean()\n            total_loss += loss.item()\n            valid_batches += 1\n            preds = torch.sigmoid(outputs)\n            all_labels.append(labels.cpu().numpy())\n            all_preds.append(preds.cpu().numpy())\n    \n    avg_loss = total_loss / valid_batches\n    all_labels = np.concatenate(all_labels, axis=0)\n    all_preds = np.concatenate(all_preds, axis=0)\n    \n    # Compute per-class optimal thresholds from validation set if requested.\n    if compute_thresholds:\n        thresholds = find_optimal_thresholds(all_labels, all_preds)\n    else:\n        thresholds = np.full((all_labels.shape[1],), 0.5)\n    \n    # Binarize predictions using per-class thresholds and compute metrics.\n    all_preds_bin = (all_preds >= thresholds).astype(int)\n    f1_scores = []\n    for i in range(all_labels.shape[1]):\n        f1 = f1_score(all_labels[:, i], all_preds_bin[:, i], zero_division=0)\n        f1_scores.append(f1)\n    avg_f1 = np.nanmean(f1_scores)\n    avg_auc, avg_precision, avg_recall = compute_extra_metrics([all_labels], [all_preds])\n    \n    print(f\"{mode} Metrics: Loss: {avg_loss:.4f}, F1: {avg_f1:.4f}, AUC: {avg_auc:.4f}, \"\n          f\"Precision: {avg_precision:.4f}, Recall: {avg_recall:.4f}\")\n    return avg_loss, avg_f1, avg_auc, avg_precision, avg_recall, thresholds\n\n# --- Training Function with Per-Class Thresholds and Checkpoint Resumption ---\ndef train_model(model, train_loader, val_loader, optimizer, num_epochs=10, patience=5, \n                accumulation_steps=4, save_interval=100,\n                checkpoint_path='recent_checkpoint.pt'):\n    \"\"\"\n    Trains the model with per-class thresholds.\n    Also includes logic to resume from a checkpoint if one exists.\n    \"\"\"\n    print(\"Training started...\")\n    writer = SummaryWriter(log_dir='tensorboard_logs')\n    total_steps = num_epochs * len(train_loader)\n    lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=total_steps)\n    scaler = GradScaler()  # Note: No device argument here\n    \n    # Initialize per-class running thresholds with the initial confidence threshold.\n    init_threshold = confidence_scheduler.get_threshold(0)\n    running_thresholds = torch.full((NUM_LABELS,), init_threshold, dtype=torch.float32, device=DEVICE)\n    threshold_momentum = 0.1  # Adjust momentum as needed\n    \n    global_step = 0\n    best_val_f1 = 0.0\n    epochs_without_improvement = 0\n\n    # Resume from checkpoint if available\n    if os.path.exists(checkpoint_path):\n        print(f\"Resuming from checkpoint: {checkpoint_path}\")\n        checkpoint = torch.load(checkpoint_path, map_location=DEVICE)\n        model.load_state_dict(checkpoint.get('model_state_dict', {}))\n        optimizer.load_state_dict(checkpoint.get('optimizer_state_dict', {}))\n        lr_scheduler.load_state_dict(checkpoint.get('scheduler_state_dict', {}))\n        global_step = checkpoint.get('global_step', 0)\n        best_val_f1 = checkpoint.get('best_val_f1', 0.0)\n        epochs_without_improvement = checkpoint.get('epochs_without_improvement', 0)\n        running_thresholds = checkpoint.get('running_thresholds', running_thresholds)\n\n    for epoch in range(num_epochs):\n        model.train()\n        running_loss = 0.0\n        train_batches = 0\n        \n        # Single progress bar for epoch\n        with tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\", unit=\"batch\") as pbar:\n            for i, (images, labels) in enumerate(pbar):\n                images = images.to(DEVICE, non_blocking=True)\n                labels = labels.to(DEVICE, non_blocking=True).float()\n                # Include all labels in training\n                \n                with autocast():\n                    outputs = model(images)\n                    loss = soft_target_loss(outputs, labels, global_step, confidence_scheduler) / accumulation_steps\n                scaler.scale(loss).backward()\n                \n                if (i + 1) % accumulation_steps == 0 or (i + 1) == len(train_loader):\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad()\n                \n                running_loss += loss.item() * accumulation_steps\n                train_batches += 1\n                global_step += 1\n                lr_scheduler.step()\n                \n                # Update per-class running thresholds from current batch.\n                with torch.no_grad():\n                    batch_preds = torch.sigmoid(outputs)\n                    batch_thresholds = batch_preds.mean(dim=0)\n                    running_thresholds = (1 - threshold_momentum) * running_thresholds + threshold_momentum * batch_thresholds\n                \n                # Use the running thresholds for calculating training accuracy.\n                batch_acc = multi_label_accuracy(labels, torch.sigmoid(outputs), running_thresholds)\n                pbar.set_postfix(loss=running_loss/train_batches, acc=batch_acc)\n                \n                # Save checkpoint at intervals.\n                if global_step % save_interval == 0:\n                    checkpoint_dict = {\n                        'global_step': global_step,\n                        'model_state_dict': model.state_dict(),\n                        'optimizer_state_dict': optimizer.state_dict(),\n                        'scheduler_state_dict': lr_scheduler.state_dict(),\n                        'best_val_f1': best_val_f1,\n                        'epochs_without_improvement': epochs_without_improvement,\n                        'running_thresholds': running_thresholds,\n                        'epoch': epoch,\n                        'batch': i + 1,\n                    }\n                    torch.save(checkpoint_dict, checkpoint_path)\n                    print(f\"\\nCheckpoint saved at epoch {epoch+1}, batch {i+1}.\")\n\n        epoch_loss = running_loss / train_batches\n        writer.add_scalar(\"Loss/Train\", epoch_loss, epoch)\n        writer.add_scalar(\"Accuracy/Train\", batch_acc, epoch)\n        \n        # Evaluate on validation set using stable (global) thresholds.\n        val_loss, val_f1, val_auc, val_precision, val_recall, val_thresholds = evaluate_model(\n            model, val_loader, mode=\"Validation\", compute_thresholds=True\n        )\n        writer.add_scalar(\"Loss/Validation\", val_loss, epoch)\n        writer.add_scalar(\"F1/Validation\", val_f1, epoch)\n        writer.add_scalar(\"AUC/Validation\", val_auc, epoch)\n        writer.add_scalar(\"Precision/Validation\", val_precision, epoch)\n        writer.add_scalar(\"Recall/Validation\", val_recall, epoch)\n\n        # Save best model checkpoint if improvement is observed.\n        if val_f1 > best_val_f1:\n            best_val_f1 = val_f1\n            epochs_without_improvement = 0\n            torch.save(model.state_dict(), \"best_model.pt\")\n            print(f\"Epoch {epoch+1}: Validation F1 improved to {val_f1:.4f}. Best checkpoint saved.\")\n        else:\n            epochs_without_improvement += 1\n            print(f\"Epoch {epoch+1}: No improvement. Patience {epochs_without_improvement}/{patience}.\")\n\n        if epochs_without_improvement >= patience:\n            print(\"Early stopping triggered.\")\n            break\n\n    writer.close()\n    return best_val_f1\n\n# --- Model Initialization and Fine-Tuning Setup ---\n# Create the MaxViT model.\nmaxvit = timm.create_model('maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k', pretrained=False)\npretrained_weights_path = \"/kaggle/input/weights-after-3-epochs/best_model.pt\"\nif os.path.exists(pretrained_weights_path):\n    print(f\"Loading pre-trained weights from {pretrained_weights_path}...\")\n    state_dict = torch.load(pretrained_weights_path, map_location=DEVICE)\n    state_dict.pop(\"head.fc.weight\", None)\n    state_dict.pop(\"head.fc.bias\", None)\n    maxvit.load_state_dict(state_dict, strict=False)\nelse:\n    print(\"No pre-trained weights found. Training from scratch.\")\n\n# Replace the classifier head for multi-label classification.\nnum_ftrs = maxvit.head.fc.in_features\nmaxvit.head.fc = nn.Linear(num_ftrs, NUM_LABELS)\n\n# Freeze all parameters except those in 'blocks' and the new head.\nfor param in maxvit.parameters():\n    param.requires_grad = False\nfor name, param in maxvit.named_parameters():\n    if \"blocks\" in name or \"head.fc\" in name:\n        param.requires_grad = True\n\nmaxvit = maxvit.to(DEVICE)\noptimizer = optim.Adam(filter(lambda p: p.requires_grad, maxvit.parameters()), lr=1e-4)\n\n# --- Start Training ---\nbest_val_f1 = train_model(maxvit, train_loader, val_loader, optimizer, num_epochs=10, patience=5,\n                          accumulation_steps=4, save_interval=100,\n                          checkpoint_path='recent_checkpoint.pt')\nprint(f\"Best Validation F1 Score: {best_val_f1:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Evaluation on Test Data with Performance Metrics","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import DataLoader\nfrom sklearn.metrics import (confusion_matrix, classification_report, accuracy_score, \n                             f1_score, hamming_loss, jaccard_score, roc_auc_score, precision_recall_curve)\nimport timm  # For MaxViT\nfrom tqdm import tqdm\nfrom torchvision import transforms\n\n# Multi-label class names (24 diseases)\nclass_names = ['Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum',\n               'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation',\n               'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other',\n               'Fracture', 'Support_Devices', 'Aortic_Enlargement', 'Calcification',\n               'ILD', 'Infiltration', 'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay']\n\nNUM_LABELS = len(class_names)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Path to saved model\nmodel_path = '/kaggle/input/weights-after-3-epochs/best_model.pt'\n\n# Load MaxViT (224x224) with correct pretraining\nmaxvit = timm.create_model('maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k', pretrained=True)\nnum_ftrs = maxvit.head.fc.in_features\nmaxvit.head.fc = nn.Linear(num_ftrs, NUM_LABELS)\n\n# Unfreeze deeper layers\nfor name, param in maxvit.named_parameters():\n    param.requires_grad = \"blocks\" in name or \"head.fc\" in name\n\nmaxvit.to(device)\nmaxvit.eval()\n\n# Load model weights\ncheckpoint = torch.load(model_path, map_location=device)\nstate_dict = checkpoint.get('model_state_dict', checkpoint)\n\n# Remove DataParallel prefixes if needed\nif any(k.startswith(\"module.\") for k in state_dict.keys()):\n    from collections import OrderedDict\n    state_dict = OrderedDict((k.replace(\"module.\", \"\"), v) for k, v in state_dict.items())\n\nmissing_keys, unexpected_keys = maxvit.load_state_dict(state_dict, strict=False)\nprint(\"Missing keys:\", missing_keys)\nprint(\"Unexpected keys:\", unexpected_keys)\n\n# Safe collate function\ndef safe_collate(batch):\n    batch = [b for b in batch if b is not None]\n    if not batch:\n        return None\n    img_shapes = [b[0].shape for b in batch]\n    if len(set(img_shapes)) > 1:\n        return None\n    return torch.utils.data.dataloader.default_collate(batch)\n\n# Model Evaluation\nall_preds, all_labels, all_probs = [], [], []\nthresh = 0.5\n\n# Assuming test_loader is defined elsewhere\nnum_batches = len(test_loader)\n\n# Evaluate batches sequentially, wrapped in torch.no_grad() to reduce GPU memory overhead\nwith torch.no_grad():\n    with tqdm(total=num_batches, desc=\"Evaluating Batches\", unit=\"batch\") as batch_pbar:\n        for batch_idx, batch in enumerate(test_loader):\n            if batch is None:\n                batch_pbar.update(1)\n                continue\n            try:\n                images, labels = batch\n                images, labels = images.to(device), labels.to(device)\n\n                outputs = maxvit(images)  # Logits\n                probs = torch.sigmoid(outputs)  # Convert logits to probabilities\n                preds = (probs > thresh).int()  # Apply threshold\n\n                # Ignore uncertain cases (0.5 values)\n                mask = (labels != 0.5).float()\n                preds = preds * mask\n                labels = labels * mask\n\n                all_preds.append(preds.cpu().numpy())\n                all_labels.append(labels.cpu().numpy())\n                all_probs.append(probs.cpu().numpy())\n            except Exception as e:\n                print(f\"🚨 Error processing batch {batch_idx}: {e}\")\n            batch_pbar.update(1)\n            \n    # Convert lists to NumPy arrays after processing all batches\n    all_preds = np.vstack(all_preds)\n    all_labels = np.vstack(all_labels)\n    all_probs = np.vstack(all_probs)\n\n# Classification Report\nreport = classification_report(all_labels, all_preds, target_names=class_names, zero_division=0)\nprint(\"\\nClassification Report:\\n\", report)\n\n# Confusion Matrices (processed sequentially)\nfig, axes = plt.subplots(6, 4, figsize=(20, 25))\nfig.suptitle(\"Confusion Matrices\", fontsize=16)\nwith tqdm(total=len(class_names), desc=\"Generating Confusion Matrices\", unit=\"class\") as cm_pbar:\n    for i, (cls, ax) in enumerate(zip(class_names, axes.flat)):\n        cm = confusion_matrix(all_labels[:, i], all_preds[:, i])\n        sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=[0, 1], yticklabels=[0, 1], ax=ax)\n        ax.set_title(cls)\n        ax.set_xlabel(\"Predicted\")\n        ax.set_ylabel(\"Actual\")\n        cm_pbar.update(1)\nplt.tight_layout(rect=[0, 0, 1, 0.97])\nplt.show()\n\n# AUC-ROC Curves (processed sequentially)\nfig, ax = plt.subplots(figsize=(10, 6))\nax.set_title(\"AUC-ROC Curves\")\nax.set_xlabel(\"False Positive Rate\")\nax.set_ylabel(\"True Positive Rate\")\nwith tqdm(total=len(class_names), desc=\"Computing AUC-ROC\", unit=\"class\") as auc_pbar:\n    for i, cls in enumerate(class_names):\n        if len(np.unique(all_labels[:, i])) < 2:\n            print(f\"Skipping AUC-ROC for {cls} (insufficient class variance)\")\n        else:\n            # Compute precision-recall curve for ROC (Note: Typically, ROC uses FPR and TPR)\n            fpr, tpr, _ = precision_recall_curve(all_labels[:, i], all_probs[:, i])\n            auc_score = roc_auc_score(all_labels[:, i], all_probs[:, i])\n            ax.plot(fpr, tpr, label=f\"{cls} (AUC={auc_score:.4f})\")\n        auc_pbar.update(1)\nax.legend()\nplt.show()\n\n# Precision-Recall Curves (processed sequentially)\nfig, axes = plt.subplots(6, 4, figsize=(20, 25))\nfig.suptitle(\"Precision-Recall Curves\", fontsize=16)\nwith tqdm(total=len(class_names), desc=\"Generating Precision-Recall Curves\", unit=\"class\") as pr_pbar:\n    for i, (cls, ax) in enumerate(zip(class_names, axes.flat)):\n        precision, recall, _ = precision_recall_curve(all_labels[:, i], all_probs[:, i])\n        ax.plot(recall, precision, marker='.')\n        ax.set_title(cls)\n        ax.set_xlabel(\"Recall\")\n        ax.set_ylabel(\"Precision\")\n        pr_pbar.update(1)\nplt.tight_layout(rect=[0, 0, 1, 0.97])\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-14T14:37:49.494379Z","iopub.execute_input":"2025-04-14T14:37:49.494758Z","iopub.status.idle":"2025-04-14T15:09:50.815512Z","shell.execute_reply.started":"2025-04-14T14:37:49.494728Z","shell.execute_reply":"2025-04-14T15:09:50.814509Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Explainable AI Integration with SHAP, GRAD-CAM ++ and ODIN for OOD ","metadata":{}},{"cell_type":"code","source":"import random\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import FancyBboxPatch\nimport matplotlib.cm as cm\nimport cv2\nimport numpy as np\nimport torch.nn.functional as F\nimport timm\nfrom tqdm import tqdm\n\n# -------------------- CONFIGURATION --------------------\nclass_names = [\n    'Normal', 'Pneumonia', 'Tuberculosis', 'Covid', 'Enlarged_Cardiomediastinum',\n    'Cardiomegaly', 'Lung_Opacity', 'Lung_Lesion', 'Edema', 'Consolidation',\n    'Atelectasis', 'Pneumothorax', 'Pleural_Effusion', 'Pleural_Other',\n    'Fracture', 'Support_Devices', 'Aortic_Enlargement', 'Calcification',\n    'ILD', 'Infiltration', 'Nodule_Mass', 'Other_Lesion', 'Pulmonary_Fibrosis', 'Non_XRay'\n]\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# -------------------- MODEL LOADING --------------------\nmodel_path = '/kaggle/input/weights-after-3-epochs/best_model.pt'\nmaxvit = timm.create_model('maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k', pretrained=True)\nnum_ftrs = maxvit.head.fc.in_features\nmaxvit.head.fc = nn.Linear(num_ftrs, len(class_names))\n\n# Freeze only blocks and classifier head\nfor name, param in maxvit.named_parameters():\n    param.requires_grad = 'blocks' in name or 'head.fc' in name\n\nmaxvit.to(device)\nmaxvit.eval()\ncheckpoint = torch.load(model_path, map_location=device)\nstate_dict = checkpoint.get('model_state_dict', checkpoint)\nmaxvit.load_state_dict(state_dict)\n\n# -------------------- ODIN OOD SCORE FUNCTION --------------------\ndef odin_ood_score(model, input_tensor, T=1000, epsilon=0.0014):\n    input_tensor = input_tensor.to(device)\n    input_tensor.requires_grad = True\n    outputs = model(input_tensor)\n    scaled = outputs / T\n    probs = torch.sigmoid(scaled)\n    entropy = -torch.sum(probs * torch.log(probs + 1e-6), dim=1)\n    entropy.backward()\n    perturb = epsilon * input_tensor.grad.sign()\n    perturbed = input_tensor - perturb\n    outputs_p = model(perturbed)\n    probs_p = torch.sigmoid(outputs_p / T)\n    ent_p = -torch.sum(probs_p * torch.log(probs_p + 1e-6), dim=1)\n    return ent_p.item()\n\n# -------------------- GRAD-CAM++ FUNCTIONS --------------------\ndef get_conv_layers(model):\n    return [m for m in model.modules() if isinstance(m, nn.Conv2d)]\n\n\ndef generate_grad_cam_plus_plus_for_layer(model, input_tensor, target_class, conv_layer):\n    feature_map, gradients = None, None\n    def forward_hook(module, inp, output):\n        nonlocal feature_map; feature_map = output.detach()\n    def backward_hook(module, grad_in, grad_out):\n        nonlocal gradients; gradients = grad_out[0].detach()\n    h_f = conv_layer.register_forward_hook(forward_hook)\n    h_b = conv_layer.register_full_backward_hook(backward_hook)\n    model.zero_grad()\n    out = model(input_tensor)\n    score = out[0, target_class]\n    score.backward(retain_graph=True)\n    alpha_num = gradients.pow(2)\n    alpha_denom = gradients.pow(2).mul(2) + (feature_map * gradients.pow(3)).sum(dim=(2,3), keepdim=True)\n    alpha_denom = torch.where(alpha_denom!=0, alpha_denom, torch.ones_like(alpha_denom))\n    alphas = alpha_num / (alpha_denom + 1e-7)\n    weights = (alphas * F.relu(gradients)).sum(dim=(2,3), keepdim=True)\n    cam = (weights * feature_map).sum(dim=1, keepdim=True)\n    cam = F.relu(cam)\n    cam = F.interpolate(cam, size=(224,224), mode='bilinear', align_corners=False)\n    cam = cam.squeeze().cpu().numpy()\n    if cam.max(): cam = cam / cam.max()\n    cam = cv2.GaussianBlur(cam, (5,5), 0)\n    h_f.remove(); h_b.remove()\n    return cam\n\n\ndef generate_multi_layer_grad_cam(model, input_tensor, target_class):\n    convs = get_conv_layers(model)\n    layers = convs[-2:] if len(convs)>=2 else convs[-1:]\n    maps = [generate_grad_cam_plus_plus_for_layer(model, input_tensor, target_class, l) for l in layers]\n    avg = np.mean(maps, axis=0)\n    return avg / avg.max() if avg.max() else avg\n\n# -------------------- GUIDED BACKPROPAGATION --------------------\nclass GuidedBackpropReLU(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, inp):\n        pos_mask = (inp>0).type_as(inp)\n        ctx.save_for_backward(inp, pos_mask)\n        return inp * pos_mask\n    @staticmethod\n    def backward(ctx, grad_out):\n        inp, pos_mask = ctx.saved_tensors\n        grad = grad_out.clone()\n        grad[inp<0] = 0; grad[grad_out<0] = 0\n        return grad\n\nclass GuidedBackpropReLUWrapper(nn.Module):\n    def forward(self, x): return GuidedBackpropReLU.apply(x)\n\ndef replace_relu_with_guided(model):\n    for name, module in model.named_modules():\n        if isinstance(module, nn.ReLU):\n            parent = model; parts = name.split('.')\n            for p in parts[:-1]: parent = getattr(parent, p)\n            setattr(parent, parts[-1], GuidedBackpropReLUWrapper())\n    return model\n\n\ndef guided_backprop(model, input_tensor, target_class):\n    gb = replace_relu_with_guided(model)\n    gb.to(device).eval()\n    inp = input_tensor.clone().detach().to(device)\n    inp.requires_grad = True\n    out = gb(inp)\n    score = out[0, target_class]\n    gb.zero_grad(); score.backward(retain_graph=True)\n    grads = inp.grad.detach().cpu().numpy()[0]\n    grads = np.maximum(grads,0)\n    grads = grads / (grads.max()+1e-8)\n    return np.mean(grads, axis=0)\n\n# -------------------- GUIDED GRAD-CAM --------------------\ndef compute_guided_gradcam(model, input_tensor, target_class, dilation_iterations=1, threshold_value=0.2):\n    cam = generate_multi_layer_grad_cam(model, input_tensor, target_class)\n    gb = guided_backprop(model, input_tensor, target_class)\n    gg = cam * gb\n    gg = cv2.dilate(gg, np.ones((3,3),np.uint8), iterations=dilation_iterations)\n    gg[gg < threshold_value] = 0\n    return gg / gg.max() if gg.max() else gg\n\n# -------------------- PREDICTION FUNCTION --------------------\ndef predict_images_with_report(testset, model, num_images=10, conf_threshold=0.75, ood_threshold=0.05):\n    results=[]\n    idxs = random.sample(range(len(testset)), num_images)\n    for i in tqdm(idxs, desc=\"Processing images\"):\n        img, true = testset[i]\n        inp = img.unsqueeze(0).to(device)\n        ood = odin_ood_score(model, inp)\n        if ood < ood_threshold:\n            results.append((img, true, ['Other_Findings'], [], [], None, 'OOD'))\n            continue\n        with torch.no_grad():\n            out = model(inp)\n            probs = torch.sigmoid(out).squeeze().cpu().numpy()\n            preds = [i for i,p in enumerate(probs) if p>conf_threshold]\n        labs = [class_names[i] for i in preds]\n        scores = [probs[i]*100 for i in preds]\n        cams = [generate_multi_layer_grad_cam(model, inp, i) for i in preds]\n        ggcs = [compute_guided_gradcam(model, inp, i) for i in preds]\n        results.append((img, true, labs, cams, ggcs, scores, None))\n    return results\n\n# -------------------- VISUALIZATION UTILITIES --------------------\n# Color map for classes\ncmap = cm.get_cmap('tab20', len(class_names))\ncolor_dict = {cls: cmap(i) for i,cls in enumerate(class_names)}\n\ndef map_true_labels(true_labels):\n    \"\"\"\n    Converts various true_labels formats to a list of class names.\n    Supports:\n     - Tensor of indices\n     - List of indices\n     - Multi-hot vector (list or tensor of length num_classes)\n     - List of class name strings\n    \"\"\"\n    # Convert tensor to list\n    if isinstance(true_labels, torch.Tensor):\n        true_labels = true_labels.detach().cpu().tolist()\n    # If list of strings, assume already class names\n    if isinstance(true_labels, list) and all(isinstance(t, str) for t in true_labels):\n        return true_labels\n    # If list/array of length == num_classes: treat as multi-hot\n    if isinstance(true_labels, (list, np.ndarray)) and len(true_labels) == len(class_names):\n        indices = [i for i, v in enumerate(true_labels) if isinstance(v, (int, float)) and v > 0]\n        return [class_names[i] for i in indices]\n    # If list of ints: direct indices\n    if isinstance(true_labels, list) and all(isinstance(t, int) for t in true_labels):\n        return [class_names[i] for i in true_labels]\n    # Fallback: convert all entries to string\n    return [str(t) for t in true_labels]\n\ndef normalize_image(image):\n    img = image.permute(1,2,0).cpu().numpy()\n    return np.clip(img * np.array([0.229,0.224,0.225]) + np.array([0.485,0.456,0.406]),0,1)\n\ndef extract_bounding_boxes(heatmap, area_threshold=50):\n    hm = np.uint8(255*heatmap)\n    _, bm = cv2.threshold(hm,0,255,cv2.THRESH_BINARY+cv2.THRESH_OTSU)\n    bm = cv2.morphologyEx(bm, cv2.MORPH_CLOSE, np.ones((3,3),np.uint8))\n    cnts,_ = cv2.findContours(bm,cv2.RETR_EXTERNAL,cv2.CHAIN_APPROX_SIMPLE)\n    boxes=[]\n    for c in cnts:\n        if cv2.contourArea(c)>area_threshold:\n            x,y,w,h = cv2.boundingRect(c)\n            boxes.append({'x':x,'y':y,'w':w,'h':h})\n    return boxes\n\ndef visualize_results(image, true_labels, predicted_labels, grad_cam_maps, conf_scores, cols=2):\n    true_names = map_true_labels(true_labels)\n    img = normalize_image(image)\n    n_pred = len(predicted_labels)\n    total = 1 + n_pred + 1\n    rows = int(np.ceil(total/cols))\n    fig, axes = plt.subplots(rows, cols, figsize=(cols*4, rows*4))\n    axes = axes.flatten()\n    # Original\n    ax = axes[0]\n    ax.imshow(img); ax.set_title(\"True: \" + \", \".join(true_names), fontsize=12, fontweight='bold'); ax.axis('off')\n    # Per-class CAM\n    for i,(disease, hm, conf) in enumerate(zip(predicted_labels, grad_cam_maps, conf_scores), start=1):\n        ax = axes[i]; ax.imshow(img); ax.imshow(hm, cmap='jet', alpha=0.4)\n        color = color_dict.get(disease, (1,1,0,1))\n        for b in extract_bounding_boxes(hm):\n            fancy = FancyBboxPatch((b['x'],b['y']),b['w'],b['h'], boxstyle='round,pad=0.2', linewidth=2, edgecolor=color, facecolor='none')\n            ax.add_patch(fancy)\n            txt = f\"{disease} ({conf:.1f}%)\"\n            ax.text(b['x']+2, b['y']+15, txt, fontsize=10, fontweight='bold', bbox=dict(boxstyle='round,pad=0.2', facecolor=color, alpha=0.6))\n        ax.set_title(f\"Grad-CAM++: {disease}\", fontsize=12); ax.axis('off')\n    # Aggregated\n    agg_idx = 1 + n_pred\n    ax = axes[agg_idx]; ax.imshow(img)\n    label_positions=[]\n    for disease, hm, conf in zip(predicted_labels, grad_cam_maps, conf_scores):\n        color = color_dict.get(disease, (1,1,0,1))\n        for b in extract_bounding_boxes(hm):\n            fancy = FancyBboxPatch((b['x'],b['y']),b['w'],b['h'], boxstyle='round,pad=0.2', linewidth=2, edgecolor=color, facecolor='none')\n            ax.add_patch(fancy)\n            lx, ly = b['x']+2, b['y']+15\n            for px, py, pw, ph in label_positions:\n                if not (lx+pw < px or lx > px+pw or ly+ph < py or ly > py+ph): ly = py+ph+5\n            label_positions.append((lx, ly, 80, 15))\n            ax.text(lx, ly, f\"{disease} ({conf:.1f}%)\", fontsize=10, fontweight='bold', bbox=dict(boxstyle='round,pad=0.2', facecolor=color, alpha=0.6))\n    ax.set_title('Aggregated Boxes', fontsize=12); ax.axis('off')\n    # Hide extras\n    for j in range(total, len(axes)): axes[j].axis('off')\n    plt.tight_layout(); plt.show()\n\n# -------------------- USAGE --------------------\n# Ensure `testset` is defined: a dataset of (image_tensor, true_label_indices)\nresults = predict_images_with_report(testset, maxvit, num_images=1, conf_threshold=0.75)\nimg, true_lbls, preds, cams, gbp_maps, scores, _ = results[0]\nvisualize_results(img, true_lbls, preds, cams, scores, cols=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T05:00:12.688578Z","iopub.execute_input":"2025-04-15T05:00:12.688892Z","iopub.status.idle":"2025-04-15T05:00:20.623632Z","shell.execute_reply.started":"2025-04-15T05:00:12.688868Z","shell.execute_reply":"2025-04-15T05:00:20.622652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport math\n\ndef show_unique_xray_per_disease_fast(testset, class_names, cols=4, figsize_per_img=3, show_all_labels=False, max_samples=3000):\n    \"\"\"\n    Show one representative X-ray image for each disease class in a multi-label dataset.\n\n    Args:\n        trainset (Dataset): Your dataset, yielding (image, label) pairs.\n        class_names (List[str]): List of class names, ordered as per label indices.\n        cols (int): Number of columns in the output grid.\n        figsize_per_img (int): Size of each image in inches.\n        show_all_labels (bool): If True, shows all present labels per image. Else, only the main class.\n        max_samples (int): Max samples to scan through from dataset to speed things up.\n    \"\"\"\n    num_classes = len(class_names)\n    found = {}\n    found_labels = {}\n    used_images = set()\n\n    for i, (img, label) in enumerate(trainset):\n        if i >= max_samples or len(found) == num_classes:\n            break\n\n        if isinstance(label, torch.Tensor):\n            pos = (label > 0).nonzero(as_tuple=False).squeeze().tolist()\n            if isinstance(pos, int):\n                pos = [pos]\n        elif isinstance(label, (list, np.ndarray)) and len(label) == num_classes:\n            pos = [j for j, v in enumerate(label) if v > 0]\n        else:\n            continue\n\n        for c in pos:\n            if c not in found and i not in used_images:\n                found[c] = img\n                found_labels[c] = label\n                used_images.add(i)\n                break\n\n    # Plotting\n    total = len(found)\n    rows = math.ceil(total / cols)\n    fig, axes = plt.subplots(rows, cols,\n                             figsize=(cols * figsize_per_img, rows * figsize_per_img))\n    axes = axes.flatten()\n\n    for idx, class_idx in enumerate(sorted(found.keys())):\n        ax = axes[idx]\n        img = found[class_idx]\n        label = found_labels[class_idx]\n\n        # Unnormalize assuming ImageNet (adjust if needed)\n        img_np = img.permute(1, 2, 0).cpu().numpy()\n        img_np = np.clip(img_np * np.array([0.229, 0.224, 0.225]) +\n                         np.array([0.485, 0.456, 0.406]), 0, 1)\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport math\nfrom tqdm import tqdm  # Progress bar\n\ndef show_unique_xray_per_disease_fast(trainset, class_names, cols=4, figsize_per_img=3, show_all_labels=False, max_samples=3000):\n    \"\"\"\n    Show one representative X-ray image for each disease class in a multi-label dataset.\n    Displays progress bar while scanning dataset.\n    \"\"\"\n    num_classes = len(class_names)\n    found = {}\n    found_labels = {}\n    used_images = set()\n\n    print(\"Scanning dataset for representative samples...\")\n    for i, (img, label) in enumerate(tqdm(trainset, total=min(len(trainset), max_samples))):\n        if i >= max_samples or len(found) == num_classes:\n            break\n\n        if isinstance(label, torch.Tensor):\n            pos = (label > 0).nonzero(as_tuple=False).squeeze().tolist()\n            if isinstance(pos, int):\n                pos = [pos]\n        elif isinstance(label, (list, np.ndarray)) and len(label) == num_classes:\n            pos = [j for j, v in enumerate(label) if v > 0]\n        else:\n            continue\n\n        for c in pos:\n            if c not in found and i not in used_images:\n                found[c] = img\n                found_labels[c] = label\n                used_images.add(i)\n                break\n\n    # Plotting\n    total = len(found)\n    rows = math.ceil(total / cols)\n    fig, axes = plt.subplots(rows, cols,\n                             figsize=(cols * figsize_per_img, rows * figsize_per_img))\n    axes = axes.flatten()\n\n    for idx, class_idx in enumerate(sorted(found.keys())):\n        ax = axes[idx]\n        img = found[class_idx]\n        label = found_labels[class_idx]\n\n        # Unnormalize assuming ImageNet stats\n        img_np = img.permute(1, 2, 0).cpu().numpy()\n        img_np = np.clip(img_np * np.array([0.229, 0.224, 0.225]) +\n                         np.array([0.485, 0.456, 0.406]), 0, 1)\n\n        ax.imshow(img_np, cmap='gray')\n        if show_all_labels:\n            present_labels = [class_names[i] for i in range(num_classes) if label[i] > 0]\n            title = f\"{class_names[class_idx]}\\n({', '.join(present_labels)})\"\n        else:\n            title = class_names[class_idx]\n\n        ax.set_title(title, fontsize=9)\n        ax.axis('off')\n\n    # Hide unused subplots\n    for j in range(total, len(axes)):\n        axes[j].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n\nshow_unique_xray_per_disease_fast(trainset, class_names, cols=4, show_all_labels=True, max_samples=15000)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T05:23:47.271648Z","iopub.execute_input":"2025-04-15T05:23:47.272019Z","iopub.status.idle":"2025-04-15T05:34:43.622268Z","shell.execute_reply.started":"2025-04-15T05:23:47.271992Z","shell.execute_reply":"2025-04-15T05:34:43.620764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}