{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"datasetVersion","sourceId":15817686,"datasetId":10139365,"databundleVersionId":16766203}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Dataset, random_split\nfrom torchvision import datasets, models, transforms\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    confusion_matrix, accuracy_score, recall_score, f1_score, precision_score,\n    roc_auc_score, cohen_kappa_score, log_loss, classification_report,\n    balanced_accuracy_score, matthews_corrcoef, roc_curve, auc,\n    precision_recall_curve, average_precision_score, brier_score_loss\n)\nfrom sklearn.manifold import TSNE\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm\nimport warnings\nimport time\nimport cv2\nfrom datetime import datetime\nfrom collections import defaultdict\nimport psutil\nimport gc\nfrom torch.profiler import profile, record_function, ProfilerActivity\nfrom scipy.stats import entropy\nfrom scipy.spatial.distance import jensenshannon\nwarnings.filterwarnings('ignore')\nimport torch.nn.functional as F\nimport numpy as np\nfrom scipy.spatial.distance import jensenshannon\n\n# Set random seeds for reproducibility\ntorch.manual_seed(42)\nnp.random.seed(42)\n\nstart_time = time.time()\nstart_date = datetime.now()\nprint(f\"Training started at: {start_date}\")\n\nBASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\n\nbrain_activities = ['Seizure', 'GPD', 'LRDA', 'Other', 'GRDA', 'LPD']\nactivity_mapping = {activity: idx for idx, activity in enumerate(brain_activities)}\n\n# # Load and split data\n# df = pd.read_csv(f\"{BASE_DIR}train.csv\")\n# # df = df.sample(frac=0.01, random_state=42)  # Uncomment for quick testing\n\n# # Split 80% Train, 20% Temp (Validation + Test)\n# train_df, temp_df = train_test_split(df, test_size=0.2, random_state=42)\n# # train_df = train_df.head(60000)\n\n# # Split 10% Validation, 10% Test from Temp\n# val_df, test_df = train_test_split(temp_df, test_size=0.5, random_state=42)\n\ntrain_df = pd.read_csv(\"/kaggle/input/datasets/ashoksingh5972/mixed-subject-split/train.csv\")\ntrain_df = train_df.sample(frac=0.2, random_state=42)\n\nval_df = pd.read_csv(\"/kaggle/input/datasets/ashoksingh5972/mixed-subject-split/val.csv\")\ntest_df = pd.read_csv(\"/kaggle/input/datasets/ashoksingh5972/mixed-subject-split/test.csv\")\n\nprint(\"Splitting done with balanced training data!\")\nprint(\"Train:\", len(train_df), \"Val:\", len(val_df), \"Test:\", len(test_df))\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-07-26T15:53:33.995600Z","iopub.execute_input":"2025-07-26T15:53:33.996216Z","iopub.status.idle":"2025-07-26T15:53:44.331739Z","shell.execute_reply.started":"2025-07-26T15:53:33.996191Z","shell.execute_reply":"2025-07-26T15:53:44.331135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class ChunkedBrainActivityDataset(Dataset):\n#     def __init__(self, csv_file, base_dir, activity_mapping, md):\n#         self.df = csv_file\n#         self.base_dir = base_dir\n#         self.activity_mapping = activity_mapping\n#         self.resize_transform = transforms.Resize((224, 224))\n#         self.md = md\n\n#     def __len__(self):\n#         return len(self.df)\n\n#     def __getitem__(self, idx):\n#         spect_id, label, offset = self.df.iloc[idx][[\"spectrogram_id\", \"expert_consensus\", \"spectrogram_label_offset_seconds\"]]\n\n#         temp_df = pd.read_parquet(f'{self.base_dir}/train_spectrograms/{spect_id}.parquet')\n#         temp_df.drop(['time'], axis=1, inplace=True)\n\n#         start = int(offset) // 2\n#         temp_df = temp_df[start:start+300]\n#         temp_df = np.log1p(temp_df)\n#         temp_df /= temp_df.max()\n#         temp_arr = np.nan_to_num(temp_df.to_numpy(), nan=1e-4)\n\n#         # Use OpenCV to apply a colormap and convert to RGB\n#         temp_arr_uint8 = np.uint8(255 * temp_arr)\n#         rgb_image = cv2.applyColorMap(temp_arr_uint8, cv2.COLORMAP_JET)\n\n#         # Normalize to [0, 1] and convert to tensor\n#         rgb_image = rgb_image.astype(np.float32) / 255.0\n#         rgb_image_tensor = torch.tensor(rgb_image).permute(2, 0, 1)  # (C, H, W)\n#         rgb_image_tensor = self.resize_transform(rgb_image_tensor)\n            \n#         y = self.activity_mapping[label]\n#         y_tensor = torch.nn.functional.one_hot(torch.tensor(y, dtype=torch.long), num_classes=6).float()\n        \n#         return rgb_image_tensor, y_tensor\n\n\n# # Create datasets and data loaders\n# train_dataset = ChunkedBrainActivityDataset(csv_file=train_df, base_dir=BASE_DIR, activity_mapping=activity_mapping, md=\"lr\")\n# val_dataset = ChunkedBrainActivityDataset(csv_file=val_df, base_dir=BASE_DIR, activity_mapping=activity_mapping, md=\"lr\")\n# test_dataset = ChunkedBrainActivityDataset(csv_file=test_df, base_dir=BASE_DIR, activity_mapping=activity_mapping, md=\"lr\")\n\n# train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2, pin_memory=True, prefetch_factor=2)\n# val_loader = DataLoader(val_dataset, batch_size=64, shuffle=False, num_workers=2, pin_memory=True, prefetch_factor=2)\n# test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers=2, pin_memory=True, prefetch_factor=2)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T15:53:44.332447Z","iopub.execute_input":"2025-07-26T15:53:44.332700Z","iopub.status.idle":"2025-07-26T15:53:44.336925Z","shell.execute_reply.started":"2025-07-26T15:53:44.332682Z","shell.execute_reply":"2025-07-26T15:53:44.336238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef extract_middle_10sec(eeg, fs):\n    total_samples = eeg.shape[0]\n    total_duration_sec = total_samples / fs\n    \n    start_time = (total_duration_sec / 2) - 5\n    end_time = (total_duration_sec / 2) + 5\n    start_idx = int(start_time * fs)\n    end_idx = int(end_time * fs)\n    \n    return eeg[start_idx:end_idx]\n\n\nimport numpy as np\nfrom scipy.signal import stft\n\ndef compute_stft_spectrogram(eeg_10sec, fs):\n\n    # If multi-channel, process each channel separately and stack results\n    if eeg_10sec.ndim == 2:\n        # Example: Use mean across channels, or adapt as needed\n        eeg_input = eeg_10sec.mean(axis=1)\n    else:\n        eeg_input = eeg_10sec\n\n    # Compute STFT: nperseg=256 window size, 50% overlap\n    f, t, Zxx = stft(eeg_input, fs=fs, nperseg=256, noverlap=128)\n\n    # Power spectral density\n    Sxx = np.abs(Zxx) ** 2\n\n    # Log scaling and normalization\n    Sxx = np.log1p(Sxx)\n    Sxx /= (Sxx.max() + 1e-8)\n    return Sxx  # shape: (freq bins, time bins)\n\n\nimport pandas as pd\nimport numpy as np\nimport os\n\nEEG_CHANNELS = ['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1',\n                'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8',\n                'T4', 'T6', 'O2']\n\ndef load_eeg_signal(row):\n    BASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\n    spect_id = row['eeg_id']\n    eeg_path = os.path.join(BASE_DIR, \"train_eegs\", f\"{spect_id}.parquet\")\n    df = pd.read_parquet(eeg_path)\n    # Load as a [n_samples, n_channels] array\n    eeg = df[EEG_CHANNELS].to_numpy()\n    return eeg  # shape: (samples, 19)\n\nimport torch\nfrom torch.utils.data import Dataset\nimport cv2\nfrom torchvision import transforms\n\nclass Middle10SecSpectrogramDataset(Dataset):\n    def __init__(self, df, eeg_loader_func, activity_mapping, fs):\n        self.df = df  # DataFrame with EEG info and labels\n        self.eeg_loader_func = eeg_loader_func  # function to load EEG array by row info\n        self.activity_mapping = activity_mapping\n        self.fs = fs\n        self.resize_transform = transforms.Resize((224, 224))\n    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        eeg = self.eeg_loader_func(row)  # Returns (samples,) array\n\n        # 1. Extract middle 10sec\n        eeg_10sec = extract_middle_10sec(eeg, self.fs)\n\n        # 2. Spectrogram\n        Sxx = compute_stft_spectrogram(eeg_10sec, self.fs)\n\n        # 3. Make into RGB image\n        Sxx_uint8 = np.uint8(255 * Sxx)\n        rgb_img = cv2.applyColorMap(Sxx_uint8, cv2.COLORMAP_JET)  # H,W,3\n        rgb_img = rgb_img.astype(np.float32) / 255.0\n\n        # 4. Tensor and resize\n        img_tensor = torch.tensor(rgb_img).permute(2,0,1)  # C,H,W\n        img_tensor = self.resize_transform(img_tensor)\n\n        # 5. Label\n        y = self.activity_mapping[row[\"expert_consensus\"]]\n        y_tensor = torch.nn.functional.one_hot(torch.tensor(y), num_classes=len(self.activity_mapping)).float()\n\n        return img_tensor, y_tensor\n\nfs = 200\n\n# Create datasets with the new class and required arguments\ntrain_dataset = Middle10SecSpectrogramDataset(df=train_df, eeg_loader_func=load_eeg_signal, activity_mapping=activity_mapping, fs=fs)\nval_dataset = Middle10SecSpectrogramDataset(df=val_df, eeg_loader_func=load_eeg_signal, activity_mapping=activity_mapping, fs=fs)\ntest_dataset = Middle10SecSpectrogramDataset(df=test_df, eeg_loader_func=load_eeg_signal, activity_mapping=activity_mapping, fs=fs)\n\n# Create DataLoaders similarly\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2, pin_memory=True, prefetch_factor=2)\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=False, num_workers=2, pin_memory=True, prefetch_factor=2)\ntest_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers=2, pin_memory=True, prefetch_factor=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T15:53:44.338782Z","iopub.execute_input":"2025-07-26T15:53:44.339044Z","iopub.status.idle":"2025-07-26T15:53:44.471333Z","shell.execute_reply.started":"2025-07-26T15:53:44.339024Z","shell.execute_reply":"2025-07-26T15:53:44.470809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_kl_divergence(y_true, y_pred_probs, epsilon=1e-8):\n    \"\"\"\n    Calculate KL divergence between true and predicted probability distributions.\n    \n    Args:\n        y_true: True labels (numpy array or tensor)\n        y_pred_probs: Predicted probabilities (numpy array or tensor)\n        epsilon: Small value to avoid log(0)\n    \n    Returns:\n        dict: Contains various KL divergence metrics\n    \"\"\"\n    # Convert to numpy if needed\n    if torch.is_tensor(y_true):\n        y_true = y_true.cpu().numpy()\n    if torch.is_tensor(y_pred_probs):\n        y_pred_probs = y_pred_probs.cpu().numpy()\n    \n    num_classes = y_pred_probs.shape[1]\n    \n    # Convert labels to one-hot encoding\n    y_true_onehot = np.eye(num_classes)[y_true]\n    \n    # Add epsilon to avoid log(0)\n    y_pred_probs_safe = np.clip(y_pred_probs, epsilon, 1.0)\n    y_true_onehot_safe = np.clip(y_true_onehot, epsilon, 1.0)\n    \n    # Calculate KL divergence for each sample\n    # KL(P||Q) = sum(P * log(P/Q))\n    kl_divs = []\n    reverse_kl_divs = []\n    js_divs = []\n    \n    for i in range(len(y_true)):\n        # Forward KL: KL(true || pred)\n        kl_forward = np.sum(y_true_onehot_safe[i] * np.log(y_true_onehot_safe[i] / y_pred_probs_safe[i]))\n        kl_divs.append(kl_forward)\n        \n        # Reverse KL: KL(pred || true)\n        kl_reverse = np.sum(y_pred_probs_safe[i] * np.log(y_pred_probs_safe[i] / y_true_onehot_safe[i]))\n        reverse_kl_divs.append(kl_reverse)\n        \n        # Jensen-Shannon divergence (symmetric)\n        js_div = jensenshannon(y_true_onehot_safe[i], y_pred_probs_safe[i]) ** 2\n        js_divs.append(js_div)\n    \n    # Calculate class-wise KL divergence\n    class_kl_divs = []\n    for class_idx in range(num_classes):\n        class_mask = (y_true == class_idx)\n        if np.sum(class_mask) > 0:\n            class_true = y_true_onehot_safe[class_mask]\n            class_pred = y_pred_probs_safe[class_mask]\n            \n            # Average KL divergence for this class\n            class_kl = np.mean([\n                np.sum(class_true[j] * np.log(class_true[j] / class_pred[j]))\n                for j in range(len(class_true))\n            ])\n            class_kl_divs.append(class_kl)\n        else:\n            class_kl_divs.append(np.nan)\n    \n    # Calculate distribution-level KL divergence\n    # Compare overall class distributions\n    true_dist = np.bincount(y_true, minlength=num_classes) / len(y_true)\n    pred_dist = np.mean(y_pred_probs, axis=0)\n    \n    # Add epsilon to distributions\n    true_dist_safe = np.clip(true_dist, epsilon, 1.0)\n    pred_dist_safe = np.clip(pred_dist, epsilon, 1.0)\n    \n    # Normalize to ensure they sum to 1\n    true_dist_safe = true_dist_safe / np.sum(true_dist_safe)\n    pred_dist_safe = pred_dist_safe / np.sum(pred_dist_safe)\n    \n    dist_kl_forward = np.sum(true_dist_safe * np.log(true_dist_safe / pred_dist_safe))\n    dist_kl_reverse = np.sum(pred_dist_safe * np.log(pred_dist_safe / true_dist_safe))\n    dist_js = jensenshannon(true_dist_safe, pred_dist_safe) ** 2\n    \n    return {\n        'mean_kl_divergence': np.mean(kl_divs),\n        'std_kl_divergence': np.std(kl_divs),\n        'median_kl_divergence': np.median(kl_divs),\n        'mean_reverse_kl_divergence': np.mean(reverse_kl_divs),\n        'mean_js_divergence': np.mean(js_divs),\n        'class_kl_divergences': class_kl_divs,\n        'distribution_kl_forward': dist_kl_forward,\n        'distribution_kl_reverse': dist_kl_reverse,\n        'distribution_js_divergence': dist_js,\n        'sample_kl_divergences': kl_divs,\n        'sample_reverse_kl_divergences': reverse_kl_divs,\n        'sample_js_divergences': js_divs\n    }\n\ndef calculate_kl_divergence_pytorch(y_true, y_pred_logits, temperature=1.0):\n    \"\"\"\n    Calculate KL divergence using PyTorch (useful for differentiable operations).\n    \n    Args:\n        y_true: True labels (tensor)\n        y_pred_logits: Predicted logits (tensor)\n        temperature: Temperature scaling parameter\n    \n    Returns:\n        dict: KL divergence metrics as tensors\n    \"\"\"\n    # Apply temperature scaling\n    y_pred_probs = F.softmax(y_pred_logits / temperature, dim=1)\n    \n    # Convert labels to one-hot\n    num_classes = y_pred_logits.shape[1]\n    y_true_onehot = F.one_hot(y_true, num_classes=num_classes).float()\n    \n    # Calculate KL divergence\n    log_pred = F.log_softmax(y_pred_logits / temperature, dim=1)\n    \n    # KL(true || pred) = sum(true * log(true / pred))\n    # Since true is one-hot, this simplifies to -log(pred[true_class])\n    kl_div = F.kl_div(log_pred, y_true_onehot, reduction='none').sum(dim=1)\n    \n    # Reverse KL: KL(pred || true)\n    log_true = torch.log(y_true_onehot + 1e-8)\n    reverse_kl_div = F.kl_div(log_true, y_pred_probs, reduction='none').sum(dim=1)\n    \n    return {\n        'mean_kl_divergence': kl_div.mean(),\n        'std_kl_divergence': kl_div.std(),\n        'mean_reverse_kl_divergence': reverse_kl_div.mean(),\n        'sample_kl_divergences': kl_div,\n        'sample_reverse_kl_divergences': reverse_kl_div\n    }\n\ndef plot_kl_divergence_analysis(kl_results, class_names):\n    \"\"\"\n    Plot KL divergence analysis results.\n    \n    Args:\n        kl_results: Results from calculate_kl_divergence function\n        class_names: List of class names\n    \"\"\"\n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Sample-wise KL divergence distribution\n    axes[0, 0].hist(kl_results['sample_kl_divergences'], bins=50, alpha=0.7, color='blue', edgecolor='black')\n    axes[0, 0].axvline(kl_results['mean_kl_divergence'], color='red', linestyle='--', \n                       label=f'Mean: {kl_results[\"mean_kl_divergence\"]:.4f}')\n    axes[0, 0].axvline(kl_results['median_kl_divergence'], color='green', linestyle='--',\n                       label=f'Median: {kl_results[\"median_kl_divergence\"]:.4f}')\n    axes[0, 0].set_title('Distribution of Sample-wise KL Divergences')\n    axes[0, 0].set_xlabel('KL Divergence')\n    axes[0, 0].set_ylabel('Frequency')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Class-wise KL divergence\n    valid_class_kl = [kl for kl in kl_results['class_kl_divergences'] if not np.isnan(kl)]\n    valid_class_names = [name for i, name in enumerate(class_names) \n                        if not np.isnan(kl_results['class_kl_divergences'][i])]\n    \n    axes[0, 1].bar(valid_class_names, valid_class_kl, color='skyblue', edgecolor='black')\n    axes[0, 1].set_title('Class-wise KL Divergences')\n    axes[0, 1].set_xlabel('Classes')\n    axes[0, 1].set_ylabel('KL Divergence')\n    axes[0, 1].tick_params(axis='x', rotation=45)\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # Comparison of different divergence measures\n    divergence_types = ['KL (Forward)', 'KL (Reverse)', 'Jensen-Shannon']\n    divergence_values = [\n        kl_results['mean_kl_divergence'],\n        kl_results['mean_reverse_kl_divergence'],\n        kl_results['mean_js_divergence']\n    ]\n    \n    axes[1, 0].bar(divergence_types, divergence_values, \n                   color=['blue', 'red', 'green'], alpha=0.7, edgecolor='black')\n    axes[1, 0].set_title('Comparison of Divergence Measures')\n    axes[1, 0].set_ylabel('Divergence Value')\n    axes[1, 0].tick_params(axis='x', rotation=45)\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # KL divergence vs prediction confidence\n    max_probs = np.max(np.array(kl_results['sample_kl_divergences']).reshape(-1, 1), axis=1) \\\n                if len(np.array(kl_results['sample_kl_divergences']).shape) == 1 else \\\n                np.max(y_pred_probs, axis=1)  # This would need y_pred_probs from calling context\n    \n    # For now, create a scatter plot of KL vs sample index\n    sample_indices = range(len(kl_results['sample_kl_divergences']))\n    axes[1, 1].scatter(sample_indices, kl_results['sample_kl_divergences'], \n                       alpha=0.6, s=10, color='purple')\n    axes[1, 1].set_title('Sample-wise KL Divergence Pattern')\n    axes[1, 1].set_xlabel('Sample Index')\n    axes[1, 1].set_ylabel('KL Divergence')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T15:53:44.472063Z","iopub.execute_input":"2025-07-26T15:53:44.472476Z","iopub.status.idle":"2025-07-26T15:53:44.489968Z","shell.execute_reply.started":"2025-07-26T15:53:44.472452Z","shell.execute_reply":"2025-07-26T15:53:44.489259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass ResNet101EncoderLogisticRegression(nn.Module):\n    def __init__(self, num_classes=6):\n        super(ResNet101EncoderLogisticRegression, self).__init__()\n        # Load pretrained ResNet101\n        self.encoder = models.resnet101(weights=models.ResNet101_Weights.DEFAULT)\n        \n        # Get the feature size from the last layer (fc layer)\n        n_features = self.encoder.fc.in_features\n        \n        # Remove the original classifier head\n        self.encoder.fc = nn.Identity()\n        \n        # Add a logistic regression layer for classification\n        self.logistic_regression = nn.Linear(n_features, num_classes)\n\n    def forward(self, x):\n        features = self.encoder(x)  # Extract features\n        logits = self.logistic_regression(features)  # Apply classifier\n        return logits\n\n    def get_features(self, x):\n        \"\"\"Extract features for t-SNE visualization\"\"\"\n        with torch.no_grad():\n            features = self.encoder(x)\n        return features\n\n\n# Early Stopping Class\nclass EarlyStopping:\n    def __init__(self, patience=7, min_delta=0, restore_best_weights=True):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.restore_best_weights = restore_best_weights\n        self.best_loss = None\n        self.counter = 0\n        self.best_weights = None\n\n    def __call__(self, val_loss, model):\n        if self.best_loss is None:\n            self.best_loss = val_loss\n            self.save_checkpoint(model)\n        elif self.best_loss - val_loss > self.min_delta:\n            self.best_loss = val_loss\n            self.counter = 0\n            self.save_checkpoint(model)\n        else:\n            self.counter += 1\n\n        if self.counter >= self.patience:\n            if self.restore_best_weights:\n                model.load_state_dict(self.best_weights)\n            return True\n        return False\n\n    def save_checkpoint(self, model):\n        self.best_weights = model.state_dict().copy()\n\n\n# Set device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Instantiate the model and move it to the appropriate device\nnum_classes = 6\nmodel = ResNet101EncoderLogisticRegression(num_classes=num_classes).to(device)\n\n# Calculate model parameters\ndef count_parameters(model):\n    return sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nparam_count = count_parameters(model)\nprint(f\"Model parameters: {param_count / 1e6:.2f}M\")\n\n# Calculate model size\ndef get_model_size_mb(model):\n    param_size = 0\n    buffer_size = 0\n    for param in model.parameters():\n        param_size += param.nelement() * param.element_size()\n    for buffer in model.buffers():\n        buffer_size += buffer.nelement() * buffer.element_size()\n    size_mb = (param_size + buffer_size) / 1024 / 1024\n    return size_mb\n\nmodel_size_mb = get_model_size_mb(model)\nprint(f\"Model size: {model_size_mb:.2f} MB\")\n\nprint(\"ResNet101 model with logistic regression classifier loaded successfully!\")\n\n# Define the loss function and optimizer\noptimizer = optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)\nscheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)\n\n# Advanced loss functions\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n\n# Initialize early stopping\nearly_stopping = EarlyStopping(patience=5, min_delta=0.001)\n\n# Training history\nhistory = {\n    'train_loss': [],\n    'train_acc': [],\n    'val_loss': [],\n    'val_acc': [],\n    'learning_rate': []\n}\n\n# Enhanced training loop with validation and early stopping\nnum_epochs = 100\nbest_val_acc = 0.0\n\nprint(\"Starting training...\")\n\nfor epoch in range(num_epochs):\n    epoch_start_time = time.time()\n    \n    # Training phase\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    for images, targets in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} - Training\"):\n        images = images.to(device)\n        targets = targets.to(device)\n        labels = torch.argmax(targets, dim=1)\n        \n        optimizer.zero_grad()\n        logits = model(images)\n        loss = criterion(logits, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * images.size(0)\n        _, preds = torch.max(logits, dim=1)\n        total += labels.size(0)\n        correct += (preds == labels).sum().item()\n    \n    train_loss = running_loss / total\n    train_acc = 100 * correct / total\n    \n    # Validation phase\n    model.eval()\n    val_loss = 0.0\n    val_correct = 0\n    val_total = 0\n    \n    with torch.no_grad():\n        for images, targets in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} - Validation\"):\n            images = images.to(device)\n            targets = targets.to(device)\n            labels = torch.argmax(targets, dim=1)\n            \n            logits = model(images)\n            loss = criterion(logits, labels)\n            \n            val_loss += loss.item() * images.size(0)\n            _, preds = torch.max(logits, dim=1)\n            val_total += labels.size(0)\n            val_correct += (preds == labels).sum().item()\n    \n    val_loss = val_loss / val_total\n    val_acc = 100 * val_correct / val_total\n    \n    # Step the scheduler\n    scheduler.step()\n    current_lr = optimizer.param_groups[0]['lr']\n    \n    # Save training history\n    history['train_loss'].append(train_loss)\n    history['train_acc'].append(train_acc)\n    history['val_loss'].append(val_loss)\n    history['val_acc'].append(val_acc)\n    history['learning_rate'].append(current_lr)\n    \n    # Save best model\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), 'best_resnet101_model.pth')\n    \n    epoch_time = time.time() - epoch_start_time\n    \n    print(f\"Epoch [{epoch+1}/{num_epochs}] - Time: {epoch_time:.2f}s\")\n    print(f\"Train - Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n    print(f\"Val - Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n    print(f\"Learning Rate: {current_lr:.6f}\")\n    print(\"-\" * 50)\n    \n    # Early stopping check\n    if early_stopping(val_loss, model):\n        print(f\"Early stopping triggered after epoch {epoch+1}\")\n        break\n\n# Load best model for evaluation\nmodel.load_state_dict(torch.load('best_resnet101_model.pth'))\n\n# Plot training curves\ndef plot_training_curves(history):\n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Loss curves\n    axes[0, 0].plot(history['train_loss'], label='Train Loss', color='blue')\n    axes[0, 0].plot(history['val_loss'], label='Validation Loss', color='red')\n    axes[0, 0].set_title('Training and Validation Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True)\n    \n    # Accuracy curves\n    axes[0, 1].plot(history['train_acc'], label='Train Accuracy', color='blue')\n    axes[0, 1].plot(history['val_acc'], label='Validation Accuracy', color='red')\n    axes[0, 1].set_title('Training and Validation Accuracy')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Accuracy (%)')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True)\n    \n    # Learning rate\n    axes[1, 0].plot(history['learning_rate'], label='Learning Rate', color='green')\n    axes[1, 0].set_title('Learning Rate Schedule')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('Learning Rate')\n    axes[1, 0].legend()\n    axes[1, 0].grid(True)\n    axes[1, 0].set_yscale('log')\n    \n    # Combined loss and accuracy\n    ax1 = axes[1, 1]\n    ax2 = ax1.twinx()\n    \n    ax1.plot(history['train_loss'], 'b-', label='Train Loss')\n    ax1.plot(history['val_loss'], 'r-', label='Val Loss')\n    ax2.plot(history['train_acc'], 'b--', label='Train Acc')\n    ax2.plot(history['val_acc'], 'r--', label='Val Acc')\n    \n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Loss', color='black')\n    ax2.set_ylabel('Accuracy (%)', color='black')\n    ax1.set_title('Combined Loss and Accuracy')\n    \n    lines1, labels1 = ax1.get_legend_handles_labels()\n    lines2, labels2 = ax2.get_legend_handles_labels()\n    ax1.legend(lines1 + lines2, labels1 + labels2, loc='center right')\n    ax1.grid(True)\n    \n    plt.tight_layout()\n    plt.show()\n\nplot_training_curves(history)\n\n\ndef calculate_kl_divergence(y_true, y_pred_probs, epsilon=1e-8):\n    \"\"\"\n    Calculate KL divergence between true and predicted probability distributions.\n    \n    Args:\n        y_true: True labels (numpy array or tensor)\n        y_pred_probs: Predicted probabilities (numpy array or tensor)\n        epsilon: Small value to avoid log(0)\n    \n    Returns:\n        dict: Contains various KL divergence metrics\n    \"\"\"\n    # Convert to numpy if needed\n    if torch.is_tensor(y_true):\n        y_true = y_true.cpu().numpy()\n    if torch.is_tensor(y_pred_probs):\n        y_pred_probs = y_pred_probs.cpu().numpy()\n    \n    num_classes = y_pred_probs.shape[1]\n    \n    # Convert labels to one-hot encoding\n    y_true_onehot = np.eye(num_classes)[y_true]\n    \n    # Add epsilon to avoid log(0)\n    y_pred_probs_safe = np.clip(y_pred_probs, epsilon, 1.0)\n    y_true_onehot_safe = np.clip(y_true_onehot, epsilon, 1.0)\n    \n    # Calculate KL divergence for each sample\n    # KL(P||Q) = sum(P * log(P/Q))\n    kl_divs = []\n    reverse_kl_divs = []\n    js_divs = []\n    \n    for i in range(len(y_true)):\n        # Forward KL: KL(true || pred)\n        kl_forward = np.sum(y_true_onehot_safe[i] * np.log(y_true_onehot_safe[i] / y_pred_probs_safe[i]))\n        kl_divs.append(kl_forward)\n        \n        # Reverse KL: KL(pred || true)\n        kl_reverse = np.sum(y_pred_probs_safe[i] * np.log(y_pred_probs_safe[i] / y_true_onehot_safe[i]))\n        reverse_kl_divs.append(kl_reverse)\n        \n        # Jensen-Shannon divergence (symmetric)\n        js_div = jensenshannon(y_true_onehot_safe[i], y_pred_probs_safe[i]) ** 2\n        js_divs.append(js_div)\n    \n    # Calculate class-wise KL divergence\n    class_kl_divs = []\n    for class_idx in range(num_classes):\n        class_mask = (y_true == class_idx)\n        if np.sum(class_mask) > 0:\n            class_true = y_true_onehot_safe[class_mask]\n            class_pred = y_pred_probs_safe[class_mask]\n            \n            # Average KL divergence for this class\n            class_kl = np.mean([\n                np.sum(class_true[j] * np.log(class_true[j] / class_pred[j]))\n                for j in range(len(class_true))\n            ])\n            class_kl_divs.append(class_kl)\n        else:\n            class_kl_divs.append(np.nan)\n    \n    # Calculate distribution-level KL divergence\n    # Compare overall class distributions\n    true_dist = np.bincount(y_true, minlength=num_classes) / len(y_true)\n    pred_dist = np.mean(y_pred_probs, axis=0)\n    \n    # Add epsilon to distributions\n    true_dist_safe = np.clip(true_dist, epsilon, 1.0)\n    pred_dist_safe = np.clip(pred_dist, epsilon, 1.0)\n    \n    # Normalize to ensure they sum to 1\n    true_dist_safe = true_dist_safe / np.sum(true_dist_safe)\n    pred_dist_safe = pred_dist_safe / np.sum(pred_dist_safe)\n    \n    dist_kl_forward = np.sum(true_dist_safe * np.log(true_dist_safe / pred_dist_safe))\n    dist_kl_reverse = np.sum(pred_dist_safe * np.log(pred_dist_safe / true_dist_safe))\n    dist_js = jensenshannon(true_dist_safe, pred_dist_safe) ** 2\n    \n    return {\n        'mean_kl_divergence': np.mean(kl_divs),\n        'std_kl_divergence': np.std(kl_divs),\n        'median_kl_divergence': np.median(kl_divs),\n        'mean_reverse_kl_divergence': np.mean(reverse_kl_divs),\n        'mean_js_divergence': np.mean(js_divs),\n        'class_kl_divergences': class_kl_divs,\n        'distribution_kl_forward': dist_kl_forward,\n        'distribution_kl_reverse': dist_kl_reverse,\n        'distribution_js_divergence': dist_js,\n        'sample_kl_divergences': kl_divs,\n        'sample_reverse_kl_divergences': reverse_kl_divs,\n        'sample_js_divergences': js_divs\n    }\n\ndef calculate_kl_divergence_pytorch(y_true, y_pred_logits, temperature=1.0):\n    \"\"\"\n    Calculate KL divergence using PyTorch (useful for differentiable operations).\n    \n    Args:\n        y_true: True labels (tensor)\n        y_pred_logits: Predicted logits (tensor)\n        temperature: Temperature scaling parameter\n    \n    Returns:\n        dict: KL divergence metrics as tensors\n    \"\"\"\n    # Apply temperature scaling\n    y_pred_probs = F.softmax(y_pred_logits / temperature, dim=1)\n    \n    # Convert labels to one-hot\n    num_classes = y_pred_logits.shape[1]\n    y_true_onehot = F.one_hot(y_true, num_classes=num_classes).float()\n    \n    # Calculate KL divergence\n    log_pred = F.log_softmax(y_pred_logits / temperature, dim=1)\n    \n    # KL(true || pred) = sum(true * log(true / pred))\n    # Since true is one-hot, this simplifies to -log(pred[true_class])\n    kl_div = F.kl_div(log_pred, y_true_onehot, reduction='none').sum(dim=1)\n    \n    # Reverse KL: KL(pred || true)\n    log_true = torch.log(y_true_onehot + 1e-8)\n    reverse_kl_div = F.kl_div(log_true, y_pred_probs, reduction='none').sum(dim=1)\n    \n    return {\n        'mean_kl_divergence': kl_div.mean(),\n        'std_kl_divergence': kl_div.std(),\n        'mean_reverse_kl_divergence': reverse_kl_div.mean(),\n        'sample_kl_divergences': kl_div,\n        'sample_reverse_kl_divergences': reverse_kl_div\n    }\n\ndef plot_kl_divergence_analysis(kl_results, class_names):\n    \"\"\"\n    Plot KL divergence analysis results.\n    \n    Args:\n        kl_results: Results from calculate_kl_divergence function\n        class_names: List of class names\n    \"\"\"\n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Sample-wise KL divergence distribution\n    axes[0, 0].hist(kl_results['sample_kl_divergences'], bins=50, alpha=0.7, color='blue', edgecolor='black')\n    axes[0, 0].axvline(kl_results['mean_kl_divergence'], color='red', linestyle='--', \n                       label=f'Mean: {kl_results[\"mean_kl_divergence\"]:.4f}')\n    axes[0, 0].axvline(kl_results['median_kl_divergence'], color='green', linestyle='--',\n                       label=f'Median: {kl_results[\"median_kl_divergence\"]:.4f}')\n    axes[0, 0].set_title('Distribution of Sample-wise KL Divergences')\n    axes[0, 0].set_xlabel('KL Divergence')\n    axes[0, 0].set_ylabel('Frequency')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Class-wise KL divergence\n    valid_class_kl = [kl for kl in kl_results['class_kl_divergences'] if not np.isnan(kl)]\n    valid_class_names = [name for i, name in enumerate(class_names) \n                        if not np.isnan(kl_results['class_kl_divergences'][i])]\n    \n    axes[0, 1].bar(valid_class_names, valid_class_kl, color='skyblue', edgecolor='black')\n    axes[0, 1].set_title('Class-wise KL Divergences')\n    axes[0, 1].set_xlabel('Classes')\n    axes[0, 1].set_ylabel('KL Divergence')\n    axes[0, 1].tick_params(axis='x', rotation=45)\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # Comparison of different divergence measures\n    divergence_types = ['KL (Forward)', 'KL (Reverse)', 'Jensen-Shannon']\n    divergence_values = [\n        kl_results['mean_kl_divergence'],\n        kl_results['mean_reverse_kl_divergence'],\n        kl_results['mean_js_divergence']\n    ]\n    \n    axes[1, 0].bar(divergence_types, divergence_values, \n                   color=['blue', 'red', 'green'], alpha=0.7, edgecolor='black')\n    axes[1, 0].set_title('Comparison of Divergence Measures')\n    axes[1, 0].set_ylabel('Divergence Value')\n    axes[1, 0].tick_params(axis='x', rotation=45)\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # KL divergence vs prediction confidence\n    max_probs = np.max(np.array(kl_results['sample_kl_divergences']).reshape(-1, 1), axis=1) \\\n                if len(np.array(kl_results['sample_kl_divergences']).shape) == 1 else \\\n                np.max(y_pred_probs, axis=1)  # This would need y_pred_probs from calling context\n    \n    # For now, create a scatter plot of KL vs sample index\n    sample_indices = range(len(kl_results['sample_kl_divergences']))\n    axes[1, 1].scatter(sample_indices, kl_results['sample_kl_divergences'], \n                       alpha=0.6, s=10, color='purple')\n    axes[1, 1].set_title('Sample-wise KL Divergence Pattern')\n    axes[1, 1].set_xlabel('Sample Index')\n    axes[1, 1].set_ylabel('KL Divergence')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n\n\n# Comprehensive evaluation functions\ndef calculate_ece(y_true, y_prob, n_bins=10):\n    \"\"\"Calculate Expected Calibration Error\"\"\"\n    bin_boundaries = np.linspace(0, 1, n_bins + 1)\n    bin_lowers = bin_boundaries[:-1]\n    bin_uppers = bin_boundaries[1:]\n    \n    ece = 0\n    for bin_lower, bin_upper in zip(bin_lowers, bin_uppers):\n        in_bin = (y_prob > bin_lower) & (y_prob <= bin_upper)\n        prop_in_bin = in_bin.mean()\n        \n        if prop_in_bin > 0:\n            accuracy_in_bin = y_true[in_bin].mean()\n            avg_confidence_in_bin = y_prob[in_bin].mean()\n            ece += np.abs(avg_confidence_in_bin - accuracy_in_bin) * prop_in_bin\n    \n    return ece\n\n\ndef reliability_diagram(y_true, y_prob, n_bins=10):\n    \"\"\"Plot reliability diagram\"\"\"\n    bin_boundaries = np.linspace(0, 1, n_bins + 1)\n    bin_lowers = bin_boundaries[:-1]\n    bin_uppers = bin_boundaries[1:]\n    \n    accuracies = []\n    confidences = []\n    \n    for bin_lower, bin_upper in zip(bin_lowers, bin_uppers):\n        in_bin = (y_prob > bin_lower) & (y_prob <= bin_upper)\n        prop_in_bin = in_bin.mean()\n        \n        if prop_in_bin > 0:\n            accuracy_in_bin = y_true[in_bin].mean()\n            avg_confidence_in_bin = y_prob[in_bin].mean()\n            accuracies.append(accuracy_in_bin)\n            confidences.append(avg_confidence_in_bin)\n    \n    plt.figure(figsize=(8, 6))\n    plt.plot([0, 1], [0, 1], 'k--', label='Perfect Calibration')\n    plt.plot(confidences, accuracies, 'o-', label='Model')\n    plt.xlabel('Mean Predicted Probability')\n    plt.ylabel('Fraction of Positives')\n    plt.title('Reliability Diagram')\n    plt.legend()\n    plt.grid(True)\n    plt.show()\n\n\ndef geometric_mean_score(y_true, y_pred):\n    \"\"\"Calculate geometric mean of class-wise recalls\"\"\"\n    cm = confusion_matrix(y_true, y_pred)\n    recalls = []\n    for i in range(cm.shape[0]):\n        if cm[i, :].sum() > 0:\n            recall = cm[i, i] / cm[i, :].sum()\n            recalls.append(recall)\n    return np.prod(recalls) ** (1.0 / len(recalls))\n\n\ndef calculate_specificity(y_true, y_pred, num_classes):\n    \"\"\"Calculate class-wise specificity\"\"\"\n    cm = confusion_matrix(y_true, y_pred)\n    specificities = []\n    for i in range(num_classes):\n        tn = cm.sum() - (cm[i, :].sum() + cm[:, i].sum() - cm[i, i])\n        fp = cm[:, i].sum() - cm[i, i]\n        specificity = tn / (tn + fp) if (tn + fp) > 0 else 0\n        specificities.append(specificity)\n    return specificities\n\n\ndef add_noise_and_evaluate(model, test_loader, device, noise_levels=[0.1, 0.2, 0.3]):\n    \"\"\"Evaluate model robustness to noise\"\"\"\n    results = {}\n    model.eval()\n    \n    for noise_level in noise_levels:\n        correct = 0\n        total = 0\n        \n        with torch.no_grad():\n            for images, targets in test_loader:\n                images = images.to(device)\n                labels = torch.argmax(targets, dim=1).to(device)\n                \n                # Add Gaussian noise\n                noise = torch.randn_like(images) * noise_level\n                noisy_images = images + noise\n                noisy_images = torch.clamp(noisy_images, 0, 1)\n                \n                outputs = model(noisy_images)\n                _, preds = torch.max(outputs, 1)\n                \n                total += labels.size(0)\n                correct += (preds == labels).sum().item()\n        \n        accuracy = 100 * correct / total\n        results[noise_level] = accuracy\n        print(f\"Noise level {noise_level}: {accuracy:.2f}%\")\n    \n    return results\n\n\ndef measure_inference_time(model, test_loader, device, num_samples=100):\n    \"\"\"Measure average inference time\"\"\"\n    model.eval()\n    times = []\n    \n    with torch.no_grad():\n        for i, (images, _) in enumerate(test_loader):\n            if i >= num_samples // images.size(0):\n                break\n                \n            images = images.to(device)\n            \n            start_time = time.time()\n            _ = model(images)\n            torch.cuda.synchronize() if device.type == 'cuda' else None\n            end_time = time.time()\n            \n            batch_time = (end_time - start_time) * 1000  # Convert to ms\n            per_sample_time = batch_time / images.size(0)\n            times.extend([per_sample_time] * images.size(0))\n    \n    return np.mean(times), np.std(times)\n\n\ndef comprehensive_evaluation(model, data_loader, device, class_names):\n    \"\"\"Calculate all evaluation metrics\"\"\"\n    model.eval()\n    all_preds = []\n    all_labels = []\n    all_probs = []\n    all_features = []\n    \n    # Memory usage before inference\n    if device.type == 'cuda':\n        torch.cuda.empty_cache()\n        memory_before = torch.cuda.memory_allocated() / 1024 / 1024  # MB\n    else:\n        memory_before = psutil.virtual_memory().used / 1024 / 1024\n    \n    with torch.no_grad():\n        for images, targets in tqdm(data_loader, desc=\"Evaluating\"):\n            images = images.to(device)\n            targets = targets.to(device)\n            labels = torch.argmax(targets, dim=1)\n            \n            # Forward pass\n            outputs = model(images)\n            probs = torch.softmax(outputs, dim=1)\n            \n            # Get features for t-SNE\n            features = model.get_features(images)\n            \n            _, preds = torch.max(outputs, dim=1)\n            \n            # Collect predictions, labels, probabilities, and features\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n            all_features.extend(features.cpu().numpy())\n    \n    # Memory usage after inference\n    if device.type == 'cuda':\n        memory_after = torch.cuda.memory_allocated() / 1024 / 1024  # MB\n    else:\n        memory_after = psutil.virtual_memory().used / 1024 / 1024\n    \n    inference_memory = memory_after - memory_before\n    \n    # Convert to numpy arrays\n    all_preds = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    all_probs = np.array(all_probs)\n    all_features = np.array(all_features)\n    \n    results = {}\n    \n    # Basic metrics\n    results['accuracy'] = accuracy_score(all_labels, all_preds)\n    results['precision_macro'] = precision_score(all_labels, all_preds, average='macro', zero_division=0)\n    results['recall_macro'] = recall_score(all_labels, all_preds, average='macro', zero_division=0)\n    results['f1_macro'] = f1_score(all_labels, all_preds, average='macro', zero_division=0)\n    results['balanced_accuracy'] = balanced_accuracy_score(all_labels, all_preds)\n    results['cohen_kappa'] = cohen_kappa_score(all_labels, all_preds)\n    results['matthews_corrcoef'] = matthews_corrcoef(all_labels, all_preds)\n    results['geometric_mean'] = geometric_mean_score(all_labels, all_preds)\n    \n    # Multiclass metrics\n    try:\n        results['auc_roc_macro'] = roc_auc_score(all_labels, all_probs, multi_class='ovr', average='macro')\n        results['log_loss'] = log_loss(all_labels, all_probs)\n    except ValueError as e:\n        print(f\"Warning: Could not calculate AUC-ROC or log loss: {e}\")\n        results['auc_roc_macro'] = np.nan\n        results['log_loss'] = np.nan\n    \n    # Brier score (for multiclass, we'll use the average)\n    brier_scores = []\n    for i in range(len(class_names)):\n        y_true_binary = (all_labels == i).astype(int)\n        y_prob_binary = all_probs[:, i]\n        brier_scores.append(brier_score_loss(y_true_binary, y_prob_binary))\n    results['brier_score'] = np.mean(brier_scores)\n    \n    # Expected Calibration Error\n    max_probs = np.max(all_probs, axis=1)\n    predicted_correctly = (all_preds == all_labels).astype(int)\n    results['ece'] = calculate_ece(predicted_correctly, max_probs)\n    \n    # Class-wise metrics\n    cm = confusion_matrix(all_labels, all_preds)\n    class_precision = precision_score(all_labels, all_preds, average=None, zero_division=0)\n    class_recall = recall_score(all_labels, all_preds, average=None, zero_division=0)\n    class_f1 = f1_score(all_labels, all_preds, average=None, zero_division=0)\n    class_specificity = calculate_specificity(all_labels, all_preds, len(class_names))\n    \n    # Class-wise AUC-ROC and AUC-PRC\n    class_auc_roc = []\n    class_auc_prc = []\n    class_avg_precision = []\n    \n    for i in range(len(class_names)):\n        try:\n            y_true_binary = (all_labels == i).astype(int)\n            y_score_binary = all_probs[:, i]\n            \n            # AUC-ROC\n            fpr, tpr, _ = roc_curve(y_true_binary, y_score_binary)\n            auc_roc = auc(fpr, tpr)\n            class_auc_roc.append(auc_roc)\n            \n            # AUC-PRC\n            precision_curve, recall_curve, _ = precision_recall_curve(y_true_binary, y_score_binary)\n            auc_prc = auc(recall_curve, precision_curve)\n            class_auc_prc.append(auc_prc)\n            \n            # Average Precision\n            avg_precision = average_precision_score(y_true_binary, y_score_binary)\n            class_avg_precision.append(avg_precision)\n            \n        except ValueError:\n            class_auc_roc.append(np.nan)\n            class_auc_prc.append(np.nan)\n            class_avg_precision.append(np.nan)\n    \n    # Store class-wise results\n    results['class_precision'] = class_precision\n    results['class_recall'] = class_recall\n    results['class_f1'] = class_f1\n    results['class_specificity'] = class_specificity\n    results['class_auc_roc'] = class_auc_roc\n    results['class_auc_prc'] = class_auc_prc\n    results['class_avg_precision'] = class_avg_precision\n    \n    # Top-k accuracy (top-3 for 6 classes)\n    top_k = min(3, len(class_names))\n    top_k_preds = np.argsort(all_probs, axis=1)[:, -top_k:]\n    top_k_correct = np.any(top_k_preds == all_labels.reshape(-1, 1), axis=1)\n    results[f'top_{top_k}_accuracy'] = np.mean(top_k_correct)\n    \n    # Macro/Micro/Weighted averages\n    results['precision_micro'] = precision_score(all_labels, all_preds, average='micro', zero_division=0)\n    results['precision_weighted'] = precision_score(all_labels, all_preds, average='weighted', zero_division=0)\n    results['recall_micro'] = recall_score(all_labels, all_preds, average='micro', zero_division=0)\n    results['recall_weighted'] = recall_score(all_labels, all_preds, average='weighted', zero_division=0)\n    results['f1_micro'] = f1_score(all_labels, all_preds, average='micro', zero_division=0)\n    results['f1_weighted'] = f1_score(all_labels, all_preds, average='weighted', zero_division=0)\n    \n    # Additional metrics\n    results['confusion_matrix'] = cm\n    results['classification_report'] = classification_report(all_labels, all_preds, target_names=class_names)\n    results['inference_memory_mb'] = inference_memory\n    \n    # Store data for visualizations\n    results['all_labels'] = all_labels\n    results['all_preds'] = all_preds  \n    results['all_probs'] = all_probs\n    results['all_features'] = all_features\n\n    # Calculate KL divergence\n    print(\"Calculating KL divergence...\")\n    kl_results = calculate_kl_divergence(all_labels, all_probs)\n    \n    # Store KL divergence results\n    results['kl_divergence_mean'] = kl_results['mean_kl_divergence']\n    results['kl_divergence_std'] = kl_results['std_kl_divergence']\n    results['kl_divergence_median'] = kl_results['median_kl_divergence']\n    results['reverse_kl_divergence_mean'] = kl_results['mean_reverse_kl_divergence']\n    results['js_divergence_mean'] = kl_results['mean_js_divergence']\n    results['distribution_kl_forward'] = kl_results['distribution_kl_forward']\n    results['distribution_kl_reverse'] = kl_results['distribution_kl_reverse']\n    results['class_kl_divergences'] = kl_results['class_kl_divergences']\n    results['kl_results_full'] = kl_results\n\n    \n    return results\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-26T15:53:44.490688Z","iopub.execute_input":"2025-07-26T15:53:44.490869Z","execution_failed":"2025-07-26T15:54:04.173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Perform comprehensive evaluation\nprint(\"Performing comprehensive evaluation...\")\n\nclass_names = ['Seizure', 'GPD', 'LRDA', 'Other', 'GRDA', 'LPD']\ntest_results = comprehensive_evaluation(model, test_loader, device, class_names)\n\n# Measure inference time and throughput\nprint(\"Measuring inference performance...\")\navg_inference_time, std_inference_time = measure_inference_time(model, test_loader, device)\nthroughput = 1000 / avg_inference_time  # samples per second\n\n# Noise robustness\nprint(\"Testing noise robustness...\")\nnoise_results = add_noise_and_evaluate(model, test_loader, device)\n\n# Calculate training time\nend_time = time.time()\ntotal_training_time = end_time - start_time\nend_date = datetime.now()\n\nprint(f\"Training ended at: {end_date}\")\n\n# Print comprehensive results\nprint(\"\\n\" + \"=\"*80)\nprint(\"COMPREHENSIVE EVALUATION RESULTS\")\nprint(\"=\"*80)\n\nprint(f\"\\n📊 BASIC METRICS:\")\nprint(f\"Accuracy: {test_results['accuracy']:.4f}\")\nprint(f\"Precision (Macro): {test_results['precision_macro']:.4f}\")\nprint(f\"Recall (Macro): {test_results['recall_macro']:.4f}\")\nprint(f\"F1 Score (Macro): {test_results['f1_macro']:.4f}\")\nprint(f\"AUC-ROC (Macro): {test_results['auc_roc_macro']:.4f}\")\nprint(f\"Cohen's Kappa: {test_results['cohen_kappa']:.4f}\")\nprint(f\"Log Loss: {test_results['log_loss']:.4f}\")\nprint(f\"Balanced Accuracy: {test_results['balanced_accuracy']:.4f}\")\nprint(f\"Matthews Correlation Coefficient: {test_results['matthews_corrcoef']:.4f}\")\nprint(f\"Geometric Mean Score: {test_results['geometric_mean']:.4f}\")\n\nprint(f\"\\n📈 CALIBRATION METRICS:\")\nprint(f\"Expected Calibration Error (ECE): {test_results['ece']:.4f}\")\nprint(f\"Brier Score: {test_results['brier_score']:.4f}\")\n\nprint(f\"\\n⚡ PERFORMANCE METRICS:\")\nprint(f\"Avg Inference Time: {avg_inference_time:.2f} ± {std_inference_time:.2f} ms/sample\")\nprint(f\"Throughput: {throughput:.2f} samples/sec\")\nprint(f\"Parameter Count: {param_count / 1e6:.2f}M\")\nprint(f\"Model Size: {model_size_mb:.2f} MB\")\nprint(f\"Training Time: {total_training_time/3600:.2f} hours\")\nprint(f\"Inference Memory: {test_results['inference_memory_mb']:.2f} MB\")\n\nprint(f\"\\n🎯 MULTI-AVERAGE METRICS:\")\nprint(f\"Precision - Micro: {test_results['precision_micro']:.4f}, Weighted: {test_results['precision_weighted']:.4f}\")\nprint(f\"Recall - Micro: {test_results['recall_micro']:.4f}, Weighted: {test_results['recall_weighted']:.4f}\")\nprint(f\"F1 Score - Micro: {test_results['f1_micro']:.4f}, Weighted: {test_results['f1_weighted']:.4f}\")\nprint(f\"Top-3 Accuracy: {test_results['top_3_accuracy']:.4f}\")\n\nprint(f\"\\n🔍 CLASS-WISE METRICS:\")\nfor i, class_name in enumerate(class_names):\n    print(f\"{class_name}:\")\n    print(f\"  Precision: {test_results['class_precision'][i]:.4f}\")\n    print(f\"  Recall: {test_results['class_recall'][i]:.4f}\")\n    print(f\"  F1 Score: {test_results['class_f1'][i]:.4f}\")\n    print(f\"  Specificity: {test_results['class_specificity'][i]:.4f}\")\n    print(f\"  AUC-ROC: {test_results['class_auc_roc'][i]:.4f}\")\n    print(f\"  AUC-PRC: {test_results['class_auc_prc'][i]:.4f}\")\n    print(f\"  Avg Precision: {test_results['class_avg_precision'][i]:.4f}\")\n\nprint(f\"\\n🛡️ ROBUSTNESS METRICS:\")\nprint(\"Noise Robustness:\")\nfor noise_level, accuracy in noise_results.items():\n    print(f\"  Noise Level {noise_level}: {accuracy:.2f}%\")\n\nprint(f\"\\n📅 TIMING INFORMATION:\")\nprint(f\"Start Date: {start_date}\")\nprint(f\"End Date: {end_date}\")\nprint(f\"Total Duration: {total_training_time/60:.2f} minutes\")\n\n# Detailed Classification Report\nprint(f\"\\n📋 CLASSIFICATION REPORT:\")\nprint(test_results['classification_report'])\n\n# Plot confusion matrix\ndef plot_confusion_matrix(cm, class_names):\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                xticklabels=class_names, yticklabels=class_names)\n    plt.xlabel('Predicted labels')\n    plt.ylabel('True labels')\n    plt.title('Confusion Matrix - ResNet101 + Logistic Regression')\n    plt.tight_layout()\n    plt.show()\n\nprint(\"\\n🔥 CONFUSION MATRIX:\")\nplot_confusion_matrix(test_results['confusion_matrix'], class_names)\n\n# Reliability diagram\nprint(\"\\n📊 RELIABILITY DIAGRAM:\")\nmax_probs = np.max(test_results['all_probs'], axis=1)\npredicted_correctly = (test_results['all_preds'] == test_results['all_labels']).astype(int)\nreliability_diagram(predicted_correctly, max_probs)\n\nprint(f\"\\n📊 KL DIVERGENCE METRICS:\")\nprint(f\"Mean KL Divergence: {test_results['kl_divergence_mean']:.4f}\")\nprint(f\"Std KL Divergence: {test_results['kl_divergence_std']:.4f}\")\nprint(f\"Median KL Divergence: {test_results['kl_divergence_median']:.4f}\")\nprint(f\"Mean Reverse KL Divergence: {test_results['reverse_kl_divergence_mean']:.4f}\")\nprint(f\"Mean Jensen-Shannon Divergence: {test_results['js_divergence_mean']:.4f}\")\nprint(f\"Distribution KL (Forward): {test_results['distribution_kl_forward']:.4f}\")\nprint(f\"Distribution KL (Reverse): {test_results['distribution_kl_reverse']:.4f}\")\n\nprint(f\"\\n🔍 CLASS-WISE KL DIVERGENCES:\")\nfor i, (class_name, kl_div) in enumerate(zip(class_names, test_results['class_kl_divergences'])):\n    if not np.isnan(kl_div):\n        print(f\"  {class_name}: {kl_div:.4f}\")\n\n# Plot KL divergence analysis\nprint(\"\\n📈 KL DIVERGENCE ANALYSIS:\")\nplot_kl_divergence_analysis(test_results['kl_results_full'], class_names)\n\n\n# ROC Curves for each class\ndef plot_multiclass_roc_curves(y_true, y_probs, class_names):\n    plt.figure(figsize=(12, 8))\n    \n    for i, class_name in enumerate(class_names):\n        y_true_binary = (y_true == i).astype(int)\n        y_score_binary = y_probs[:, i]\n        \n        fpr, tpr, _ = roc_curve(y_true_binary, y_score_binary)\n        roc_auc = auc(fpr, tpr)\n        \n        plt.plot(fpr, tpr, label=f'{class_name} (AUC = {roc_auc:.3f})')\n    \n    plt.plot([0, 1], [0, 1], 'k--', label='Random Classifier')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title('ROC Curves for Each Class')\n    plt.legend(loc=\"lower right\")\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n\nprint(\"\\n📈 ROC CURVES:\")\nplot_multiclass_roc_curves(test_results['all_labels'], test_results['all_probs'], class_names)\n\n# Precision-Recall Curves\ndef plot_multiclass_pr_curves(y_true, y_probs, class_names):\n    plt.figure(figsize=(12, 8))\n    \n    for i, class_name in enumerate(class_names):\n        y_true_binary = (y_true == i).astype(int)\n        y_score_binary = y_probs[:, i]\n        \n        precision, recall, _ = precision_recall_curve(y_true_binary, y_score_binary)\n        pr_auc = auc(recall, precision)\n        avg_precision = average_precision_score(y_true_binary, y_score_binary)\n        \n        plt.plot(recall, precision, label=f'{class_name} (AP = {avg_precision:.3f})')\n    \n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('Recall')\n    plt.ylabel('Precision')\n    plt.title('Precision-Recall Curves for Each Class')\n    plt.legend(loc=\"lower left\")\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n\nprint(\"\\n📊 PRECISION-RECALL CURVES:\")\nplot_multiclass_pr_curves(test_results['all_labels'], test_results['all_probs'], class_names)\n\n# t-SNE Visualization\ndef plot_tsne(features, labels, class_names, title=\"t-SNE Visualization\"):\n    print(\"Computing t-SNE... This may take a while...\")\n    \n    # Subsample for faster computation if dataset is large\n    if len(features) > 2000:\n        indices = np.random.choice(len(features), 2000, replace=False)\n        features = features[indices]\n        labels = labels[indices]\n    \n    # Compute t-SNE\n    tsne = TSNE(n_components=2, random_state=42, perplexity=30, n_iter=1000)\n    features_2d = tsne.fit_transform(features)\n    \n    # Plot\n    plt.figure(figsize=(12, 8))\n    colors = plt.cm.Set3(np.linspace(0, 1, len(class_names)))\n    \n    for i, (class_name, color) in enumerate(zip(class_names, colors)):\n        mask = labels == i\n        plt.scatter(features_2d[mask, 0], features_2d[mask, 1], \n                   c=[color], label=class_name, alpha=0.7, s=20)\n    \n    plt.xlabel('t-SNE Component 1')\n    plt.ylabel('t-SNE Component 2')\n    plt.title(title)\n    plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')\n    plt.grid(True, alpha=0.3)\n    plt.tight_layout()\n    plt.show()\n\nprint(\"\\n🗺️ t-SNE VISUALIZATION:\")\nplot_tsne(test_results['all_features'], test_results['all_labels'], class_names, \n          \"t-SNE Visualization of Learned Features\")\n\n# Feature importance heatmap (using gradients)\ndef plot_feature_importance_heatmap(model, test_loader, device, class_names):\n    model.eval()\n    \n    # Get one batch for gradient analysis\n    images, targets = next(iter(test_loader))\n    images = images.to(device)\n    images.requires_grad_(True)\n    \n    # Forward pass\n    outputs = model(images)\n    \n    # Calculate gradients for each class\n    gradients = []\n    for class_idx in range(len(class_names)):\n        model.zero_grad()\n        class_output = outputs[:, class_idx].sum()\n        class_output.backward(retain_graph=True)\n        \n        # Get gradient magnitude\n        grad = torch.abs(images.grad).mean(dim=(0, 1)).cpu().numpy()\n        gradients.append(grad)\n    \n    # Plot heatmap\n    gradients = np.array(gradients)\n    plt.figure(figsize=(12, 8))\n    sns.heatmap(gradients, xticklabels=False, yticklabels=class_names, \n                cmap='viridis', cbar=True)\n    plt.title('Feature Importance Heatmap (Gradient-based)')\n    plt.xlabel('Input Features')\n    plt.ylabel('Classes')\n    plt.tight_layout()\n    plt.show()\n\nprint(\"\\n🔥 FEATURE IMPORTANCE:\")\ntry:\n    plot_feature_importance_heatmap(model, test_loader, device, class_names)\nexcept Exception as e:\n    print(f\"Could not generate feature importance heatmap: {e}\")\n\n# Learning curve analysis\ndef plot_detailed_learning_curves(history):\n    fig, axes = plt.subplots(2, 3, figsize=(18, 10))\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    # Loss comparison\n    axes[0, 0].plot(epochs, history['train_loss'], 'b-', label='Training Loss', linewidth=2)\n    axes[0, 0].plot(epochs, history['val_loss'], 'r-', label='Validation Loss', linewidth=2)\n    axes[0, 0].set_title('Training vs Validation Loss')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Accuracy comparison\n    axes[0, 1].plot(epochs, history['train_acc'], 'b-', label='Training Accuracy', linewidth=2)\n    axes[0, 1].plot(epochs, history['val_acc'], 'r-', label='Validation Accuracy', linewidth=2)\n    axes[0, 1].set_title('Training vs Validation Accuracy')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Accuracy (%)')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # Learning rate schedule\n    axes[0, 2].plot(epochs, history['learning_rate'], 'g-', linewidth=2)\n    axes[0, 2].set_title('Learning Rate Schedule')\n    axes[0, 2].set_xlabel('Epoch')\n    axes[0, 2].set_ylabel('Learning Rate')\n    axes[0, 2].set_yscale('log')\n    axes[0, 2].grid(True, alpha=0.3)\n    \n    # Loss difference (overfitting indicator)\n    loss_diff = np.array(history['val_loss']) - np.array(history['train_loss'])\n    axes[1, 0].plot(epochs, loss_diff, 'purple', linewidth=2)\n    axes[1, 0].axhline(y=0, color='black', linestyle='--', alpha=0.5)\n    axes[1, 0].set_title('Validation - Training Loss (Overfitting Indicator)')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('Loss Difference')\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # Accuracy difference\n    acc_diff = np.array(history['train_acc']) - np.array(history['val_acc'])\n    axes[1, 1].plot(epochs, acc_diff, 'orange', linewidth=2)\n    axes[1, 1].axhline(y=0, color='black', linestyle='--', alpha=0.5)\n    axes[1, 1].set_title('Training - Validation Accuracy')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('Accuracy Difference (%)')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    # Combined metrics\n    ax = axes[1, 2]\n    ax2 = ax.twinx()\n    \n    line1 = ax.plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2)\n    line2 = ax2.plot(epochs, history['val_acc'], 'b-', label='Val Accuracy', linewidth=2)\n    \n    ax.set_xlabel('Epoch')\n    ax.set_ylabel('Validation Loss', color='red')\n    ax2.set_ylabel('Validation Accuracy (%)', color='blue')\n    ax.set_title('Validation Metrics Combined')\n    \n    # Combine legends\n    lines = line1 + line2\n    labels = [l.get_label() for l in lines]\n    ax.legend(lines, labels, loc='center right')\n    \n    ax.grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n\nprint(\"\\n📊 DETAILED LEARNING CURVES:\")\nplot_detailed_learning_curves(history)\n\n# Performance summary table\ndef create_performance_summary():\n    summary_data = {\n        'Metric': [\n            'Accuracy', 'Precision (Macro)', 'Recall (Macro)', 'F1 Score (Macro)',\n            'AUC-ROC (Macro)', 'Cohen\\'s Kappa', 'Matthews Corr Coef', 'Balanced Accuracy',\n            'Top-3 Accuracy', 'Expected Calibration Error', 'Brier Score',\n            'Avg Inference Time (ms)', 'Throughput (samples/sec)', 'Model Size (MB)',\n            'Training Time (hrs)', 'Parameter Count (M)'\n        ],\n        'Value': [\n            f\"{test_results['accuracy']:.4f}\",\n            f\"{test_results['precision_macro']:.4f}\",\n            f\"{test_results['recall_macro']:.4f}\",\n            f\"{test_results['f1_macro']:.4f}\",\n            f\"{test_results['auc_roc_macro']:.4f}\",\n            f\"{test_results['cohen_kappa']:.4f}\",\n            f\"{test_results['matthews_corrcoef']:.4f}\",\n            f\"{test_results['balanced_accuracy']:.4f}\",\n            f\"{test_results['top_3_accuracy']:.4f}\",\n            f\"{test_results['ece']:.4f}\",\n            f\"{test_results['brier_score']:.4f}\",\n            f\"{avg_inference_time:.2f}\",\n            f\"{throughput:.2f}\",\n            f\"{model_size_mb:.2f}\",\n            f\"{total_training_time/3600:.2f}\",\n            f\"{param_count/1e6:.2f}\"\n        ]\n    }\n    \n    summary_df = pd.DataFrame(summary_data)\n    print(\"\\n📋 PERFORMANCE SUMMARY TABLE:\")\n    print(summary_df.to_string(index=False))\n    \n    return summary_df\n\nsummary_df = create_performance_summary()\n\n# Class-wise performance table\ndef create_classwise_summary():\n    classwise_data = {\n        'Class': class_names,\n        'Precision': [f\"{p:.4f}\" for p in test_results['class_precision']],\n        'Recall': [f\"{r:.4f}\" for r in test_results['class_recall']],\n        'F1-Score': [f\"{f:.4f}\" for f in test_results['class_f1']],\n        'Specificity': [f\"{s:.4f}\" for s in test_results['class_specificity']],\n        'AUC-ROC': [f\"{a:.4f}\" for a in test_results['class_auc_roc']],\n        'AUC-PRC': [f\"{a:.4f}\" for a in test_results['class_auc_prc']],\n        'Avg Precision': [f\"{a:.4f}\" for a in test_results['class_avg_precision']]\n    }\n    \n    classwise_df = pd.DataFrame(classwise_data)\n    print(\"\\n📊 CLASS-WISE PERFORMANCE TABLE:\")\n    print(classwise_df.to_string(index=False))\n    \n    return classwise_df\n\nclasswise_df = create_classwise_summary()\n\n# Save results to files\nprint(\"\\n💾 SAVING RESULTS...\")\n\n# Save performance summary\nsummary_df.to_csv('performance_summary.csv', index=False)\nclasswise_df.to_csv('classwise_performance.csv', index=False)\n\n# Save detailed results as JSON\nimport json\nresults_dict = {\n    'basic_metrics': {\n        'accuracy': float(test_results['accuracy']),\n        'precision_macro': float(test_results['precision_macro']),\n        'recall_macro': float(test_results['recall_macro']),\n        'f1_macro': float(test_results['f1_macro']),\n        'auc_roc_macro': float(test_results['auc_roc_macro']) if not np.isnan(test_results['auc_roc_macro']) else None,\n        'cohen_kappa': float(test_results['cohen_kappa']),\n        'log_loss': float(test_results['log_loss']) if not np.isnan(test_results['log_loss']) else None,\n        'balanced_accuracy': float(test_results['balanced_accuracy']),\n        'matthews_corrcoef': float(test_results['matthews_corrcoef']),\n        'geometric_mean': float(test_results['geometric_mean'])\n    },\n    'calibration_metrics': {\n        'ece': float(test_results['ece']),\n        'brier_score': float(test_results['brier_score'])\n    },\n    'performance_metrics': {\n        'avg_inference_time_ms': float(avg_inference_time),\n        'throughput_samples_per_sec': float(throughput),\n        'parameter_count_M': float(param_count / 1e6),\n        'model_size_MB': float(model_size_mb),\n        'training_time_hours': float(total_training_time / 3600),\n        'inference_memory_MB': float(test_results['inference_memory_mb'])\n    },\n    'multi_average_metrics': {\n        'precision_micro': float(test_results['precision_micro']),\n        'precision_weighted': float(test_results['precision_weighted']),\n        'recall_micro': float(test_results['recall_micro']),\n        'recall_weighted': float(test_results['recall_weighted']),\n        'f1_micro': float(test_results['f1_micro']),\n        'f1_weighted': float(test_results['f1_weighted']),\n        'top_3_accuracy': float(test_results['top_3_accuracy'])\n    },\n    'noise_robustness': {f'noise_{k}': float(v) for k, v in noise_results.items()},\n    'timing': {\n        'start_date': start_date.isoformat(),\n        'end_date': end_date.isoformat(),\n        'total_duration_minutes': float(total_training_time / 60)\n    }\n}\n\nwith open('detailed_results.json', 'w') as f:\n    json.dump(results_dict, f, indent=2)\n\n# Save training history\nhistory_df = pd.DataFrame(history)\nhistory_df.to_csv('training_history.csv', index=False)\n\nprint(\"✅ All results saved successfully!\")\nprint(\"\\nFiles created:\")\nprint(\"- performance_summary.csv\")\nprint(\"- classwise_performance.csv\") \nprint(\"- detailed_results.json\")\nprint(\"- training_history.csv\")\nprint(\"- best_resnet101_model.pth\")\n\nprint(\"\\n🎉 EVALUATION COMPLETE!\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-07-26T15:54:04.173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}