{"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":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"}],"dockerImageVersionId":30887,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom PIL import Image\nimport random\nimport seaborn as sns\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:37:58.573744Z","iopub.execute_input":"2025-02-15T04:37:58.574026Z","iopub.status.idle":"2025-02-15T04:38:00.859887Z","shell.execute_reply.started":"2025-02-15T04:37:58.574003Z","shell.execute_reply":"2025-02-15T04:38:00.859233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_and_display_samples(train_df, image_dir, n_samples=5):\n    \"\"\"\n    Load and display sample images for each cancer subtype\n    \"\"\"\n    # Set style for better visualization\n    plt.style.use('seaborn')\n    \n    # Get unique subtypes\n    subtypes = train_df['label'].unique()\n    \n    # Create a figure\n    plt.figure(figsize=(20, 4*len(subtypes)))\n    \n    # For each subtype\n    for idx, subtype in enumerate(subtypes):\n        # Get sample images for this subtype\n        subtype_df = train_df[train_df['label'] == subtype]\n        valid_images = 0\n        \n        # Keep sampling until we get enough valid images\n        for _, row in subtype_df.sample(frac=1).iterrows():  # Shuffle and iterate\n            if valid_images >= n_samples:\n                break\n                \n            image_path = os.path.join(image_dir, f\"{str(row['image_id'])}_thumbnail.png\")\n            \n            try:\n                if os.path.exists(image_path):\n                    img = Image.open(image_path)\n                    plt.subplot(len(subtypes), n_samples, idx*n_samples + valid_images + 1)\n                    plt.imshow(img)\n                    plt.axis('off')\n                    if valid_images == 0:  # Only add label for first image in row\n                        plt.title(f'{subtype}\\n(n={len(subtype_df)})', fontsize=12, pad=20)\n                    valid_images += 1\n            except:\n                continue\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:38:04.034356Z","iopub.execute_input":"2025-02-15T04:38:04.0348Z","iopub.status.idle":"2025-02-15T04:38:04.042008Z","shell.execute_reply.started":"2025-02-15T04:38:04.034774Z","shell.execute_reply":"2025-02-15T04:38:04.041051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Load the training data\n    train_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n    \n    # Print initial data info\n    print(\"Dataset Overview:\")\n    print(f\"Total number of images: {len(train_df)}\")\n    print(\"\\nDistribution of subtypes:\")\n    print(train_df['label'].value_counts())\n    print(\"\\nSample of image IDs:\")\n    print(train_df['image_id'].head())\n    \n    # Define image directory \n    train_image_dir = '/kaggle/input/UBC-OCEAN/train_thumbnails'\n    \n    # Verify directory exists\n    if not os.path.exists(train_image_dir):\n        print(f\"Error: Directory '{train_image_dir}' not found\")\n        return\n    \n    # Display visualization\n    print(\"\\nDisplaying cancer subtype samples...\")\n    load_and_display_samples(train_df, train_image_dir)\n    \n    # Display image size statistics\n    print(\"\\nImage size statistics:\")\n    print(\"\\nWidth:\")\n    print(train_df['image_width'].describe())\n    print(\"\\nHeight:\")\n    print(train_df['image_height'].describe())\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:38:07.120973Z","iopub.execute_input":"2025-02-15T04:38:07.12137Z","iopub.status.idle":"2025-02-15T04:38:29.228758Z","shell.execute_reply.started":"2025-02-15T04:38:07.121337Z","shell.execute_reply":"2025-02-15T04:38:29.222471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:12.979865Z","iopub.execute_input":"2025-02-15T04:39:12.980229Z","iopub.status.idle":"2025-02-15T04:39:13.002938Z","shell.execute_reply.started":"2025-02-15T04:39:12.980201Z","shell.execute_reply":"2025-02-15T04:39:13.002231Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:17.882712Z","iopub.execute_input":"2025-02-15T04:39:17.883101Z","iopub.status.idle":"2025-02-15T04:39:17.887721Z","shell.execute_reply.started":"2025-02-15T04:39:17.883069Z","shell.execute_reply":"2025-02-15T04:39:17.886842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def perform_eda(train_df):\n    \"\"\"\n    Perform comprehensive EDA on the UBC-OCEAN dataset\n    \"\"\"\n    # Set the style for better visualizations\n    plt.style.use('seaborn')\n    \n    # 1. Class Distribution Analysis\n    plt.figure(figsize=(10, 6))\n    sns.barplot(x=train_df['label'].value_counts().index, \n                y=train_df['label'].value_counts().values)\n    plt.title('Distribution of Cancer Subtypes')\n    plt.xlabel('Subtype')\n    plt.ylabel('Count')\n    plt.xticks(rotation=45)\n    #plt.savefig(\"class_distribution.png\")\n    plt.show()\n    \n    # 2. Image Size Distribution\n    plt.figure(figsize=(15, 5))\n    \n    # Width distribution\n    plt.subplot(1, 2, 1)\n    sns.histplot(data=train_df, x='image_width', bins=30)\n    plt.title('Distribution of Image Widths')\n    plt.xlabel('Width (pixels)')\n    \n    # Height distribution\n    plt.subplot(1, 2, 2)\n    sns.histplot(data=train_df, x='image_height', bins=30)\n    plt.title('Distribution of Image Heights')\n    plt.xlabel('Height (pixels)')\n    plt.tight_layout()\n    plt.show()\n    \n    # 3. TMA vs WSI Analysis\n    plt.figure(figsize=(10, 5))\n    sns.countplot(data=train_df, x='label', hue='is_tma')\n    plt.title('Distribution of TMA vs WSI across Subtypes')\n    plt.xlabel('Subtype')\n    plt.ylabel('Count')\n    plt.xticks(rotation=45)\n    plt.legend(title='Is TMA')\n    plt.tight_layout()\n    #plt.savefig(\"tma_wsi.png\")\n    plt.show()\n    \n    # 4. Image Aspect Ratio Analysis\n    train_df['aspect_ratio'] = train_df['image_width'] / train_df['image_height']\n    \n    plt.figure(figsize=(10, 6))\n    sns.boxplot(data=train_df, x='label', y='aspect_ratio')\n    plt.title('Image Aspect Ratios by Subtype')\n    plt.xlabel('Subtype')\n    plt.ylabel('Aspect Ratio (Width/Height)')\n    plt.xticks(rotation=45)\n    plt.tight_layout()\n    #plt.savefig(\"image_ratio.png\")\n    plt.show()\n    \n    # 5. Print Statistical Summary\n    print(\"\\nStatistical Summary by Subtype:\")\n    summary_stats = train_df.groupby('label').agg({\n        'image_width': ['mean', 'std', 'min', 'max'],\n        'image_height': ['mean', 'std', 'min', 'max'],\n        'is_tma': 'sum'\n    }).round(2)\n    print(summary_stats)\n    \n    # 6. Image Size Scatter Plot\n    plt.figure(figsize=(10, 8))\n    sns.scatterplot(data=train_df, x='image_width', y='image_height', \n                    hue='label', style='is_tma', alpha=0.6)\n    plt.title('Image Dimensions by Subtype and Type (TMA vs WSI)')\n    plt.xlabel('Width (pixels)')\n    plt.ylabel('Height (pixels)')\n    plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')\n    plt.tight_layout()\n    #plt.savefig(\"scatter_plot.png\")\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:20.965629Z","iopub.execute_input":"2025-02-15T04:39:20.965911Z","iopub.status.idle":"2025-02-15T04:39:20.975718Z","shell.execute_reply.started":"2025-02-15T04:39:20.96589Z","shell.execute_reply":"2025-02-15T04:39:20.97484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Load the data\n    train_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n    \n    print(\"Starting Exploratory Data Analysis...\")\n    print(\"\\nDataset Overview:\")\n    print(f\"Total number of samples: {len(train_df)}\")\n    print(f\"Number of unique subtypes: {train_df['label'].nunique()}\")\n    print(f\"Number of TMA images: {train_df['is_tma'].sum()}\")\n    print(f\"Number of WSI images: {(~train_df['is_tma']).sum()}\")\n    \n    # Perform EDA\n    perform_eda(train_df)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:25.664392Z","iopub.execute_input":"2025-02-15T04:39:25.664817Z","iopub.status.idle":"2025-02-15T04:39:27.048217Z","shell.execute_reply.started":"2025-02-15T04:39:25.664788Z","shell.execute_reply":"2025-02-15T04:39:27.047192Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Outlier Detection Analysis","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom scipy import stats\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:32.267076Z","iopub.execute_input":"2025-02-15T04:39:32.267473Z","iopub.status.idle":"2025-02-15T04:39:32.271614Z","shell.execute_reply.started":"2025-02-15T04:39:32.267443Z","shell.execute_reply":"2025-02-15T04:39:32.270717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def detect_outliers(train_df):\n    \"\"\"\n    Perform comprehensive outlier detection analysis\n    \"\"\"\n    # Create figure for multiple plots\n    plt.figure(figsize=(20, 15))\n    \n    # 1. Z-score based outlier detection for image dimensions\n    train_df['width_zscore'] = np.abs(stats.zscore(train_df['image_width']))\n    train_df['height_zscore'] = np.abs(stats.zscore(train_df['image_height']))\n    \n    # 2. Calculate additional features for outlier detection\n    train_df['area'] = train_df['image_width'] * train_df['image_height']\n    train_df['aspect_ratio'] = train_df['image_width'] / train_df['image_height']\n    \n    # Plot 1: Area vs Aspect Ratio with outlier boundaries\n    plt.subplot(2, 2, 1)\n    sns.scatterplot(data=train_df, x='area', y='aspect_ratio', hue='label', alpha=0.6)\n    plt.title('Image Area vs Aspect Ratio\\nPotential Outliers Detection')\n    plt.xlabel('Area (pixels²)')\n    plt.ylabel('Aspect Ratio')\n    \n    # Plot 2: Width vs Height with Outlier Boundaries\n    plt.subplot(2, 2, 2)\n    sns.scatterplot(data=train_df, x='width_zscore', y='height_zscore', \n                    hue='label', alpha=0.6)\n    plt.axhline(y=3, color='r', linestyle='--', alpha=0.3)\n    plt.axvline(x=3, color='r', linestyle='--', alpha=0.3)\n    plt.title('Z-scores of Width vs Height\\nRed lines indicate z-score = 3')\n    plt.xlabel('Width Z-score')\n    plt.ylabel('Height Z-score')\n    \n    # Plot 3: Box plot of image areas by subtype\n    plt.subplot(2, 2, 3)\n    sns.boxplot(data=train_df, x='label', y='area')\n    plt.title('Distribution of Image Areas by Subtype')\n    plt.xticks(rotation=45)\n    plt.ylabel('Area (pixels²)')\n    \n    # Plot 4: Density plot of aspect ratios\n    plt.subplot(2, 2, 4)\n    sns.kdeplot(data=train_df, x='aspect_ratio', hue='label')\n    plt.title('Density Distribution of Aspect Ratios')\n    plt.xlabel('Aspect Ratio')\n    \n    plt.tight_layout()\n    #plt.savefig(\"outlier.png\")\n    plt.show()\n    \n    # Print statistical outliers\n    print(\"\\nPotential Outliers Analysis:\")\n    \n    # Z-score based outliers (|z| > 3)\n    width_outliers = train_df[train_df['width_zscore'] > 3]\n    height_outliers = train_df[train_df['height_zscore'] > 3]\n    \n    print(f\"\\nImages with unusual width (Z-score > 3): {len(width_outliers)}\")\n    print(f\"Images with unusual height (Z-score > 3): {len(height_outliers)}\")\n    \n    # IQR based outlier detection for area\n    Q1 = train_df['area'].quantile(0.25)\n    Q3 = train_df['area'].quantile(0.75)\n    IQR = Q3 - Q1\n    area_outliers = train_df[(train_df['area'] < (Q1 - 1.5 * IQR)) | \n                            (train_df['area'] > (Q3 + 1.5 * IQR))]\n    \n    print(f\"\\nImages with unusual area (IQR method): {len(area_outliers)}\")\n    \n    # Print summary of extreme cases\n    print(\"\\nExtreme Cases Summary:\")\n    extremes = pd.DataFrame({\n        'Metric': ['Smallest Area', 'Largest Area', 'Most Square', 'Least Square'],\n        'Image ID': [\n            train_df.loc[train_df['area'].idxmin(), 'image_id'],\n            train_df.loc[train_df['area'].idxmax(), 'image_id'],\n            train_df.loc[(train_df['aspect_ratio'] - 1).abs().idxmin(), 'image_id'],\n            train_df.loc[(train_df['aspect_ratio'] - 1).abs().idxmax(), 'image_id']\n        ],\n        'Subtype': [\n            train_df.loc[train_df['area'].idxmin(), 'label'],\n            train_df.loc[train_df['area'].idxmax(), 'label'],\n            train_df.loc[(train_df['aspect_ratio'] - 1).abs().idxmin(), 'label'],\n            train_df.loc[(train_df['aspect_ratio'] - 1).abs().idxmax(), 'label']\n        ]\n    })\n    print(extremes)\n    \n    return width_outliers, height_outliers, area_outliers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:35.306489Z","iopub.execute_input":"2025-02-15T04:39:35.30679Z","iopub.status.idle":"2025-02-15T04:39:35.318532Z","shell.execute_reply.started":"2025-02-15T04:39:35.306766Z","shell.execute_reply":"2025-02-15T04:39:35.317647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Load the data\n    train_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n    \n    print(\"Starting Outlier Detection Analysis...\")\n    width_outliers, height_outliers, area_outliers = detect_outliers(train_df)\n    \n    # Save outlier information for future reference\n    outlier_summary = pd.DataFrame({\n        'image_id': list(set(width_outliers['image_id'].tolist() + \n                           height_outliers['image_id'].tolist() + \n                           area_outliers['image_id'].tolist())),\n        'is_outlier': True\n    })\n    \n    print(\"\\nTotal unique outliers detected:\", len(outlier_summary))\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:39.494079Z","iopub.execute_input":"2025-02-15T04:39:39.494434Z","iopub.status.idle":"2025-02-15T04:39:40.700873Z","shell.execute_reply.started":"2025-02-15T04:39:39.494407Z","shell.execute_reply":"2025-02-15T04:39:40.700195Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Class Imblanace Analysis","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:45.08941Z","iopub.execute_input":"2025-02-15T04:39:45.089729Z","iopub.status.idle":"2025-02-15T04:39:45.093406Z","shell.execute_reply.started":"2025-02-15T04:39:45.089703Z","shell.execute_reply":"2025-02-15T04:39:45.092612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_class_imbalance(train_df):\n    \"\"\"\n    Perform comprehensive class imbalance analysis\n    \"\"\"\n    # Set style\n    plt.style.use('seaborn')\n    \n    # Create figure for multiple plots\n    plt.figure(figsize=(15, 10))\n    \n    # 1. Class Distribution Plot\n    plt.subplot(2, 2, 1)\n    class_counts = train_df['label'].value_counts()\n    sns.barplot(x=class_counts.index, y=class_counts.values)\n    plt.title('Class Distribution')\n    plt.xlabel('Subtype')\n    plt.ylabel('Count')\n    \n    # Add count labels on top of bars\n    for i, v in enumerate(class_counts.values):\n        plt.text(i, v, str(v), ha='center', va='bottom')\n    \n    # 2. Percentage Distribution\n    plt.subplot(2, 2, 2)\n    class_percentages = (class_counts / len(train_df) * 100).round(2)\n    sns.barplot(x=class_percentages.index, y=class_percentages.values)\n    plt.title('Class Distribution (%)')\n    plt.xlabel('Subtype')\n    plt.ylabel('Percentage')\n    \n    # Add percentage labels on top of bars\n    for i, v in enumerate(class_percentages.values):\n        plt.text(i, v, f'{v:.1f}%', ha='center', va='bottom')\n    \n    # 3. Pie Chart\n    plt.subplot(2, 2, 3)\n    plt.pie(class_counts.values, labels=class_counts.index, autopct='%1.1f%%',\n            colors=sns.color_palette('husl', n_colors=len(class_counts)))\n    plt.title('Class Distribution (Pie Chart)')\n    \n    # 4. Imbalance Metrics Table\n    plt.subplot(2, 2, 4)\n    plt.axis('off')\n    \n    # Calculate imbalance metrics\n    majority_class = class_counts.max()\n    minority_class = class_counts.min()\n    imbalance_ratio = majority_class / minority_class\n    \n    metrics_text = (\n        f'Imbalance Analysis:\\n\\n'\n        f'Total Samples: {len(train_df)}\\n'\n        f'Number of Classes: {len(class_counts)}\\n'\n        f'Majority Class (HGSC): {majority_class}\\n'\n        f'Minority Class (MC): {minority_class}\\n'\n        f'Imbalance Ratio: {imbalance_ratio:.2f}:1\\n\\n'\n        f'Class Distribution:\\n'\n    )\n    \n    for class_name, percentage in class_percentages.items():\n        metrics_text += f'{class_name}: {percentage:.1f}%\\n'\n    \n    plt.text(0.1, 0.9, metrics_text, fontsize=10, va='top')\n    \n    plt.tight_layout()\n    #plt.savefig(\"class_imbalance.png\")\n    plt.show()\n    \n    # Print additional analysis\n    print(\"\\nDetailed Class Imbalance Analysis:\")\n    print(\"\\nClass Counts:\")\n    print(class_counts)\n    \n    print(\"\\nClass Percentages:\")\n    print(class_percentages)\n    \n    print(\"\\nImbalance Ratios (relative to majority class):\")\n    imbalance_ratios = majority_class / class_counts\n    print(imbalance_ratios)\n    \n    # Suggest potential strategies\n    print(\"\\nRecommended Strategies based on Imbalance:\")\n    if imbalance_ratio > 4:\n        print(\"- Consider using class weights in model\")\n        print(\"- Implement oversampling techniques (e.g., SMOTE) for minority classes\")\n        print(\"- Use stratified sampling in train/validation split\")\n    if imbalance_ratio > 2:\n        print(\"- Use balanced accuracy or F1-score as metrics\")\n        print(\"- Consider ensemble methods with balanced class weights\")\n    \n    return class_counts, class_percentages, imbalance_ratio","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:47.351105Z","iopub.execute_input":"2025-02-15T04:39:47.35144Z","iopub.status.idle":"2025-02-15T04:39:47.48653Z","shell.execute_reply.started":"2025-02-15T04:39:47.351414Z","shell.execute_reply":"2025-02-15T04:39:47.485617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Load the data\n    train_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n    \n    print(\"Starting Class Imbalance Analysis...\")\n    class_counts, class_percentages, imbalance_ratio = analyze_class_imbalance(train_df)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:39:52.811497Z","iopub.execute_input":"2025-02-15T04:39:52.811819Z","iopub.status.idle":"2025-02-15T04:39:53.302852Z","shell.execute_reply.started":"2025-02-15T04:39:52.811791Z","shell.execute_reply":"2025-02-15T04:39:53.302068Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Preprocessing and Data Preparation Pipeline","metadata":{}},{"cell_type":"code","source":"!pip install -q -U albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:40:02.608305Z","iopub.execute_input":"2025-02-15T04:40:02.608605Z","iopub.status.idle":"2025-02-15T04:40:08.698663Z","shell.execute_reply.started":"2025-02-15T04:40:02.608582Z","shell.execute_reply":"2025-02-15T04:40:08.69756Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q tiatoolbox","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:40:13.925452Z","iopub.execute_input":"2025-02-15T04:40:13.925799Z","iopub.status.idle":"2025-02-15T04:40:37.400603Z","shell.execute_reply.started":"2025-02-15T04:40:13.925767Z","shell.execute_reply":"2025-02-15T04:40:37.399635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport cv2\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom typing import List, Tuple, Dict, Optional\nfrom collections import Counter\nimport albumentations as A\nfrom imblearn.over_sampling import SMOTE\nfrom scipy.ndimage import gaussian_filter\nfrom torchvision import transforms as T\nfrom torchvision.transforms import InterpolationMode\nfrom torchvision import transforms\nimport torchvision.models as models\nfrom tiatoolbox import logger\nfrom tiatoolbox.tools import stainnorm, patchextraction\nfrom tiatoolbox.tools.stainaugment import StainAugmentor\nfrom imblearn.over_sampling import BorderlineSMOTE\nfrom imblearn.under_sampling import RandomUnderSampler\nfrom imblearn.pipeline import Pipeline\nfrom torch.utils.data import WeightedRandomSampler\nfrom tqdm import tqdm\nimport PIL.Image\nimport gc\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:40:41.981449Z","iopub.execute_input":"2025-02-15T04:40:41.981775Z","iopub.status.idle":"2025-02-15T04:41:18.922962Z","shell.execute_reply.started":"2025-02-15T04:40:41.981747Z","shell.execute_reply":"2025-02-15T04:41:18.922331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EnhancedPreprocessor:\n    def __init__(self,\n                 target_size: Tuple[int, int] = (224, 224), \n                 wsi_magnification: float = 20.0,\n                 tma_magnification: float = 40.0,\n                 stain_norm_method: str = 'reinhard'):\n        self.target_size = target_size\n        self.wsi_magnification = wsi_magnification\n        self.tma_magnification = tma_magnification\n        \n        # Initialize stain normalizer\n        if stain_norm_method == 'macenko':\n            self.normalizer = stainnorm.MacenkoNormalizer()\n        elif stain_norm_method == 'vahadane':\n            self.normalizer = stainnorm.VahadaneNormalizer()\n        elif stain_norm_method == 'reinhard':\n            self.normalizer = stainnorm.ReinhardNormalizer()\n        elif stain_norm_method == 'ruifrok':\n            self.normalizer = stainnorm.RuifrokNormalizer()\n        else:\n            raise ValueError(f\"Unknown stain normalization method: {stain_norm_method}\")\n\n    def detect_image_type(self, image: np.ndarray) -> str:\n        \"\"\"Determine if image is WSI or TMA based on size\"\"\"\n        height, width = image.shape[:2]\n        if height <= 5000 and width <= 5000:\n            return 'TMA'\n        return 'WSI'\n    \n    def normalize_magnification(self, image: np.ndarray, image_type: str) -> np.ndarray:\n        \"\"\"Normalize image magnification\"\"\"\n        if image_type == 'TMA':\n            scale_factor = self.wsi_magnification / self.tma_magnification\n            new_size = (int(image.shape[1] * scale_factor), \n                       int(image.shape[0] * scale_factor))\n            return cv2.resize(image, new_size, interpolation=cv2.INTER_AREA)\n        return image\n\n    def apply_stain_normalization(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"Apply stain normalization with error handling\"\"\"\n        try:\n            self.normalizer.fit(image)\n            normalized = self.normalizer.transform(image)\n            return normalized\n        except Exception as e:\n            print(f\"Error in stain normalization: {str(e)}\")\n            return image  # Return original image if normalization fails\n\n    def detect_tissue(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"Improved tissue detection using LAB color space and adaptive thresholding\"\"\"\n        # Convert to LAB color space\n        lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n        l_channel = lab[:, :, 0]\n        \n        # Adaptive thresholding\n        mask = cv2.adaptiveThreshold(\n            l_channel, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C,\n            cv2.THRESH_BINARY_INV, 11, 2\n        )\n        \n        # Morphological operations\n        kernel = np.ones((5, 5), np.uint8)\n        mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel)\n        mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN, kernel)\n        \n        return mask > 0\n\n    def extract_tissue_region(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"Extract main tissue region\"\"\"\n        try:\n            # Get tissue mask\n            tissue_mask = self.detect_tissue(image)\n            \n            # Find contours\n            contours, _ = cv2.findContours(tissue_mask.astype(np.uint8), \n                                         cv2.RETR_EXTERNAL, \n                                         cv2.CHAIN_APPROX_SIMPLE)\n            \n            if not contours:\n                return image\n            \n            # Find largest contour\n            largest_contour = max(contours, key=cv2.contourArea)\n            x, y, w, h = cv2.boundingRect(largest_contour)\n            \n            # Extract region with padding\n            pad = 10\n            x_start = max(0, x - pad)\n            y_start = max(0, y - pad)\n            x_end = min(image.shape[1], x + w + pad)\n            y_end = min(image.shape[0], y + h + pad)\n            \n            return image[y_start:y_end, x_start:x_end]\n            \n        except Exception as e:\n            print(f\"Error in tissue extraction: {str(e)}\")\n            return image\n\n    def handle_image_dimensions(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"Handle different image dimensions based on size\"\"\"\n        height, width = image.shape[:2]\n        \n        # Handle very large WSIs\n        if width > 50000 or height > 50000:\n            scale_factor = min(50000 / width, 50000 / height)\n            new_width = int(width * scale_factor)\n            new_height = int(height * scale_factor)\n            print(f\"Resizing large WSI from {width}x{height} to {new_width}x{new_height}\")\n            return cv2.resize(image, (new_width, new_height), interpolation=cv2.INTER_AREA)\n        \n        return image\n\n    def preprocess_image(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"Complete preprocessing pipeline\"\"\"\n        try:\n            # Handle large image dimensions first\n            image = self.handle_image_dimensions(image)\n            \n            # Determine image type\n            image_type = self.detect_image_type(image)\n            \n            # Normalize magnification\n            image = self.normalize_magnification(image, image_type)\n            \n            # Extract tissue region\n            image = self.extract_tissue_region(image)\n            \n            # Apply stain normalization\n            image = self.apply_stain_normalization(image)\n            \n            # Resize to target size\n            image = cv2.resize(image, self.target_size, interpolation=cv2.INTER_AREA)\n            \n            return image\n\n        except Exception as e:\n            print(f\"Error in preprocessing: {str(e)}\")\n            # Return resized original image as fallback\n            return cv2.resize(image, self.target_size, interpolation=cv2.INTER_AREA)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:41:53.189625Z","iopub.execute_input":"2025-02-15T04:41:53.190469Z","iopub.status.idle":"2025-02-15T04:41:53.641439Z","shell.execute_reply.started":"2025-02-15T04:41:53.190434Z","shell.execute_reply":"2025-02-15T04:41:53.640269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_transforms(image_size: Tuple[int, int] = (224, 224), stain_augment_prob: float = 0.5):\n    \"\"\"Create augmentation transforms with advanced techniques\"\"\"\n    train_transform = A.Compose([\n        A.Resize(height=image_size[0], width=image_size[1], always_apply=True),\n        # Color augmentations\n        A.OneOf([\n            A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=1.0),\n            A.RandomGamma(p=1.0)\n        ], p=0.5),\n        # Geometric augmentations\n        A.OneOf([\n            A.ElasticTransform(p=0.5),\n            A.GridDistortion(p=0.5),\n            A.OpticalDistortion(p=0.5)\n        ], p=0.3),\n        # Cutout augmentation\n        A.CoarseDropout(max_holes=8, max_height=16, max_width=16, p=0.5),\n        # Basic transforms\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5)\n    ])\n    \n    val_transform = A.Compose([\n        A.Resize(height=image_size[0], width=image_size[1], always_apply=True)\n    ])\n    \n    return train_transform, val_transform\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:41:58.912405Z","iopub.execute_input":"2025-02-15T04:41:58.912735Z","iopub.status.idle":"2025-02-15T04:41:58.919147Z","shell.execute_reply.started":"2025-02-15T04:41:58.912707Z","shell.execute_reply":"2025-02-15T04:41:58.918207Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HistoDataset(Dataset):\n    \"\"\"Enhanced dataset with robust preprocessing\"\"\"\n    def __init__(self, df: pd.DataFrame, image_dir: str,\n                 transform: Optional[transforms.Compose] = None,\n                 is_training: bool = True,\n                 apply_smote: bool = True,\n                 stain_norm_method: str = 'reinhard'):\n        \n        self.df = df.copy()  # Make a copy to prevent modifications\n        self.image_dir = image_dir\n        self.transform = transform\n        self.is_training = is_training\n        self.apply_smote = apply_smote and is_training\n        \n        # Verify required columns exist\n        required_columns = ['image_id', 'encoded_label']\n        if not all(col in self.df.columns for col in required_columns):\n            raise ValueError(f\"DataFrame must contain columns: {required_columns}\")\n        \n        # Initialize enhanced preprocessor\n        self.preprocessor = EnhancedPreprocessor(\n            stain_norm_method=stain_norm_method\n        )\n        \n        # Load and preprocess images\n        print(\"Loading and preprocessing images...\")\n        self._load_images()\n        \n        if self.apply_smote and len(self.processed_images) > 0:\n            self._apply_smote_preprocessing()\n    \n    def _load_images(self):\n        \"\"\"Load and preprocess all images\"\"\"\n        valid_indices = []\n        self.processed_images = []\n        self.labels = []\n        \n        for idx in tqdm(range(len(self.df)), desc=\"Processing images\"):\n            try:\n                image_path = os.path.join(self.image_dir,\n                                        f\"{self.df.iloc[idx]['image_id']}_thumbnail.png\")\n                if os.path.exists(image_path):\n                    # Load image\n                    image = cv2.imread(image_path)\n                    if image is None:\n                        print(f\"Warning: Could not read image - {image_path}\")\n                        continue\n                    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n                    \n                    # Preprocess image\n                    processed_image = self.preprocessor.preprocess_image(image)\n                    \n                    self.processed_images.append(processed_image)\n                    valid_indices.append(idx)\n                    self.labels.append(self.df.iloc[idx]['encoded_label'])\n                else:\n                    print(f\"Warning: Image not found - {image_path}\")\n            \n            except Exception as e:\n                print(f\"Error processing image at index {idx}: {str(e)}\")\n                continue\n        \n        if len(valid_indices) == 0:\n            raise ValueError(\"No valid images were loaded\")\n            \n        self.df = self.df.iloc[valid_indices].reset_index(drop=True)\n        self.labels = np.array(self.labels)\n        \n        print(f\"Successfully processed {len(self.processed_images)} images\")\n\n    def _apply_smote_preprocessing(self):\n        \"\"\"Apply Borderline-SMOTE and random undersampling\"\"\"\n        print(\"\\nApplying SMOTE and undersampling...\")\n        print(\"Class distribution before resampling:\")\n        print(self.df['label'].value_counts())\n        \n        # Reshape images for SMOTE\n        features = [img.reshape(-1) for img in self.processed_images]\n        features = np.array(features)\n        \n        # Define resampling pipeline\n        smote = BorderlineSMOTE(random_state=42)\n        under = RandomUnderSampler(random_state=42)\n        pipeline = Pipeline([('smote', smote), ('under', under)])\n        \n        # Apply resampling\n        features_resampled, labels_resampled = pipeline.fit_resample(features, self.labels)\n        \n        # Reconstruct images\n        self.images = [feat.reshape(self.preprocessor.target_size[0], \n                                  self.preprocessor.target_size[1], 3) \n                      for feat in features_resampled]\n        \n        # Create new balanced dataframe\n        new_data = []\n        for idx, label in enumerate(labels_resampled):\n            new_data.append({\n                'image_id': f'synthetic_{idx}' if idx >= len(self.df) else self.df.iloc[idx]['image_id'],\n                'label': self.df['label'].unique()[label],\n                'encoded_label': label,\n                'is_synthetic': idx >= len(self.df)\n            })\n        \n        self.df = pd.DataFrame(new_data)\n        print(\"\\nClass distribution after resampling:\")\n        print(self.df['label'].value_counts())\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        try:\n            if hasattr(self, 'images'):  # If SMOTE was applied\n                image = self.images[idx]\n                label = self.df.iloc[idx]['encoded_label']\n            else:  # Original image loading\n                image = self.processed_images[idx]\n                label = self.labels[idx]\n            \n            if self.transform:\n                transformed = self.transform(image=image)\n                image = transformed['image']\n            \n            # Convert to tensor\n            image = torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0\n            label = torch.tensor(label, dtype=torch.long)\n            \n            return image, label\n            \n        except Exception as e:\n            print(f\"Error loading image at index {idx}: {str(e)}\")\n            return torch.zeros((3, 224, 224)), torch.tensor(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:02.09592Z","iopub.execute_input":"2025-02-15T04:42:02.096349Z","iopub.status.idle":"2025-02-15T04:42:02.110838Z","shell.execute_reply.started":"2025-02-15T04:42:02.096313Z","shell.execute_reply":"2025-02-15T04:42:02.109982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_data(df: pd.DataFrame, image_dir: str, batch_size: int = 32):\n    \"\"\"Prepare data loaders with TIAToolbox preprocessing\"\"\"\n    try:\n        # Create label encodings\n        label_encoder = {'HGSC': 0, 'EC': 1, 'CC': 2, 'LGSC': 3, 'MC': 4}\n        df['encoded_label'] = df['label'].map(label_encoder)\n        \n        # Stratified split\n        train_df, val_df = train_test_split(\n            df,\n            test_size=0.2,\n            stratify=df['label'],\n            random_state=42\n        )\n        \n        # Create transforms\n        train_transform, val_transform = create_transforms(image_size=(224, 224), stain_augment_prob=0.5 )\n        \n        print(\"Creating training dataset...\")\n        train_dataset = HistoDataset(\n            df=train_df,\n            image_dir=image_dir,\n            transform=train_transform,\n            is_training=True,\n            apply_smote=True,\n            stain_norm_method='reinhard',# try 'macenko' or 'reinhard' or 'ruifrok'\n        )\n        \n        print(\"\\nCreating validation dataset...\")\n        val_dataset = HistoDataset(\n            df=val_df,\n            image_dir=image_dir,\n            transform=val_transform,\n            is_training=False,\n            apply_smote=False,\n            stain_norm_method='reinhard',\n        )\n        \n        # Create data loaders\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=batch_size,\n            shuffle=True,\n            num_workers=0,\n            pin_memory=True\n        )\n        \n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=batch_size,\n            shuffle=False,\n            num_workers=0,\n            pin_memory=True\n        )\n        \n        return train_loader, val_loader\n        \n    except Exception as e:\n        print(f\"Error in prepare_data: {str(e)}\")\n        raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:10.736227Z","iopub.execute_input":"2025-02-15T04:42:10.736538Z","iopub.status.idle":"2025-02-15T04:42:10.742788Z","shell.execute_reply.started":"2025-02-15T04:42:10.736513Z","shell.execute_reply":"2025-02-15T04:42:10.741923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_sample_images(loader):\n    \"\"\"Display sample images from the data loader\"\"\"\n    plt.figure(figsize=(15, 5))\n    images, labels = next(iter(loader))\n    for i in range(min(5, len(images))):\n        plt.subplot(1, 5, i + 1)\n        img = images[i].numpy().transpose(1, 2, 0)\n        img = np.clip(img, 0, 1)\n        plt.imshow(img)\n        plt.title(f'Label: {labels[i].item()}')\n        plt.axis('off')\n    plt.tight_layout()\n    plt.savefig(\"preprocessing.png\")\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:15.042524Z","iopub.execute_input":"2025-02-15T04:42:15.04282Z","iopub.status.idle":"2025-02-15T04:42:15.047869Z","shell.execute_reply.started":"2025-02-15T04:42:15.042797Z","shell.execute_reply":"2025-02-15T04:42:15.047044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_stain_normalization():\n    image = cv2.imread(\"/kaggle/input/UBC-OCEAN/train_thumbnails/10077_thumbnail.png\")\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    preprocessor = EnhancedPreprocessor(stain_norm_method='reinhard')\n    normalized_image = preprocessor.apply_stain_normalization(image)\n    \n    plt.figure(figsize=(10, 5))\n    plt.subplot(1, 2, 1)\n    plt.imshow(image)\n    plt.title(\"Original Image\")\n    plt.axis('off')\n    \n    plt.subplot(1, 2, 2)\n    plt.imshow(normalized_image)\n    plt.title(\"Normalized Image\")\n    plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n#if __name__ == \"__main__\":\n    #test_stain_normalization()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:16.812532Z","iopub.execute_input":"2025-02-15T04:42:16.812876Z","iopub.status.idle":"2025-02-15T04:42:16.817931Z","shell.execute_reply.started":"2025-02-15T04:42:16.812844Z","shell.execute_reply":"2025-02-15T04:42:16.817049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_tissue_detection():\n    image = cv2.imread(\"/kaggle/input/UBC-OCEAN/train_thumbnails/10896_thumbnail.png\")\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    preprocessor = EnhancedPreprocessor()\n    tissue_mask = preprocessor.detect_tissue(image)\n    \n    plt.figure(figsize=(10, 5))\n    plt.subplot(1, 2, 1)\n    plt.imshow(image)\n    plt.title(\"Original Image\")\n    plt.axis('off')\n    \n    plt.subplot(1, 2, 2)\n    plt.imshow(tissue_mask, cmap='gray')\n    plt.title(\"Tissue Mask\")\n    plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n#if __name__ == \"__main__\":\n    #test_tissue_detection()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:19.362336Z","iopub.execute_input":"2025-02-15T04:42:19.36268Z","iopub.status.idle":"2025-02-15T04:42:19.367722Z","shell.execute_reply.started":"2025-02-15T04:42:19.362649Z","shell.execute_reply":"2025-02-15T04:42:19.366908Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_data_augmentation():\n    image = cv2.imread(\"/kaggle/input/UBC-OCEAN/train_thumbnails/12222_thumbnail.png\")\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    train_transform, _ = create_transforms()\n    augmented = train_transform(image=image)['image']\n    \n    plt.figure(figsize=(10, 5))\n    plt.subplot(1, 2, 1)\n    plt.imshow(image)\n    plt.title(\"Original Image\")\n    plt.axis('off')\n    \n    plt.subplot(1, 2, 2)\n    plt.imshow(augmented)\n    plt.title(\"Augmented Image\")\n    plt.axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n#if __name__ == \"__main__\":\n    #test_data_augmentation()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:21.908133Z","iopub.execute_input":"2025-02-15T04:42:21.908477Z","iopub.status.idle":"2025-02-15T04:42:21.913534Z","shell.execute_reply.started":"2025-02-15T04:42:21.908451Z","shell.execute_reply":"2025-02-15T04:42:21.912633Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_smote_undersampling():\n    \"\"\"Test SMOTE and undersampling with proper label encoding\"\"\"\n    try:\n        # Load data\n        df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n        image_dir = '/kaggle/input/UBC-OCEAN/train_thumbnails'\n        \n        # Add label encoding\n        label_encoder = {'HGSC': 0, 'EC': 1, 'CC': 2, 'LGSC': 3, 'MC': 4}\n        df['encoded_label'] = df['label'].map(label_encoder)\n        \n        print(\"Initial class distribution:\")\n        print(df['label'].value_counts())\n        \n        # Create dataset with SMOTE\n        dataset = HistoDataset(\n            df=df,\n            image_dir=image_dir,\n            transform=None,\n            is_training=True,\n            apply_smote=True\n        )\n        \n        print(\"\\nFinal class distribution after SMOTE:\")\n        print(dataset.df['label'].value_counts())\n        \n    except Exception as e:\n        print(f\"Error in test_smote_undersampling: {str(e)}\")\n\n\n#if __name__ == \"__main__\":\n    #test_smote_undersampling()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:24.323192Z","iopub.execute_input":"2025-02-15T04:42:24.32352Z","iopub.status.idle":"2025-02-15T04:42:24.328756Z","shell.execute_reply.started":"2025-02-15T04:42:24.323496Z","shell.execute_reply":"2025-02-15T04:42:24.327868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main function to prepare and store data loaders in global scope\"\"\"\n    try:\n        # Load data\n        df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n        image_dir = '/kaggle/input/UBC-OCEAN/train_thumbnails'\n        \n        print(\"Initial class distribution:\")\n        print(df['label'].value_counts())\n        \n        # Create data loaders with full pipeline and store in global scope\n        global train_loader, val_loader\n        train_loader, val_loader = prepare_data(df, image_dir)\n        \n        # Test batch loading\n        images, labels = next(iter(train_loader))\n        print(f\"\\nBatch shapes:\")\n        print(f\"Images: {images.shape}\")\n        print(f\"Labels: {labels.shape}\")\n        \n        print(\"\\nClass distribution in batch:\")\n        print(pd.Series(labels.numpy()).value_counts())\n        \n        # Show sample images\n        print(\"\\nDisplaying sample processed images:\")\n        show_sample_images(train_loader)\n        \n        # Memory cleanup\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n    except Exception as e:\n        print(f\"Error in main: {str(e)}\")\n        # Memory cleanup even if there's an error\n        gc.collect()\n        torch.cuda.empty_cache()\n\n#if __name__ == \"__main__\":\n    #main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:27.04861Z","iopub.execute_input":"2025-02-15T04:42:27.048905Z","iopub.status.idle":"2025-02-15T04:42:27.054778Z","shell.execute_reply.started":"2025-02-15T04:42:27.048883Z","shell.execute_reply":"2025-02-15T04:42:27.053821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    # Test preprocessing components\n    test_stain_normalization()\n    test_tissue_detection()\n    test_data_augmentation()\n\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:42:29.916032Z","iopub.execute_input":"2025-02-15T04:42:29.916382Z","iopub.status.idle":"2025-02-15T04:45:33.707079Z","shell.execute_reply.started":"2025-02-15T04:42:29.91635Z","shell.execute_reply":"2025-02-15T04:45:33.706024Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Development","metadata":{}},{"cell_type":"code","source":"from tiatoolbox.models.architecture import vanilla\nfrom torchvision import models, transforms\nfrom torch import nn\nimport torch\nimport torch.nn.functional as F\nfrom pathlib import Path\nimport logging\nfrom typing import Dict, Optional\nimport timm\nimport PIL.Image\nfrom typing import Optional, Union, Dict\nfrom tiatoolbox.models.models_abc import ModelABC\nfrom torchvision.models import efficientnet_b0, EfficientNet_B0_Weights\nfrom torchvision.models import efficientnet_b3, EfficientNet_B3_Weights\nfrom torchvision.models import resnet101, ResNet101_Weights\nfrom torchvision.models import resnet152, ResNet152_Weights\nfrom tiatoolbox.models.engine.patch_predictor import PatchPredictor, IOPatchPredictorConfig\nimport gc\nimport math\nimport traceback\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:19.581173Z","iopub.execute_input":"2025-02-15T04:46:19.581533Z","iopub.status.idle":"2025-02-15T04:46:19.587717Z","shell.execute_reply.started":"2025-02-15T04:46:19.581495Z","shell.execute_reply":"2025-02-15T04:46:19.586883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HistoPathModel(nn.Module):\n    def __init__(self, num_classes=5):\n        super().__init__()\n        # Initialize both backbones\n        self.resnet = resnet101(weights=ResNet101_Weights.DEFAULT)\n        self.efficientnet = efficientnet_b3(weights=EfficientNet_B3_Weights.DEFAULT)\n        \n        # Remove original classifier layers\n        self.resnet.fc = nn.Identity()\n        self.efficientnet.classifier = nn.Identity()\n        \n        # Get feature dimensions\n        self.resnet_dim = 2048  # ResNet101's output dimension\n        self.efficient_dim = 1536  # EfficientNet-B3's output dimension\n        self.feature_dim = 512\n        \n        # Calculate combined features dimension for ResNet\n        # Global features (2048) + 3 processors (512 each) = 3584\n        self.combined_resnet_dim = self.resnet_dim + (512 * 3)\n        \n        # Feature reduction layers\n        self.resnet_reducer = nn.Sequential(\n            nn.Linear(self.combined_resnet_dim, self.feature_dim),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        \n        self.efficient_reducer = nn.Sequential(\n            nn.Linear(self.efficient_dim, self.feature_dim * 2),\n            nn.BatchNorm1d(self.feature_dim * 2),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(self.feature_dim * 2, self.feature_dim)\n        )\n        \n        # Modify first conv layer for histology images\n        self.resnet.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        \n        # Spatial attention module\n        self.spatial_attention = nn.Sequential(\n            nn.Conv2d(2048, 512, kernel_size=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n        # Channel attention module\n        self.channel_attention = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(2048, 512, kernel_size=1),\n            nn.ReLU(),\n            nn.Conv2d(512, 2048, kernel_size=1),\n            nn.Sigmoid()\n        )\n        # Add channel attention for EfficientNet\n        self.efficient_channel_attention = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(1536, 512, kernel_size=1),  # 1536 for EfficientNet-B3\n            nn.ReLU(),\n            nn.Conv2d(512, 1536, kernel_size=1),\n            nn.Sigmoid()\n        )\n\n        # Add fusion attention\n        self.fusion_attention = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim // 16),\n            nn.ReLU(),\n            nn.Linear(self.feature_dim // 16, self.feature_dim),\n            nn.Sigmoid()\n        )\n\n                \n        # Feature processors for ResNet features\n        self.feature_processors = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(2048, 512, 1),\n                nn.BatchNorm2d(512),\n                nn.ReLU()\n            ),\n            nn.Sequential(\n                nn.Conv2d(2048, 512, 3, padding=1),\n                nn.BatchNorm2d(512),\n                nn.ReLU()\n            ),\n            nn.Sequential(\n                nn.Conv2d(2048, 512, 5, padding=2),\n                nn.BatchNorm2d(512),\n                nn.ReLU()\n            )\n        ])\n        \n        self.feature_fusion = nn.Sequential(\n            nn.Linear(self.feature_dim * 2, self.feature_dim * 2),\n            nn.LayerNorm(self.feature_dim * 2),\n            nn.ReLU(),\n            nn.Dropout(0.2),  # Reduced dropout for stability\n            nn.Linear(self.feature_dim * 2, self.feature_dim * 2),\n            nn.LayerNorm(self.feature_dim * 2),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(self.feature_dim * 2, self.feature_dim)\n        )\n    \n  \n        # Add feature gates\n        self.feature_gates = nn.Sequential(\n            nn.Linear(self.feature_dim * 2, 2),\n            nn.Softmax(dim=1)\n        )\n\n    \n        # Add a squeeze-excitation block\n        self.se_block = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim // 16),\n            nn.ReLU(),\n            nn.Linear(self.feature_dim // 16, self.feature_dim),\n            nn.Sigmoid()\n         )\n    \n        \n        # Self-attention for global context\n        self.self_attention = nn.MultiheadAttention(\n            embed_dim=self.feature_dim,\n            num_heads=8,\n            dropout=0.1,\n            batch_first=True\n        )\n        \n        # Main classifier\n        self.main_classifier = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim),\n            nn.LayerNorm(self.feature_dim),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(self.feature_dim, num_classes)\n        )\n        \n        # Auxiliary classifier\n        self.aux_classifier = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim),\n            nn.LayerNorm(self.feature_dim),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(self.feature_dim, num_classes)\n        )\n\n    def extract_resnet_features(self, x):\n        # Initial layers\n        x = self.resnet.conv1(x)\n        x = self.resnet.bn1(x)\n        x = self.resnet.relu(x)\n        x = self.resnet.maxpool(x)\n        \n        # ResNet blocks\n        x = self.resnet.layer1(x)\n        x = self.resnet.layer2(x)\n        x = self.resnet.layer3(x)\n        x = self.resnet.layer4(x)  # Shape: [B, 2048, H, W]\n        \n        # Apply attention\n        spatial_weights = self.spatial_attention(x)\n        channel_weights = self.channel_attention(x)\n        attended_features = x * spatial_weights * channel_weights\n        \n        # Global features\n        global_features = F.adaptive_avg_pool2d(attended_features, 1).flatten(1)\n        \n        # Process features through each processor\n        processed_features = []\n        for processor in self.feature_processors:\n            features = processor(attended_features)\n            pooled = F.adaptive_avg_pool2d(features, 1).flatten(1)\n            processed_features.append(pooled)\n        \n        # Concatenate global and processed features\n        combined_features = torch.cat([global_features] + processed_features, dim=1)\n        \n        # Reduce features\n        return self.resnet_reducer(combined_features)\n\n    def extract_efficient_features(self, x):\n        features = self.efficientnet.features(x)\n    \n        # Apply channel attention\n        channel_weights = self.efficient_channel_attention(features)\n        features = features * channel_weights\n    \n        features = self.efficientnet.avgpool(features)\n        features = torch.flatten(features, 1)\n        return self.efficient_reducer(features)\n    \n\n    def extract_features(self, x):\n        resnet_features = self.extract_resnet_features(x)\n        torch.cuda.empty_cache()\n    \n        efficient_features = self.extract_efficient_features(x)\n        torch.cuda.empty_cache()\n    \n        combined_features = torch.cat([resnet_features, efficient_features], dim=1)\n        fused_features = self.feature_fusion(combined_features)\n    \n        # Residual connection\n        if hasattr(self, 'feature_gates'):\n            gates = self.feature_gates(combined_features)\n            residual = resnet_features * gates[:, 0].unsqueeze(1) + efficient_features * gates[:, 1].unsqueeze(1)\n            fused_features = fused_features + residual\n    \n        \n        # Apply self-attention with gradient clipping\n        attended_features, _ = self.self_attention(\n            fused_features.unsqueeze(1),\n            fused_features.unsqueeze(1),\n            fused_features.unsqueeze(1)\n        )\n        \n        return fused_features + 0.1 * attended_features.squeeze(1)\n        \n    \n    def forward(self, x):\n        # Extract combined features\n        features = self.extract_features(x)\n        \n        if self.training:\n            # During training, return both main and auxiliary outputs\n            main_logits = self.main_classifier(features)\n            aux_logits = self.aux_classifier(features)\n            return main_logits, aux_logits\n        else:\n            # During inference, only return main classifier output\n            return self.main_classifier(features)\n\n    def __del__(self):\n        # Clean up CUDA memory\n        torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:22.276023Z","iopub.execute_input":"2025-02-15T04:46:22.276385Z","iopub.status.idle":"2025-02-15T04:46:22.29407Z","shell.execute_reply.started":"2025-02-15T04:46:22.276354Z","shell.execute_reply":"2025-02-15T04:46:22.293234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_model(device, learning_rate=5e-4):\n    \"\"\"Create model and associated components.\"\"\"\n    # Create enhanced model\n    model = HistoPathModel(num_classes=5)\n    model = model.to(device)\n    \n    # Create predictor configuration\n    wsi_config = IOPatchPredictorConfig(\n        input_resolutions=[{\"units\": \"mpp\", \"resolution\": 0.5}],\n        patch_input_shape=[224, 224],\n        stride_shape=[224, 224]\n    )\n    \n    # Create patch predictor\n    predictor = PatchPredictor(\n        model=model,\n        batch_size=32,\n        num_loader_workers=4\n    )\n    \n    # Loss function with class weights\n    #class_weights = torch.tensor([1.0, 1.8, 2.2, 4.7, 4.8]).to(device)\n    class_weights = torch.tensor([1.2, 2.0, 2.4, 4.8, 4.9]).to(device)\n    criterion = nn.CrossEntropyLoss(weight=class_weights)\n    \n    # Optimizer with weight decay\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=learning_rate,\n        weight_decay=0.01\n    )\n    \n    # Learning rate scheduler\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=learning_rate,\n        steps_per_epoch=28,\n        epochs=50,\n        pct_start=0.3,\n        div_factor=10,\n        final_div_factor=100\n    )\n    \n    return model, criterion, optimizer, scheduler, predictor, wsi_config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:30.191946Z","iopub.execute_input":"2025-02-15T04:46:30.192303Z","iopub.status.idle":"2025-02-15T04:46:30.198089Z","shell.execute_reply.started":"2025-02-15T04:46:30.192271Z","shell.execute_reply":"2025-02-15T04:46:30.197241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main function to test model architecture\"\"\"\n    try:\n        print(\"Testing Model Architecture...\")\n        \n        # Set device\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"Using device: {device}\")\n        \n        # Create model and get all components\n        model, criterion, optimizer, scheduler, predictor, wsi_config = create_model(\n            device, learning_rate=5e-4)\n        \n        # Print model parameters\n        total_params = sum(p.numel() for p in model.parameters())\n        trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        \n        print(\"\\nModel Parameters:\")\n        print(f\"Total: {total_params:,}\")\n        print(f\"Trainable: {trainable_params:,}\")\n        \n        # Test forward pass\n        batch_size = 4\n        x = torch.randn(batch_size, 3, 224, 224).to(device)\n        \n        print(f\"\\nInput shape: {x.shape}\")\n        \n        # Set model to eval mode for testing\n        model.eval()\n        with torch.no_grad():\n            outputs = model(x)\n            if isinstance(outputs, tuple):\n                main_logits = outputs[0]  # Get main classifier output\n                probs = F.softmax(main_logits, dim=1)\n                print(f\"Output shape: {main_logits.shape}\")\n                \n                # Print mean probabilities for each class\n                mean_probs = probs.mean(dim=0)\n                print(\"\\nMean class probabilities:\")\n                for i, p in enumerate(mean_probs):\n                    print(f\"Class {i}: {p:.4f}\")\n            else:\n                print(f\"Output shape: {outputs.shape}\")\n        \n        print(\"\\nChecking PatchPredictor configuration:\")\n        print(f\"Batch size: {predictor.batch_size}\")\n        print(f\"WSI config input shape: {wsi_config.patch_input_shape}\")\n        \n        print(\"\\nModel architecture test completed successfully!\")\n        \n        # Clean up\n        del model, predictor\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n    except Exception as e:\n        print(f\"Error in model test: {str(e)}\")\n        traceback.print_exc()\n        \n        # Clean up even if there's an error\n        try:\n            del model, predictor\n            gc.collect()\n            torch.cuda.empty_cache()\n        except:\n            pass\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:35.475436Z","iopub.execute_input":"2025-02-15T04:46:35.475756Z","iopub.status.idle":"2025-02-15T04:46:39.620574Z","shell.execute_reply.started":"2025-02-15T04:46:35.47573Z","shell.execute_reply":"2025-02-15T04:46:39.619685Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Pipeline","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom sklearn.metrics import classification_report, balanced_accuracy_score\nimport numpy as np\nimport pandas as pd\nfrom typing import Dict, List, Tuple\nimport gc\nimport os\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:45.754812Z","iopub.execute_input":"2025-02-15T04:46:45.75512Z","iopub.status.idle":"2025-02-15T04:46:45.759812Z","shell.execute_reply.started":"2025-02-15T04:46:45.755096Z","shell.execute_reply":"2025-02-15T04:46:45.758946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MetricTracker:\n    \"\"\"Track training metrics\"\"\"\n    def __init__(self):\n        self.reset()\n    \n    def reset(self):\n        self.metrics = {\n            'loss': [],\n            'acc': [],\n            'balanced_acc': [],\n            'val_loss': [],\n            'val_acc': [],\n            'val_balanced_acc': [],\n            'best_val_acc': 0.0,\n            'best_val_balanced_acc': 0.0\n        }\n    \n    def update(self, phase: str, loss: float, acc: float, balanced_acc: float):\n        if phase == 'train':\n            self.metrics['loss'].append(loss)\n            self.metrics['acc'].append(acc)\n            self.metrics['balanced_acc'].append(balanced_acc)\n        else:\n            self.metrics['val_loss'].append(loss)\n            self.metrics['val_acc'].append(acc)\n            self.metrics['val_balanced_acc'].append(balanced_acc)\n            if acc > self.metrics['best_val_acc']:\n                self.metrics['best_val_acc'] = acc\n            if balanced_acc > self.metrics['best_val_balanced_acc']:\n                self.metrics['best_val_balanced_acc'] = balanced_acc\n    \n    def get_best_metrics(self) -> Dict:\n        return {\n            'best_val_acc': self.metrics['best_val_acc'],\n            'best_val_balanced_acc': self.metrics['best_val_balanced_acc'],\n            'best_epoch_val_loss': min(self.metrics['val_loss']) if self.metrics['val_loss'] else float('inf'),\n            'best_epoch_val_acc': max(self.metrics['val_acc']) if self.metrics['val_acc'] else 0,\n            'best_epoch_val_balanced_acc': max(self.metrics['val_balanced_acc']) if self.metrics['val_balanced_acc'] else 0\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:48.136683Z","iopub.execute_input":"2025-02-15T04:46:48.136989Z","iopub.status.idle":"2025-02-15T04:46:48.144372Z","shell.execute_reply.started":"2025-02-15T04:46:48.136965Z","shell.execute_reply":"2025-02-15T04:46:48.143468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, train_loader, val_loader, device, criterion,\n                 optimizer, scheduler, predictor, wsi_config, epochs=50,\n                 save_dir='./model_checkpoints'):\n        self.model = model\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.device = device\n        self.criterion = criterion\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.predictor = predictor\n        self.wsi_config = wsi_config\n        self.epochs = epochs\n        self.save_dir = save_dir\n        self.tracker = MetricTracker()\n        \n        os.makedirs(save_dir, exist_ok=True)\n    \n    def compute_loss(self, main_logits, aux_logits, labels):\n        \"\"\"Compute combined loss from main and auxiliary outputs\"\"\"\n        # Main classification loss\n        main_loss = self.criterion(main_logits, labels)\n        \n        # Auxiliary classification loss\n        aux_loss = self.criterion(aux_logits, labels)\n        \n        # Combined loss with weighting\n        total_loss = main_loss + 0.3 * aux_loss\n        \n        return total_loss, main_loss, aux_loss\n\n    def train_epoch(self):\n        self.model.train()\n        running_loss = 0.0\n        all_preds = []\n        all_labels = []\n        correct = 0\n        total = 0\n        \n        pbar = tqdm(self.train_loader, desc='Training')\n        for inputs, labels in pbar:\n            inputs = inputs.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n            \n            self.optimizer.zero_grad(set_to_none=True)\n            \n            # Get both main and auxiliary outputs\n            main_logits, aux_logits = self.model(inputs)\n            \n            # Compute combined loss\n            loss, main_loss, aux_loss = self.compute_loss(main_logits, aux_logits, labels)\n            \n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)\n            self.optimizer.step()\n            \n            if self.scheduler is not None:\n                self.scheduler.step()\n            \n            running_loss += loss.item()\n            _, predicted = main_logits.max(1)  # Use main classifier predictions\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            \n            pbar.set_postfix({\n                'loss': f'{loss.item():.4f}',\n                'acc': f'{100.*correct/total:.2f}%'\n            })\n        \n        balanced_acc = 100. * balanced_accuracy_score(all_labels, all_preds)\n        return running_loss / len(self.train_loader), 100. * correct / total, balanced_acc\n\n    @torch.no_grad()\n    def validate(self):\n        self.model.eval()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n        all_preds = []\n        all_labels = []\n        \n        pbar = tqdm(self.val_loader, desc='Validating')\n        for inputs, labels in pbar:\n            inputs = inputs.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n            \n            with torch.cuda.amp.autocast():\n                # During validation, model only returns main classifier output\n                outputs = self.model(inputs)\n                loss = self.criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            \n            pbar.set_postfix({\n                'loss': f'{loss.item():.4f}',\n                'acc': f'{100.*correct/total:.2f}%'\n            })\n        \n        balanced_acc = 100. * balanced_accuracy_score(all_labels, all_preds)\n        \n        class_names = ['HGSC', 'EC', 'CC', 'LGSC', 'MC']\n        report = classification_report(\n            all_labels,\n            all_preds,\n            target_names=class_names,\n            digits=3,\n            output_dict=True\n        )\n        \n        return running_loss / len(self.val_loader), 100. * correct / total, balanced_acc, report\n\n    def train(self) -> Dict:\n        print(f\"\\nStarting training for {self.epochs} epochs...\")\n        best_val_balanced_acc = 0.0\n    \n        for epoch in range(self.epochs):\n            print(f'\\nEpoch {epoch+1}/{self.epochs}')\n            print('-' * 20)\n        \n            # Get training metrics (now includes main and auxiliary losses)\n            train_loss, train_acc, train_balanced_acc = self.train_epoch()\n            self.tracker.update('train', train_loss, train_acc, train_balanced_acc)\n        \n            # Validation phase remains the same\n            val_loss, val_acc, val_balanced_acc, report = self.validate()\n            self.tracker.update('val', val_loss, val_acc, val_balanced_acc)\n        \n            # Print detailed training metrics\n            print(f'\\nTraining Results:')\n            print(f'Total Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%, Balanced Acc: {train_balanced_acc:.2f}%')\n        \n            # Print validation metrics\n            print(f'\\nValidation Results:')\n            print(f'Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%, Balanced Acc: {val_balanced_acc:.2f}%')\n        \n            # Print class-wise performance\n            print('\\nClass-wise Performance:')\n            for cls_name in ['HGSC', 'EC', 'CC', 'LGSC', 'MC']:\n                cls_metrics = report[cls_name]\n                print(f'{cls_name} - Precision: {cls_metrics[\"precision\"]:.3f}, '\n                      f'Recall: {cls_metrics[\"recall\"]:.3f}, '\n                      f'F1: {cls_metrics[\"f1-score\"]:.3f}')\n            if val_balanced_acc > best_val_balanced_acc:\n                best_val_balanced_acc = val_balanced_acc\n                model_path = os.path.join(self.save_dir, 'best_model.pth')\n            \n                # Save model with additional metrics\n                torch.save({\n                    'epoch': epoch + 1,\n                    'model_state_dict': self.model.state_dict(),\n                    'optimizer_state_dict': self.optimizer.state_dict(),\n                    'val_acc': val_acc,\n                    'val_balanced_acc': val_balanced_acc,\n                    'val_loss': val_loss,\n                    'train_acc': train_acc,\n                    'train_balanced_acc': train_balanced_acc,\n                    'train_loss': train_loss,\n                 }, model_path)\n                print(f'Saved new best model with validation balanced accuracy: {val_balanced_acc:.2f}%')\n            \n            # Memory cleanup after each epoch\n            torch.cuda.empty_cache()\n            gc.collect()\n\n        # Return final metrics\n        best_metrics = self.tracker.get_best_metrics()\n    \n        print(\"\\nTraining completed!\")\n        print(\"Best metrics achieved:\")\n        print(f\"Best validation accuracy: {best_metrics['best_val_acc']:.2f}%\")\n        print(f\"Best validation balanced accuracy: {best_metrics['best_val_balanced_acc']:.2f}%\")\n        print(f\"Best epoch validation loss: {best_metrics['best_epoch_val_loss']:.4f}\")\n    \n        return best_metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:46:51.937711Z","iopub.execute_input":"2025-02-15T04:46:51.93804Z","iopub.status.idle":"2025-02-15T04:46:51.953999Z","shell.execute_reply.started":"2025-02-15T04:46:51.938015Z","shell.execute_reply":"2025-02-15T04:46:51.953066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_main():\n    try:\n        print(\"Initializing Training Pipeline...\")\n        \n        # Set device\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"Using device: {device}\")\n        \n        # Use global data loaders\n        global train_loader, val_loader\n        \n        print(\"\\nInitializing model...\")\n        model, criterion, optimizer, scheduler, predictor, wsi_config = create_model(\n            device, \n            learning_rate=5e-4,\n        )\n        \n        # Print model summary\n        total_params = sum(p.numel() for p in model.parameters())\n        trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        print(f\"Total parameters: {total_params:,}\")\n        print(f\"Trainable parameters: {trainable_params:,}\")\n        \n        # Initialize trainer with new components\n        trainer = Trainer(\n            model=model,\n            train_loader=train_loader,\n            val_loader=val_loader,\n            device=device,\n            criterion=criterion,\n            optimizer=optimizer,\n            scheduler=scheduler,\n            predictor=predictor,\n            wsi_config=wsi_config,\n            epochs=50,  \n            save_dir='./model_checkpoints'\n        )\n        \n        # Start training\n        print(\"\\nStarting training process...\")\n        print(f\"Training on {len(train_loader.dataset)} samples\")\n        print(f\"Validating on {len(val_loader.dataset)} samples\")\n        \n        best_metrics = trainer.train()\n        \n        # Print final results\n        print(\"\\nTraining completed!\")\n        print(\"Best metrics achieved:\")\n        print(f\"Best validation accuracy: {best_metrics['best_val_acc']:.2f}%\")\n        print(f\"Best validation balanced accuracy: {best_metrics['best_val_balanced_acc']:.2f}%\")\n        print(f\"Best epoch validation loss: {best_metrics['best_epoch_val_loss']:.4f}\")\n        \n        # Additional metrics reporting\n        print(\"\\nDetailed metrics:\")\n        print(f\"Best epoch validation accuracy: {best_metrics['best_epoch_val_acc']:.2f}%\")\n        print(f\"Best epoch validation balanced accuracy: {best_metrics['best_epoch_val_balanced_acc']:.2f}%\")\n        \n        # Clean up\n        del model, trainer\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        return best_metrics\n        \n    except Exception as e:\n        print(f\"Error in training pipeline: {str(e)}\")\n        traceback.print_exc()\n        \n        # Clean up even if there's an error\n        try:\n            del model, trainer\n            gc.collect()\n            torch.cuda.empty_cache()\n        except:\n            pass\n        return None\n\nif __name__ == \"__main__\":\n    train_main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T04:47:00.814267Z","iopub.execute_input":"2025-02-15T04:47:00.814596Z","iopub.status.idle":"2025-02-15T05:07:33.798846Z","shell.execute_reply.started":"2025-02-15T04:47:00.814571Z","shell.execute_reply":"2025-02-15T05:07:33.79786Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Image Prediction and Analysis","metadata":{}},{"cell_type":"code","source":"!pip install -q umap-learn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:08:49.79271Z","iopub.execute_input":"2025-02-15T05:08:49.793023Z","iopub.status.idle":"2025-02-15T05:08:53.212484Z","shell.execute_reply.started":"2025-02-15T05:08:49.792999Z","shell.execute_reply":"2025-02-15T05:08:53.2115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import umap\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom sklearn.manifold import TSNE\nfrom sklearn.metrics import balanced_accuracy_score, confusion_matrix, classification_report\nimport seaborn as sns\nimport torch\nimport gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:08:55.439781Z","iopub.execute_input":"2025-02-15T05:08:55.440119Z","iopub.status.idle":"2025-02-15T05:08:55.445012Z","shell.execute_reply.started":"2025-02-15T05:08:55.440091Z","shell.execute_reply":"2025-02-15T05:08:55.443975Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_features(model, dataloader, device):\n    \"\"\"Extract features from the model's intermediate layer\"\"\"\n    features = []\n    positions = []\n    labels = []\n    \n    model.eval()\n    with torch.no_grad():\n        for i, (images, batch_labels) in enumerate(dataloader):\n            images = images.to(device)\n            # Use model's extract_features method directly\n            batch_features = model.extract_features(images)\n            \n            features.append(batch_features.cpu().numpy())\n            labels.append(batch_labels.numpy())\n            positions.append(np.array([(i * images.shape[0] + j, j) for j in range(images.shape[0])]))\n    \n    features = np.concatenate(features)\n    labels = np.concatenate(labels)\n    positions = np.concatenate(positions)\n    return features, labels, positions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:08:57.952235Z","iopub.execute_input":"2025-02-15T05:08:57.95255Z","iopub.status.idle":"2025-02-15T05:08:57.95884Z","shell.execute_reply.started":"2025-02-15T05:08:57.952527Z","shell.execute_reply":"2025-02-15T05:08:57.957958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_sample_predictions(model, val_loader, device, num_samples=10):\n    \"\"\"Display sample images with their predictions\"\"\"\n    model.eval()\n    classes = ['HGSC', 'EC', 'CC', 'LGSC', 'MC'] \n    \n    # Get a batch of images\n    images, labels = next(iter(val_loader))\n    \n    # Get predictions\n    with torch.no_grad():\n        outputs = model(images.to(device))\n        if isinstance(outputs, tuple):  # Handle training mode output\n            outputs = outputs[0]\n        _, preds = torch.max(outputs, 1)\n    \n    # Create a figure to display images\n    fig = plt.figure(figsize=(20, 4))\n    for idx in range(min(num_samples, len(images))):\n        ax = fig.add_subplot(1, num_samples, idx + 1, xticks=[], yticks=[])\n        \n        # Convert tensor to image\n        img = images[idx].numpy().transpose((1, 2, 0))\n        img = np.clip(img, 0, 1)\n        \n        # Display image\n        ax.imshow(img)\n        \n        # Add title with true and predicted labels\n        true_label = classes[labels[idx]]\n        pred_label = classes[preds[idx].cpu()]\n        color = 'green' if true_label == pred_label else 'red'\n        ax.set_title(f'True: {true_label}\\nPred: {pred_label}', color=color)\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:09:00.777313Z","iopub.execute_input":"2025-02-15T05:09:00.777656Z","iopub.status.idle":"2025-02-15T05:09:00.784614Z","shell.execute_reply.started":"2025-02-15T05:09:00.777625Z","shell.execute_reply":"2025-02-15T05:09:00.783769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_feature_distribution(features, labels):\n    \"\"\"Plot the distribution of features across classes\"\"\"\n    # Reduce dimensionality to 2D using UMAP\n    reducer = umap.UMAP(n_neighbors=15, min_dist=0.1, metric='euclidean')\n    embedding = reducer.fit_transform(features)\n    \n    # Create scatter plot with different colors for each class\n    plt.figure(figsize=(12, 8))\n    classes = ['HGSC', 'EC', 'CC', 'LGSC', 'MC']\n    colors = ['blue', 'red', 'green', 'purple', 'orange']\n    \n    for i, cls in enumerate(classes):\n        mask = labels == i\n        plt.scatter(embedding[mask, 0], embedding[mask, 1],\n                   c=colors[i], label=cls, alpha=0.6)\n    \n    plt.title('Feature Distribution Across Classes')\n    plt.legend()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:09:04.421304Z","iopub.execute_input":"2025-02-15T05:09:04.42163Z","iopub.status.idle":"2025-02-15T05:09:04.42711Z","shell.execute_reply.started":"2025-02-15T05:09:04.421606Z","shell.execute_reply":"2025-02-15T05:09:04.42626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_feature_space(features, labels, title=\"Feature Space Visualization\"):\n    \"\"\"Create UMAP visualization of feature space\"\"\"\n    # Reduce dimensionality to 2D\n    reducer = umap.UMAP(n_neighbors=15, min_dist=0.1, metric='euclidean')\n    embedding = reducer.fit_transform(features)\n    \n    # Create plot\n    plt.figure(figsize=(12, 8))\n    scatter = plt.scatter(embedding[:, 0], embedding[:, 1], c=labels, cmap='Spectral')\n    plt.colorbar(scatter, label='True Class')\n    plt.title(title)\n    plt.xlabel(\"UMAP Dimension 1\")\n    plt.ylabel(\"UMAP Dimension 2\")\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:09:07.731253Z","iopub.execute_input":"2025-02-15T05:09:07.731631Z","iopub.status.idle":"2025-02-15T05:09:07.736907Z","shell.execute_reply.started":"2025-02-15T05:09:07.7316Z","shell.execute_reply":"2025-02-15T05:09:07.735884Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(model, dataloader, device):\n    \"\"\"Visualize model predictions vs actual labels with enhanced metrics\"\"\"\n    predictions = []\n    actuals = []\n    probabilities = []\n    \n    model.eval()\n    with torch.no_grad():\n        for images, labels in dataloader:\n            images = images.to(device)\n            outputs = model(images)\n            if isinstance(outputs, tuple):  # Handle training mode output\n                outputs = outputs[0]\n            probs = torch.softmax(outputs, dim=1)\n            _, preds = torch.max(outputs, 1)\n            predictions.extend(preds.cpu().numpy())\n            actuals.extend(labels.numpy())\n            probabilities.extend(probs.cpu().numpy())\n    \n    predictions = np.array(predictions)\n    actuals = np.array(actuals)\n    probabilities = np.array(probabilities)\n    \n    # Calculate balanced accuracy\n    balanced_acc = balanced_accuracy_score(actuals, predictions) * 100\n    \n    # Create confusion matrix plot\n    cm = confusion_matrix(actuals, predictions)\n    plt.figure(figsize=(12, 8))\n    classes = ['HGSC', 'EC', 'CC', 'LGSC', 'MC']\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=classes, yticklabels=classes)\n    plt.title(f'Confusion Matrix\\nBalanced Accuracy: {balanced_acc:.2f}%')\n    plt.ylabel('True Label')\n    plt.xlabel('Predicted Label')\n    plt.show()\n    \n    # Print detailed classification report\n    print(\"\\nClassification Report:\")\n    print(classification_report(actuals, predictions, target_names=classes))\n    \n    # Plot per-class prediction confidence\n    plt.figure(figsize=(12, 6))\n    \n    for i, cls in enumerate(classes):\n        true_mask = actuals == i\n        pred_mask = predictions == i\n        \n        # Correct predictions\n        correct_mask = np.logical_and(true_mask, pred_mask)\n        if np.any(correct_mask):\n            plt.scatter(np.full(np.sum(correct_mask), i+0.1),\n                       probabilities[correct_mask, i],\n                       c='green', alpha=0.5, label='Correct' if i == 0 else '')\n        \n        # Wrong predictions\n        wrong_mask = np.logical_and(true_mask, ~pred_mask)\n        if np.any(wrong_mask):\n            plt.scatter(np.full(np.sum(wrong_mask), i-0.1),\n                       probabilities[wrong_mask, i],\n                       c='red', alpha=0.5, label='Wrong' if i == 0 else '')\n    \n    plt.xticks(range(len(classes)), classes, rotation=45)\n    plt.ylabel('Prediction Confidence')\n    plt.title('Per-class Prediction Confidence Distribution')\n    plt.legend()\n    plt.tight_layout()\n    plt.show()\n    \n    return balanced_acc, cm, predictions, actuals, probabilities","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:09:10.042615Z","iopub.execute_input":"2025-02-15T05:09:10.042952Z","iopub.status.idle":"2025-02-15T05:09:10.053015Z","shell.execute_reply.started":"2025-02-15T05:09:10.042922Z","shell.execute_reply":"2025-02-15T05:09:10.052064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_model(model, train_loader, val_loader, device):\n    \"\"\"Run comprehensive model analysis with enhanced visualizations\"\"\"\n    print(\"1. Displaying sample predictions...\")\n    visualize_sample_predictions(model, val_loader, device)\n    \n    print(\"\\n2. Extracting features...\")\n    train_features, train_labels, _ = extract_features(model, train_loader, device)\n    val_features, val_labels, _ = extract_features(model, val_loader, device)\n    \n    print(\"\\n3. Plotting feature distributions...\")\n    plot_feature_distribution(train_features, train_labels)\n    \n    print(\"\\n4. Analyzing predictions and metrics...\")\n    balanced_acc, cm, predictions, actuals, probs = visualize_predictions(model, val_loader, device)\n    \n    print(f\"\\nOverall Balanced Accuracy: {balanced_acc:.2f}%\")\n    \n    print(\"\\n5. Plotting UMAP embeddings...\")\n    plot_feature_space(train_features, train_labels, \"Training Data Feature Space\")\n    plot_feature_space(val_features, val_labels, \"Validation Data Feature Space\")\n    \n    # Additional per-class analysis\n    print(\"\\nPer-class Performance Summary:\")\n    for i, cls in enumerate(['HGSC', 'EC', 'CC', 'LGSC', 'MC']):\n        class_mask = actuals == i\n        class_acc = balanced_accuracy_score([1 if x == i else 0 for x in actuals],\n                                          [1 if x == i else 0 for x in predictions]) * 100\n        class_conf = probs[class_mask, i].mean() * 100\n        print(f\"{cls}:\")\n        print(f\" Balanced Accuracy: {class_acc:.2f}%\")\n        print(f\" Average Confidence: {class_conf:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:09:15.065308Z","iopub.execute_input":"2025-02-15T05:09:15.065645Z","iopub.status.idle":"2025-02-15T05:09:15.072249Z","shell.execute_reply.started":"2025-02-15T05:09:15.065618Z","shell.execute_reply":"2025-02-15T05:09:15.071232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_analysis():\n    \"\"\"Run the complete analysis pipeline\"\"\"\n    try:\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"Using device: {device}\")\n        \n        # Load your best model\n        model = HistoPathModel()\n        checkpoint = torch.load('./model_checkpoints/best_model.pth')\n        model.load_state_dict(checkpoint['model_state_dict'])\n        model.to(device)\n        \n        print(\"\\nStarting model analysis...\")\n        print(f\"Model checkpoint metrics:\")\n        print(f\"Validation Accuracy: {checkpoint['val_acc']:.2f}%\")\n        print(f\"Validation Balanced Accuracy: {checkpoint['val_balanced_acc']:.2f}%\")\n        print(f\"Validation Loss: {checkpoint['val_loss']:.4f}\")\n        \n        analyze_model(model, train_loader, val_loader, device)\n        \n        # Clean up\n        del model\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n    except Exception as e:\n        print(f\"Error in analysis: {str(e)}\")\n        raise\n\nif __name__ == \"__main__\":\n    run_analysis()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:09:18.491729Z","iopub.execute_input":"2025-02-15T05:09:18.492034Z","iopub.status.idle":"2025-02-15T05:09:42.416502Z","shell.execute_reply.started":"2025-02-15T05:09:18.492009Z","shell.execute_reply":"2025-02-15T05:09:42.415761Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Outlier Detection","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torch.nn import functional as F\nfrom torchvision.models import efficientnet_b0, EfficientNet_B0_Weights\nfrom torchvision.models import efficientnet_b3, EfficientNet_B3_Weights\nfrom torchvision.models import resnet101, ResNet101_Weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:10:22.688107Z","iopub.execute_input":"2025-02-15T05:10:22.68843Z","iopub.status.idle":"2025-02-15T05:10:22.692875Z","shell.execute_reply.started":"2025-02-15T05:10:22.688404Z","shell.execute_reply":"2025-02-15T05:10:22.692045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class OutlierHistoPathModel(nn.Module):\n    def __init__(self, num_classes=5, feature_dim=512):\n        super().__init__()\n        # Base backbones (same as before)\n        self.resnet = resnet101(weights=ResNet101_Weights.DEFAULT)\n        self.efficientnet = efficientnet_b3(weights=EfficientNet_B3_Weights.DEFAULT)\n        \n        # Remove original classifier layers\n        self.resnet.fc = nn.Identity()\n        self.efficientnet.classifier = nn.Identity()\n        \n        # Feature dimensions\n        self.resnet_dim = 2048\n        self.efficient_dim = 1536\n        self.feature_dim = feature_dim\n        self.combined_resnet_dim = self.resnet_dim + (512 * 3)\n        \n        # Feature reduction layers (same as before)\n        self.resnet_reducer = nn.Sequential(\n            nn.Linear(self.combined_resnet_dim, self.feature_dim),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        \n        self.efficient_reducer = nn.Sequential(\n            nn.Linear(self.efficient_dim, self.feature_dim * 2),\n            nn.BatchNorm1d(self.feature_dim * 2),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(self.feature_dim * 2, self.feature_dim)\n         )\n        \n        # Modify first conv layer\n        self.resnet.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        \n        # Attention modules (same as before)\n        self.spatial_attention = nn.Sequential(\n            nn.Conv2d(2048, 512, kernel_size=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n        self.channel_attention = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(2048, 512, kernel_size=1),\n            nn.ReLU(),\n            nn.Conv2d(512, 2048, kernel_size=1),\n            nn.Sigmoid()\n        )\n        # Add channel attention for EfficientNet\n        self.efficient_channel_attention = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(1536, 512, kernel_size=1),  # 1536 for EfficientNet-B3\n            nn.ReLU(),\n            nn.Conv2d(512, 1536, kernel_size=1),\n            nn.Sigmoid()\n        )\n\n        # Add fusion attention\n        self.fusion_attention = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim // 16),\n            nn.ReLU(),\n            nn.Linear(self.feature_dim // 16, self.feature_dim),\n            nn.Sigmoid()\n        )\n\n        \n        # Feature processors (same as before)\n        self.feature_processors = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(2048, 512, 1),\n                nn.BatchNorm2d(512),\n                nn.ReLU()\n            ),\n            nn.Sequential(\n                nn.Conv2d(2048, 512, 3, padding=1),\n                nn.BatchNorm2d(512),\n                nn.ReLU()\n            ),\n            nn.Sequential(\n                nn.Conv2d(2048, 512, 5, padding=2),\n                nn.BatchNorm2d(512),\n                nn.ReLU()\n            )\n        ])\n        \n        self.feature_fusion = nn.Sequential(\n            nn.Linear(self.feature_dim * 2, self.feature_dim * 2),\n            nn.LayerNorm(self.feature_dim * 2),\n            nn.ReLU(),\n            nn.Dropout(0.2),  # Reduced dropout for stability\n            nn.Linear(self.feature_dim * 2, self.feature_dim * 2),\n            nn.LayerNorm(self.feature_dim * 2),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(self.feature_dim * 2, self.feature_dim)\n        )\n    \n  \n        # Add feature gates\n        self.feature_gates = nn.Sequential(\n            nn.Linear(self.feature_dim * 2, 2),\n            nn.Softmax(dim=1)\n        )\n\n    \n        # Add a squeeze-excitation block\n        self.se_block = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim // 16),\n            nn.ReLU(),\n            nn.Linear(self.feature_dim // 16, self.feature_dim),\n            nn.Sigmoid()\n         )\n    \n        \n        # Self-attention\n        self.self_attention = nn.MultiheadAttention(\n            embed_dim=self.feature_dim,\n            num_heads=8,\n            dropout=0.1,\n            batch_first=True\n        )\n        \n        # Classification head\n        self.classifier = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim),\n            nn.LayerNorm(self.feature_dim),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(self.feature_dim, num_classes)\n        )\n        \n        # Outlier detection head\n        self.outlier_detector = nn.Sequential(\n            nn.Linear(self.feature_dim, self.feature_dim),\n            nn.LayerNorm(self.feature_dim),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(self.feature_dim, self.feature_dim * 2)  # Mean and log variance\n        )\n\n    def extract_resnet_features(self, x):\n        # Initial layers\n        x = self.resnet.conv1(x)\n        x = self.resnet.bn1(x)\n        x = self.resnet.relu(x)\n        x = self.resnet.maxpool(x)\n        \n        # ResNet blocks\n        x = self.resnet.layer1(x)\n        x = self.resnet.layer2(x)\n        x = self.resnet.layer3(x)\n        x = self.resnet.layer4(x)\n        \n        # Apply attention\n        spatial_weights = self.spatial_attention(x)\n        channel_weights = self.channel_attention(x)\n        attended_features = x * spatial_weights * channel_weights\n        \n        # Global features\n        global_features = F.adaptive_avg_pool2d(attended_features, 1).flatten(1)\n        \n        # Process features through each processor\n        processed_features = []\n        for processor in self.feature_processors:\n            features = processor(attended_features)\n            pooled = F.adaptive_avg_pool2d(features, 1).flatten(1)\n            processed_features.append(pooled)\n        \n        # Concatenate global and processed features\n        combined_features = torch.cat([global_features] + processed_features, dim=1)\n        \n        return self.resnet_reducer(combined_features)\n\n    def extract_efficient_features(self, x):\n        features = self.efficientnet.features(x)\n    \n        # Apply channel attention\n        channel_weights = self.efficient_channel_attention(features)\n        features = features * channel_weights\n    \n        features = self.efficientnet.avgpool(features)\n        features = torch.flatten(features, 1)\n        return self.efficient_reducer(features)\n    \n\n    def extract_features(self, x):\n        resnet_features = self.extract_resnet_features(x)\n        torch.cuda.empty_cache()\n    \n        efficient_features = self.extract_efficient_features(x)\n        torch.cuda.empty_cache()\n    \n        combined_features = torch.cat([resnet_features, efficient_features], dim=1)\n        fused_features = self.feature_fusion(combined_features)\n    \n        # Residual connection\n        if hasattr(self, 'feature_gates'):\n            gates = self.feature_gates(combined_features)\n            residual = resnet_features * gates[:, 0].unsqueeze(1) + efficient_features * gates[:, 1].unsqueeze(1)\n            fused_features = fused_features + residual\n    \n        \n        # Apply self-attention with gradient clipping\n        attended_features, _ = self.self_attention(\n            fused_features.unsqueeze(1),\n            fused_features.unsqueeze(1),\n            fused_features.unsqueeze(1)\n        )\n        \n        return fused_features + 0.1 * attended_features.squeeze(1)\n\n    def compute_outlier_score(self, x):\n        \"\"\"Compute outlier score for input images\"\"\"\n        with torch.no_grad():\n            # Extract features\n            features = self.extract_features(x)\n            \n            # Get distribution parameters\n            dist_params = self.outlier_detector(features)\n            mean, log_var = torch.chunk(dist_params, 2, dim=1)\n            \n            # Compute Mahalanobis distance as outlier score\n            var = torch.exp(log_var)\n            z_score = (features - mean) / torch.sqrt(var + 1e-6)\n            outlier_score = torch.sum(z_score ** 2, dim=1)\n            \n            return outlier_score\n\n    def forward(self, x):\n        \"\"\"Forward pass with both classification and outlier detection\"\"\"\n        # Extract features\n        features = self.extract_features(x)\n        \n        # Classification logits\n        logits = self.classifier(features)\n        \n        # Outlier detection parameters\n        dist_params = self.outlier_detector(features)\n        mean, log_var = torch.chunk(dist_params, 2, dim=1)\n        \n        if self.training:\n            return logits, mean, log_var\n        else:\n            return logits\n\n    def __del__(self):\n        torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:17:25.036127Z","iopub.execute_input":"2025-02-15T05:17:25.036548Z","iopub.status.idle":"2025-02-15T05:17:25.056683Z","shell.execute_reply.started":"2025-02-15T05:17:25.036521Z","shell.execute_reply":"2025-02-15T05:17:25.055922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class OutlierLoss(nn.Module):\n    \"\"\"Combined loss for classification and outlier detection\"\"\"\n    def __init__(self, num_classes=5, outlier_weight=0.1):\n        super().__init__()\n        self.ce_loss = nn.CrossEntropyLoss()\n        self.outlier_weight = outlier_weight\n        \n    def forward(self, logits, mean, log_var, labels):\n        # Classification loss\n        ce_loss = self.ce_loss(logits, labels)\n        \n        # Feature distribution regularization\n        kl_loss = -0.5 * torch.mean(1 + log_var - mean.pow(2) - log_var.exp())\n        \n        # Combined loss\n        total_loss = ce_loss + self.outlier_weight * kl_loss\n        \n        return total_loss, ce_loss, kl_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:17:28.293469Z","iopub.execute_input":"2025-02-15T05:17:28.293787Z","iopub.status.idle":"2025-02-15T05:17:28.298937Z","shell.execute_reply.started":"2025-02-15T05:17:28.293761Z","shell.execute_reply":"2025-02-15T05:17:28.298028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_outlier_model(device, learning_rate=5e-4, epochs=30):\n    \"\"\"Create model with outlier detection capabilities\"\"\"\n    model = OutlierHistoPathModel(num_classes=5)\n    model = model.to(device)\n    \n    # Create predictor configuration\n    wsi_config = IOPatchPredictorConfig(\n        input_resolutions=[{\"units\": \"mpp\", \"resolution\": 0.5}],\n        patch_input_shape=[224, 224],\n        stride_shape=[224, 224]\n    )\n    \n    # Create patch predictor\n    predictor = PatchPredictor(\n        model=model,\n        batch_size=32,\n        num_loader_workers=4\n    )\n    \n    # Combined loss function\n    criterion = OutlierLoss(num_classes=5)\n    \n    # Optimizer\n    optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)\n    \n    # Learning rate scheduler\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=learning_rate,\n        steps_per_epoch=28,\n        epochs=epochs,\n        pct_start=0.3\n    )\n    \n    return model, criterion, optimizer, scheduler, predictor, wsi_config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:18:00.824073Z","iopub.execute_input":"2025-02-15T05:18:00.824391Z","iopub.status.idle":"2025-02-15T05:18:00.829779Z","shell.execute_reply.started":"2025-02-15T05:18:00.824366Z","shell.execute_reply":"2025-02-15T05:18:00.828951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class OutlierTrainer:\n    def __init__(self, model, train_loader, val_loader, device, criterion,\n                 optimizer, scheduler, predictor, wsi_config, epochs=30,\n                 save_dir='./outlier_model_checkpoints', batch_size=16):\n        self.model = model\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.device = device\n        self.criterion = criterion\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.predictor = predictor\n        self.wsi_config = wsi_config\n        self.epochs = epochs\n        self.save_dir = save_dir\n        self.batch_size = batch_size\n        \n        # Create gradient scaler for mixed precision training\n        self.scaler = torch.cuda.amp.GradScaler()\n        \n        os.makedirs(save_dir, exist_ok=True)\n\n    def compute_auxiliary_loss(self, features, labels):\n        \"\"\"Auxiliary task: Feature clustering loss\"\"\"\n        # Process in chunks to save memory\n        chunk_size = 4\n        total_center_loss = 0\n        \n        for i in range(0, len(features), chunk_size):\n            chunk_features = features[i:i + chunk_size]\n            chunk_labels = labels[i:i + chunk_size]\n            \n            # Compute center for each class in chunk\n            centers = {}\n            for cls in torch.unique(chunk_labels):\n                centers[cls.item()] = chunk_features[chunk_labels == cls].mean(0)\n            \n            # Compute center loss for chunk\n            chunk_loss = 0\n            for cls in centers:\n                cls_features = chunk_features[chunk_labels == cls]\n                if len(cls_features) > 0:\n                    chunk_loss += F.mse_loss(cls_features, \n                                           centers[cls].expand(len(cls_features), -1))\n            \n            total_center_loss += chunk_loss\n            \n            # Clear cache\n            torch.cuda.empty_cache()\n        \n        return total_center_loss\n    \n    def compute_consistency_loss(self, image, augmented_image):\n        \"\"\"Consistency between different views of same image\"\"\"\n        # Process in chunks\n        chunk_size = 4\n        total_consist_loss = 0\n        \n        for i in range(0, len(image), chunk_size):\n            # Extract features for original and augmented chunks\n            with torch.cuda.amp.autocast():\n                orig_features = self.model.extract_features(image[i:i + chunk_size])\n                aug_features = self.model.extract_features(augmented_image[i:i + chunk_size])\n                chunk_loss = F.mse_loss(orig_features, aug_features)\n            \n            total_consist_loss += chunk_loss\n            \n            # Clear cache\n            torch.cuda.empty_cache()\n        \n        return total_consist_loss\n    \n    def augment_batch(self, images):\n        \"\"\"Memory efficient augmentation\"\"\"\n        augmented = []\n        chunk_size = 4\n        \n        for i in range(0, len(images), chunk_size):\n            chunk = images[i:i + chunk_size]\n            chunk_aug = []\n            \n            for img in chunk:\n                transform = A.Compose([\n                    A.RandomRotate90(p=0.5),\n                    A.HorizontalFlip(p=0.5),\n                    A.VerticalFlip(p=0.5),\n                    A.ColorJitter(brightness=0.2, contrast=0.2, p=0.5)\n                ])\n                \n                # Move to CPU for transformation\n                img_np = img.cpu().numpy().transpose(1, 2, 0)\n                aug_img = transform(image=img_np)['image']\n                chunk_aug.append(torch.from_numpy(aug_img.transpose(2, 0, 1)))\n            \n            # Move augmented chunk back to GPU\n            chunk_tensor = torch.stack(chunk_aug).to(self.device)\n            augmented.append(chunk_tensor)\n            \n            # Clear cache\n            torch.cuda.empty_cache()\n        \n        return torch.cat(augmented, dim=0)\n\n    def train_epoch(self):\n        self.model.train()\n        running_total_loss = 0.0\n        running_ce_loss = 0.0\n        running_kl_loss = 0.0\n        running_aux_loss = 0.0\n        running_consist_loss = 0.0\n        correct = 0\n        total = 0\n        \n        pbar = tqdm(self.train_loader, desc='Training')\n        for inputs, labels in pbar:\n            # Limit batch size\n            if len(inputs) > self.batch_size:\n                inputs = inputs[:self.batch_size]\n                labels = labels[:self.batch_size]\n            \n            inputs = inputs.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n            \n            # Clear cache before augmentation\n            torch.cuda.empty_cache()\n            \n            # Get augmented version\n            augmented_inputs = self.augment_batch(inputs)\n            \n            self.optimizer.zero_grad(set_to_none=True)\n            \n            # Use mixed precision training\n            with torch.cuda.amp.autocast():\n                # Get model outputs\n                logits, mean, log_var = self.model(inputs)\n                features = self.model.extract_features(inputs)\n                \n                # Calculate losses\n                total_loss, ce_loss, kl_loss = self.criterion(logits, mean, log_var, labels)\n                aux_loss = self.compute_auxiliary_loss(features, labels)\n                consist_loss = self.compute_consistency_loss(inputs, augmented_inputs)\n                \n                # Combined loss\n                total_loss = total_loss + 0.1 * aux_loss + 0.1 * consist_loss\n            \n            # Scaled backward pass\n            self.scaler.scale(total_loss).backward()\n            self.scaler.unscale_(self.optimizer)\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)\n            self.scaler.step(self.optimizer)\n            self.scaler.update()\n            \n            if self.scheduler is not None:\n                self.scheduler.step()\n            \n            # Update metrics\n            running_total_loss += total_loss.item()\n            running_ce_loss += ce_loss.item()\n            running_kl_loss += kl_loss.item()\n            running_aux_loss += aux_loss.item()\n            running_consist_loss += consist_loss.item()\n            \n            _, predicted = logits.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n            # Clear cache\n            torch.cuda.empty_cache()\n            \n            pbar.set_postfix({\n                'total_loss': f'{total_loss.item():.4f}',\n                'acc': f'{100.*correct/total:.2f}%'\n            })\n        \n        return {\n            'total_loss': running_total_loss / len(self.train_loader),\n            'ce_loss': running_ce_loss / len(self.train_loader),\n            'kl_loss': running_kl_loss / len(self.train_loader),\n            'aux_loss': running_aux_loss / len(self.train_loader),\n            'consist_loss': running_consist_loss / len(self.train_loader),\n            'accuracy': 100. * correct / total\n        }\n\n    @torch.no_grad()\n    def validate(self):\n        self.model.eval()\n        running_total_loss = 0.0\n        running_ce_loss = 0.0\n        running_kl_loss = 0.0\n        correct = 0\n        total = 0\n        \n        # For outlier detection metrics\n        all_scores = []\n        all_preds = []\n        all_labels = []\n        \n        pbar = tqdm(self.val_loader, desc='Validating')\n        for inputs, labels in pbar:\n            # Limit batch size\n            if len(inputs) > self.batch_size:\n                inputs = inputs[:self.batch_size]\n                labels = labels[:self.batch_size]\n            \n            inputs = inputs.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n            \n            # Use mixed precision for validation too\n            with torch.cuda.amp.autocast():\n                # Get model outputs (model returns tuple in training mode)\n                self.model.train()  # Temporarily set to train mode to get all outputs\n                logits, mean, log_var = self.model(inputs)\n                self.model.eval()  # Set back to eval mode\n                \n                # Calculate losses\n                total_loss, ce_loss, kl_loss = self.criterion(logits, mean, log_var, labels)\n                \n                # Calculate outlier scores\n                outlier_scores = self.model.compute_outlier_score(inputs)\n            \n            # Update metrics\n            running_total_loss += total_loss.item()\n            running_ce_loss += ce_loss.item()\n            running_kl_loss += kl_loss.item()\n            \n            _, predicted = logits.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n            # Store predictions and scores\n            all_scores.extend(outlier_scores.cpu().numpy())\n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            \n            # Clear cache\n            torch.cuda.empty_cache()\n            \n            pbar.set_postfix({\n                'total_loss': f'{total_loss.item():.4f}',\n                'acc': f'{100.*correct/total:.2f}%'\n            })\n        \n        # Calculate average metrics\n        avg_total_loss = running_total_loss / len(self.val_loader)\n        avg_ce_loss = running_ce_loss / len(self.val_loader)\n        avg_kl_loss = running_kl_loss / len(self.val_loader)\n        accuracy = 100. * correct / total\n        \n        outlier_metrics = {\n            'scores': np.array(all_scores),\n            'predictions': np.array(all_preds),\n            'labels': np.array(all_labels)\n        }\n        \n        return avg_total_loss, avg_ce_loss, avg_kl_loss, accuracy, outlier_metrics\n\n    def train(self):\n        print(f\"\\nStarting outlier detection training for {self.epochs} epochs...\")\n        best_metrics = {\n            'best_val_loss': float('inf'),\n            'best_val_acc': 0.0,\n            'best_epoch': 0\n        }\n        \n        for epoch in range(self.epochs):\n            print(f'\\nEpoch {epoch+1}/{self.epochs}')\n            print('-' * 20)\n            \n            # Training phase\n            train_metrics = self.train_epoch()\n            \n            # Clear cache before validation\n            torch.cuda.empty_cache()\n            \n            # Validation phase\n            val_total_loss, val_ce_loss, val_kl_loss, val_acc, outlier_metrics = self.validate()\n            \n            # Print epoch results\n            print(f'\\nTraining Results:')\n            print(f\"Total Loss: {train_metrics['total_loss']:.4f}\")\n            print(f\"CE Loss: {train_metrics['ce_loss']:.4f}\")\n            print(f\"KL Loss: {train_metrics['kl_loss']:.4f}\")\n            print(f\"Auxiliary Loss: {train_metrics['aux_loss']:.4f}\")\n            print(f\"Consistency Loss: {train_metrics['consist_loss']:.4f}\")\n            print(f\"Accuracy: {train_metrics['accuracy']:.2f}%\")\n            \n            print(f'\\nValidation Results:')\n            print(f'Total Loss: {val_total_loss:.4f}')\n            print(f'CE Loss: {val_ce_loss:.4f}')\n            print(f'KL Loss: {val_kl_loss:.4f}')\n            print(f'Accuracy: {val_acc:.2f}%')\n            \n            # Save best model\n            if val_total_loss < best_metrics['best_val_loss']:\n                best_metrics['best_val_loss'] = val_total_loss\n                best_metrics['best_val_acc'] = val_acc\n                best_metrics['best_epoch'] = epoch + 1\n                \n                model_path = os.path.join(self.save_dir, 'best_model.pth')\n                torch.save({\n                    'epoch': epoch + 1,\n                    'model_state_dict': self.model.state_dict(),\n                    'optimizer_state_dict': self.optimizer.state_dict(),\n                    'scaler_state_dict': self.scaler.state_dict(),\n                    'val_loss': val_total_loss,\n                    'val_acc': val_acc,\n                    'outlier_metrics': outlier_metrics\n                }, model_path)\n                print(f'Saved new best model with validation loss: {val_total_loss:.4f}')\n                print(f'Saved new best model with validation accuracy: {val_acc:.2f}')\n            \n            # Memory cleanup after each epoch\n            torch.cuda.empty_cache()\n            gc.collect()\n        \n        return best_metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:18:14.404941Z","iopub.execute_input":"2025-02-15T05:18:14.405277Z","iopub.status.idle":"2025-02-15T05:18:14.430036Z","shell.execute_reply.started":"2025-02-15T05:18:14.40524Z","shell.execute_reply":"2025-02-15T05:18:14.429342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_outlier_main():\n    try:\n        print(\"Initializing Outlier Detection Pipeline...\")\n        \n        # Set device\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"Using device: {device}\")\n        \n        # Use global data loaders\n        global train_loader, val_loader\n        \n        print(\"\\nInitializing model with outlier detection...\")\n        # Create model with outlier detection capabilities\n        model, criterion, optimizer, scheduler, predictor, wsi_config = create_outlier_model(\n            device, \n            learning_rate=5e-4,  \n            epochs=30\n        )\n        \n        # Print model summary\n        total_params = sum(p.numel() for p in model.parameters())\n        trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        print(f\"Total parameters: {total_params:,}\")\n        print(f\"Trainable parameters: {trainable_params:,}\")\n        \n        # Initialize trainer with memory optimizations\n        trainer = OutlierTrainer(\n            model=model,\n            train_loader=train_loader,\n            val_loader=val_loader,\n            device=device,\n            criterion=criterion,\n            optimizer=optimizer,\n            scheduler=scheduler,\n            predictor=predictor,\n            wsi_config=wsi_config,\n            epochs=30,\n            save_dir='./outlier_model_checkpoints',\n            batch_size=16  # Add reduced batch size for memory efficiency\n        )\n        \n        # Start training\n        print(\"\\nStarting training process...\")\n        print(f\"Training on {len(train_loader.dataset)} samples\")\n        print(f\"Validating on {len(val_loader.dataset)} samples\")\n        \n        best_metrics = trainer.train()\n        \n        # Print final results\n        print(\"\\nTraining completed!\")\n        print(\"Best metrics achieved:\")\n        print(f\"Best validation loss: {best_metrics['best_val_loss']:.4f}\")\n        print(f\"Best validation accuracy: {best_metrics['best_val_acc']:.2f}%\")\n        print(f\"Best epoch: {best_metrics['best_epoch']}\")\n        \n        # Clean up\n        del model, trainer\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        return best_metrics\n        \n    except Exception as e:\n        print(f\"Error in outlier detection pipeline: {str(e)}\")\n        traceback.print_exc()\n        \n        # Clean up even if there's an error\n        try:\n            del model, trainer\n            gc.collect()\n            torch.cuda.empty_cache()\n        except:\n            pass\n        return None\n\nif __name__ == \"__main__\":\n    train_outlier_main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:18:34.03486Z","iopub.execute_input":"2025-02-15T05:18:34.035183Z","iopub.status.idle":"2025-02-15T05:54:42.862095Z","shell.execute_reply.started":"2025-02-15T05:18:34.035119Z","shell.execute_reply":"2025-02-15T05:54:42.861361Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Outlier Analysis","metadata":{}},{"cell_type":"code","source":"def analyze_outliers(model, dataloader, device, threshold=3.0):\n    \"\"\"Analyze potential outliers in the dataset\"\"\"\n    model.eval()\n    outlier_scores = []\n    predictions = []\n    labels = []\n    \n    with torch.no_grad():\n        for images, batch_labels in dataloader:\n            images = images.to(device)\n            # Compute outlier scores directly\n            scores = model.compute_outlier_score(images)\n            # Get predictions\n            logits = model(images)  # Model returns only logits in eval mode\n            \n            outlier_scores.extend(scores.cpu().numpy())\n            predictions.extend(torch.argmax(logits, dim=1).cpu().numpy())\n            labels.extend(batch_labels.numpy())\n    \n    outlier_scores = np.array(outlier_scores)\n    predictions = np.array(predictions)\n    labels = np.array(labels)\n    \n    # Identify outliers using Z-score\n    z_scores = (outlier_scores - np.mean(outlier_scores)) / np.std(outlier_scores)\n    outliers = z_scores > threshold\n    \n    # Plot results\n    plt.figure(figsize=(15, 5))\n    \n    # Plot 1: Outlier scores distribution\n    plt.subplot(1, 2, 1)\n    plt.hist(z_scores, bins=50)\n    plt.axvline(threshold, color='r', linestyle='--', label=f'Threshold ({threshold})')\n    plt.title('Distribution of Outlier Scores')\n    plt.xlabel('Z-score')\n    plt.ylabel('Count')\n    plt.legend()\n    \n    # Plot 2: Scatter plot of outlier scores vs predictions\n    plt.subplot(1, 2, 2)\n    scatter = plt.scatter(predictions, z_scores, c=labels, cmap='viridis', alpha=0.6)\n    plt.axhline(threshold, color='r', linestyle='--', label=f'Threshold ({threshold})')\n    plt.title('Outlier Scores vs Predictions')\n    plt.xlabel('Predicted Class')\n    plt.ylabel('Outlier Score (Z-score)')\n    plt.legend()\n    plt.colorbar(scatter, label='True Class')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Print summary\n    print(f\"\\nFound {np.sum(outliers)} potential outliers out of {len(outlier_scores)} samples\")\n    print(f\"Outlier percentage: {100 * np.sum(outliers) / len(outlier_scores):.2f}%\")\n    \n    return outlier_scores, z_scores, outliers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:55:21.183793Z","iopub.execute_input":"2025-02-15T05:55:21.184116Z","iopub.status.idle":"2025-02-15T05:55:21.193276Z","shell.execute_reply.started":"2025-02-15T05:55:21.184092Z","shell.execute_reply":"2025-02-15T05:55:21.192385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_outliers(model, dataloader, device, threshold=3.0, num_samples=10):\n    \"\"\"Display sample images identified as outliers\"\"\"\n    model.eval()\n    classes = ['HGSC', 'EC', 'CC', 'LGSC', 'MC']\n    \n    # Collect images and their outlier scores\n    all_images = []\n    all_scores = []\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for images, labels in dataloader:\n            images = images.to(device)\n            # Get predictions\n            logits = model(images)  # Model returns only logits in eval mode\n            # Get outlier scores\n            scores = model.compute_outlier_score(images)\n            preds = torch.argmax(logits, dim=1)\n            \n            # Store batch data\n            all_images.extend(images.cpu())\n            all_scores.extend(scores.cpu().numpy())\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.numpy())\n    \n    # Convert to numpy arrays\n    all_scores = np.array(all_scores)\n    all_preds = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    \n    # Calculate z-scores\n    z_scores = (all_scores - np.mean(all_scores)) / np.std(all_scores)\n    \n    # Find outlier indices\n    outlier_indices = np.where(z_scores > threshold)[0]\n    \n    if len(outlier_indices) == 0:\n        print(\"No outliers found with the current threshold.\")\n        return\n    \n    # Sort outliers by score for most extreme cases\n    sorted_indices = outlier_indices[np.argsort(-z_scores[outlier_indices])]\n    \n    # Display top outliers\n    n_cols = 5\n    n_rows = (min(num_samples, len(sorted_indices)) + n_cols - 1) // n_cols\n    fig = plt.figure(figsize=(20, 4*n_rows))\n    \n    for idx, outlier_idx in enumerate(sorted_indices[:num_samples]):\n        ax = fig.add_subplot(n_rows, n_cols, idx + 1, xticks=[], yticks=[])\n        \n        # Get image and convert from tensor\n        img = all_images[outlier_idx].numpy().transpose((1, 2, 0))\n        img = np.clip(img, 0, 1)\n        \n        # Display image\n        ax.imshow(img)\n        \n        # Add title with prediction and outlier score\n        true_label = classes[all_labels[outlier_idx]]\n        pred_label = classes[all_preds[outlier_idx]]\n        score = z_scores[outlier_idx]\n        \n        title = f'True: {true_label}\\nPred: {pred_label}\\nOutlier Score: {score:.2f}'\n        ax.set_title(title, color='red', fontsize=10)\n    \n    plt.suptitle(f'Top {num_samples} Outliers (Threshold = {threshold})', fontsize=16)\n    plt.tight_layout()\n    plt.show()\n    \n    # Print summary statistics\n    print(f\"\\nTotal outliers found: {len(outlier_indices)} out of {len(z_scores)} images\")\n    print(f\"Percentage of outliers: {100 * len(outlier_indices) / len(z_scores):.2f}%\")\n    \n    # Show class distribution of outliers\n    print(\"\\nClass distribution of outliers:\")\n    for i, cls in enumerate(classes):\n        outlier_count = np.sum(all_labels[outlier_indices] == i)\n        total_count = np.sum(all_labels == i)\n        if total_count > 0:\n            percentage = 100 * outlier_count / total_count\n            print(f\"{cls}: {outlier_count}/{total_count} ({percentage:.2f}%)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:55:26.540994Z","iopub.execute_input":"2025-02-15T05:55:26.541511Z","iopub.status.idle":"2025-02-15T05:55:26.551544Z","shell.execute_reply.started":"2025-02-15T05:55:26.541478Z","shell.execute_reply":"2025-02-15T05:55:26.550589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_outlier_analysis():\n    \"\"\"Run the complete outlier analysis pipeline\"\"\"\n    try:\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"Using device: {device}\")\n        \n        # Load trained model\n        model = OutlierHistoPathModel()\n        model.load_state_dict(torch.load('./outlier_model_checkpoints/best_model.pth')['model_state_dict'])\n        model.to(device)\n        \n        # Different thresholds for comparison\n        thresholds = [2.0, 2.5, 3.0]\n        results = {}\n        \n        for threshold in thresholds:\n            print(f\"\\nAnalyzing outliers with threshold {threshold}...\")\n            outlier_scores, z_scores, outliers = analyze_outliers(\n                model,\n                val_loader,\n                device,\n                threshold=threshold\n            )\n            results[threshold] = {\n                'scores': outlier_scores,\n                'z_scores': z_scores,\n                'outliers': outliers\n            }\n            \n            # Visualize outliers for each threshold\n            print(f\"\\nVisualizing outliers for threshold {threshold}...\")\n            visualize_outliers(model, val_loader, device, threshold=threshold)\n        \n        # Clean up\n        del model\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        return results\n    \n    except Exception as e:\n        print(f\"Error in outlier analysis: {str(e)}\")\n        traceback.print_exc()\n        return None\n\nrun_outlier_analysis()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:55:30.646525Z","iopub.execute_input":"2025-02-15T05:55:30.646819Z","iopub.status.idle":"2025-02-15T05:55:44.336839Z","shell.execute_reply.started":"2025-02-15T05:55:30.646791Z","shell.execute_reply":"2025-02-15T05:55:44.336042Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test Prediction","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/test.csv\")\ntest_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:56:09.000898Z","iopub.execute_input":"2025-02-15T05:56:09.001221Z","iopub.status.idle":"2025-02-15T05:56:09.025149Z","shell.execute_reply.started":"2025-02-15T05:56:09.001194Z","shell.execute_reply":"2025-02-15T05:56:09.024454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_df = pd.read_csv(\"/kaggle/input/UBC-OCEAN/sample_submission.csv\")\nsample_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:56:12.59073Z","iopub.execute_input":"2025-02-15T05:56:12.591018Z","iopub.status.idle":"2025-02-15T05:56:12.604527Z","shell.execute_reply.started":"2025-02-15T05:56:12.590995Z","shell.execute_reply":"2025-02-15T05:56:12.603682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport cv2\nimport numpy as np\nfrom torchvision import models, transforms\nfrom torch.nn import functional as F\nfrom torchvision.models import ResNet101_Weights\nimport albumentations as A\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:56:16.027262Z","iopub.execute_input":"2025-02-15T05:56:16.027611Z","iopub.status.idle":"2025-02-15T05:56:16.031616Z","shell.execute_reply.started":"2025-02-15T05:56:16.027581Z","shell.execute_reply":"2025-02-15T05:56:16.030777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_single_image(image_path, model_path='./model_checkpoints/best_model.pth', device='cuda'):\n    \"\"\"\n    Predict class for a single test image using the saved best model\n    \"\"\"\n    try:\n        # Load and preprocess the image\n        image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # Initialize preprocessor\n        preprocessor = EnhancedPreprocessor()\n        preprocessed_image = preprocessor.preprocess_image(image)\n        \n        # Convert to tensor and normalize\n        transform = transforms.Compose([\n            transforms.ToTensor(),\n            transforms.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225]\n            )\n        ])\n        image_tensor = transform(preprocessed_image).unsqueeze(0)\n        \n        # Load the saved model\n        model = HistoPathModel()  # Initialize model architecture\n        checkpoint = torch.load(model_path)\n        model.load_state_dict(checkpoint['model_state_dict'])  # Load weights\n        model = model.to(device)\n        model.eval()\n        \n        # Make prediction\n        with torch.no_grad():\n            image_tensor = image_tensor.to(device)\n            outputs = model(image_tensor)\n            probabilities = F.softmax(outputs, dim=1)\n        \n        # Get prediction\n        classes = ['HGSC', 'EC', 'CC', 'LGSC', 'MC']\n        pred_class = classes[torch.argmax(probabilities).item()]\n        confidence = torch.max(probabilities).item()\n        \n        # Get probabilities for all classes\n        class_probs = {cls: prob.item() for cls, prob in zip(classes, probabilities[0])}\n        \n        return pred_class, confidence, class_probs\n        \n    except Exception as e:\n        print(f\"Error in prediction: {str(e)}\")\n        raise","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:56:44.822232Z","iopub.execute_input":"2025-02-15T05:56:44.822546Z","iopub.status.idle":"2025-02-15T05:56:44.829465Z","shell.execute_reply.started":"2025-02-15T05:56:44.82252Z","shell.execute_reply":"2025-02-15T05:56:44.828453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_prediction_to_csv(image_id, predicted_class, output_path='submission.csv'):\n    \"\"\"\n    Save prediction to CSV in the required format\n    \"\"\"\n    df = pd.DataFrame({\n        'image_id': [image_id],\n        'label': [predicted_class]\n    })\n    df.to_csv(output_path, index=False)\n    print(f\"\\nPrediction saved to {output_path}\")\n    print(f\"Content preview:\")\n    print(df.to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:56:54.327914Z","iopub.execute_input":"2025-02-15T05:56:54.328253Z","iopub.status.idle":"2025-02-15T05:56:54.332638Z","shell.execute_reply.started":"2025-02-15T05:56:54.328223Z","shell.execute_reply":"2025-02-15T05:56:54.33173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    # Paths\n    image_path = '/kaggle/input/UBC-OCEAN/test_thumbnails/41_thumbnail.png'\n    model_path = './model_checkpoints/best_model.pth'\n    \n    # Get image ID (removing leading zeros)\n    image_id = str(int(os.path.basename(image_path).split('_')[0]))\n    \n    try:\n        # Make prediction\n        pred_class, confidence, class_probs = predict_single_image(\n            image_path, \n            model_path, \n            device\n        )\n        \n        # Print detailed results\n        print(f\"\\nPredicted class: {pred_class}\")\n        print(f\"Confidence: {confidence:.2%}\")\n        print(\"\\nClass probabilities:\")\n        for cls, prob in class_probs.items():\n            print(f\"{cls}: {prob:.2%}\")\n        \n        # Save to CSV\n        save_prediction_to_csv(image_id, pred_class)\n        \n    except Exception as e:\n        print(f\"Error during prediction: {str(e)}\")\n        raise\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-15T05:56:57.731717Z","iopub.execute_input":"2025-02-15T05:56:57.732008Z","iopub.status.idle":"2025-02-15T05:57:01.191148Z","shell.execute_reply.started":"2025-02-15T05:56:57.731984Z","shell.execute_reply":"2025-02-15T05:57:01.190424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}