{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport warnings\nimport os\nimport time\nimport torch\nimport torch.nn as nn\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.utils import to_categorical\nfrom tensorflow.keras.callbacks import ReduceLROnPlateau, EarlyStopping, TensorBoard\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (accuracy_score, classification_report, confusion_matrix, roc_curve, auc, cohen_kappa_score)\nfrom sklearn.preprocessing import StandardScaler  \nwarnings.filterwarnings('ignore')\n\nimport cv2\nimport random\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nfrom collections import defaultdict\nfrom tqdm import tqdm\n# from easyfsl.samplers import TaskSampler","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-05T07:09:54.483960Z","iopub.execute_input":"2025-06-05T07:09:54.484359Z","iopub.status.idle":"2025-06-05T07:10:18.990613Z","shell.execute_reply.started":"2025-06-05T07:09:54.484336Z","shell.execute_reply":"2025-06-05T07:10:18.990012Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Directories\n# MY_DIR = \"/home/rs/20CS91P02/projects/22CS10066_deepak/\"\nBASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\nPREPROCESSED_DIR = \"/kaggle/working/preprocessed/\"\nos.makedirs(PREPROCESSED_DIR, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:51:40.980949Z","iopub.execute_input":"2025-06-04T13:51:40.981667Z","iopub.status.idle":"2025-06-04T13:51:40.985291Z","shell.execute_reply.started":"2025-06-04T13:51:40.981642Z","shell.execute_reply":"2025-06-04T13:51:40.984574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df= pd.read_csv(f\"{BASE_DIR}train.csv\")\ndf.head()\n\ndf_org = pd.read_csv(f\"{BASE_DIR}train.csv\")\n# Print the total number of rows in the dataset\nprint(f\"Total rows in the dataset: {len(df)}\")\n\n#Randomly select 10,000 rows for a quick training check\ndf_subset = df_org.sample(n=14286, random_state=42)\n# df_subset = df_org.sample(n=1000, random_state=42)\nprint(f\"Total rows in the dataset: {len(df_subset)}\")\n# # Display the first few rows of the sampled dataframe\ndf_subset.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:51:42.823604Z","iopub.execute_input":"2025-06-04T13:51:42.824267Z","iopub.status.idle":"2025-06-04T13:51:43.210832Z","shell.execute_reply.started":"2025-06-04T13:51:42.824242Z","shell.execute_reply":"2025-06-04T13:51:43.210148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#code to split df in train,test val 70, 15,15 \ntrain_df, temp_df = train_test_split(df, test_size=0.30, random_state=42)\ntest_df, val_df = train_test_split(temp_df, test_size=0.50, random_state=42)\n\nprint(f\"Training set size: {len(train_df)}\")\nprint(f\"Validation set size: {len(val_df)}\")\nprint(f\"Test set size: {len(test_df)}\")\n\n# Save the datasets\ntrain_csv = \"/kaggle/working/train10K_70.csv\"\nval_csv = \"/kaggle/working/val10K_15.csv\"\ntest_csv = \"/kaggle/working/test10K_15.csv\"\n\ntrain_df.to_csv(train_csv, index=False)\nval_df.to_csv(val_csv, index=False)\ntest_df.to_csv(test_csv, index=False)\n\nprint(f\"Train CSV saved to: {train_csv}\")\nprint(f\"Validation CSV saved to: {val_csv}\")\nprint(f\"Test CSV saved to: {test_csv}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:51:56.078709Z","iopub.execute_input":"2025-06-04T13:51:56.078949Z","iopub.status.idle":"2025-06-04T13:51:56.175223Z","shell.execute_reply.started":"2025-06-04T13:51:56.078934Z","shell.execute_reply":"2025-06-04T13:51:56.174204Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract EEGid, labels, and offsets\n# EEGid_label_list = df_subset[[\"eeg_id\", \"expert_consensus\", \"eeg_label_offset_seconds\"]].values.tolist()\n\n# X = []\n# y = []\n# prev_eegId = \"\"\n\nbrain_activities = ['Seizure', 'GPD', 'LRDA', 'Other', 'GRDA', 'LPD']\nactivity_mapping = {activity: idx for idx, activity in enumerate(brain_activities)}\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed(42)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:51:59.328930Z","iopub.execute_input":"2025-06-04T13:51:59.329518Z","iopub.status.idle":"2025-06-04T13:51:59.397416Z","shell.execute_reply.started":"2025-06-04T13:51:59.329496Z","shell.execute_reply":"2025-06-04T13:51:59.396722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class OptimizedBrainDataset(Dataset):\n    def __init__(self, csv_file, base_dir, activity_mapping, preprocessed_dir=\"preprocessed\"):\n        self.df = pd.read_csv(csv_file)\n        self.base_dir = base_dir\n        self.activity_mapping = activity_mapping\n        self.preprocessed_dir = preprocessed_dir\n        self.resize_transform = transforms.Resize((224, 224))\n        \n        # Create directory if needed\n        os.makedirs(self.preprocessed_dir, exist_ok=True)\n        \n        # Memory map preprocessed files\n        self.spect_mmaps = {}\n        spect_ids = self.df[\"spectrogram_id\"].unique()\n        for spect_id in spect_ids:\n            npy_path = f\"{self.preprocessed_dir}/{spect_id}.npy\"\n            if not os.path.exists(npy_path):\n                self._preprocess_and_save(spect_id)\n            self.spect_mmaps[spect_id] = np.load(npy_path, mmap_mode='r')\n\n    def __len__(self):\n        return len(self.df)\n\n    def _preprocess_and_save(self, spect_id):\n        \"\"\"Batch process and save spectrogram once\"\"\"\n        parquet_path = f'{self.base_dir}/train_spectrograms/{spect_id}.parquet'\n        temp_df = pd.read_parquet(parquet_path).drop('time', axis=1)\n        \n        # Process entire spectrogram\n        arr = temp_df.to_numpy()\n        arr = np.log1p(arr)\n        arr /= arr.max()\n        arr = np.nan_to_num(arr, nan=1e-4)\n        arr_uint8 = (255 * arr).astype(np.uint8)\n        \n        np.save(f\"{self.preprocessed_dir}/{spect_id}.npy\", arr_uint8)\n\n    def __getitem__(self, idx):\n        spect_id, label, offset = self.df.iloc[idx][[\"spectrogram_id\", \"expert_consensus\", \"spectrogram_label_offset_seconds\"]]\n        start = int(offset) // 2\n        \n        # Direct memory access\n        spectrogram = self.spect_mmaps[spect_id]\n        segment = spectrogram[start:start+300]\n        \n        # Convert to RGB\n        rgb_image = cv2.applyColorMap(segment, cv2.COLORMAP_JET)\n        rgb_image = rgb_image.astype(np.float32) / 255.0\n        tensor_image = torch.tensor(rgb_image).permute(2, 0, 1)\n        return self.resize_transform(tensor_image), torch.tensor(self.activity_mapping[label])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:52:00.888220Z","iopub.execute_input":"2025-06-04T13:52:00.888751Z","iopub.status.idle":"2025-06-04T13:52:00.896878Z","shell.execute_reply.started":"2025-06-04T13:52:00.888731Z","shell.execute_reply":"2025-06-04T13:52:00.895929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EpisodeGenerator:\n    def __init__(self, dataset):\n        self.dataset = dataset\n        self.class_to_indices = defaultdict(list)\n        \n        # Precompute once\n        for idx, (_, label) in enumerate(dataset):\n            self.class_to_indices[label.item()].append(idx)\n\n        # Ensure all classes have sufficient samples\n        min_samples = min(len(indices) for indices in self.class_to_indices.values())\n        # print(f\"Minimum samples per class: {min_samples}\")\n            \n    def get_episode(self, n_way=6, k_shot=3, q_queries=3):\n        all_classes = list(self.class_to_indices.keys())\n        if len(all_classes) < n_way:\n            raise ValueError(f\"Dataset only has {len(all_classes)} classes, need {n_way}\")\n        # selected_classes = random.sample(list(self.class_to_indices.keys()), n_way)\n        selected_classes = all_classes[:n_way]\n        support, query = [], []\n        \n        for cls in selected_classes:\n            available_indices = self.class_to_indices[cls]\n            if len(available_indices) < (k_shot + q_queries):\n                # Handle insufficient samples\n                indices = np.random.choice(available_indices, \n                                         size=k_shot+q_queries, replace=True)\n            else:\n                indices = np.random.choice(available_indices, \n                                         size=k_shot+q_queries, replace=False)\n            # indices = np.random.choice(self.class_to_indices[cls], \n            #                          size=k_shot+q_queries, replace=False)\n            support.extend(indices[:k_shot])\n            query.extend(indices[k_shot:])\n        \n        # Batch loading\n        sX = torch.stack([self.dataset[i][0] for i in support])\n        qX = torch.stack([self.dataset[i][0] for i in query])\n        sy = torch.tensor([self.dataset[i][1] for i in support])\n        qy = torch.tensor([self.dataset[i][1] for i in query])\n        \n        return sX, sy, qX, qy\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:52:05.445431Z","iopub.execute_input":"2025-06-04T13:52:05.445723Z","iopub.status.idle":"2025-06-04T13:52:05.453306Z","shell.execute_reply.started":"2025-06-04T13:52:05.445701Z","shell.execute_reply":"2025-06-04T13:52:05.452493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pregenerate_episodes(dataset, cache_size=1000):\n    import gc\n    generator = EpisodeGenerator(dataset)\n    episode_cache = []\n    \n    # Warmup GPU\n    _ = torch.rand(1).to(device)\n    \n    print(f\"=> Pre-generating {cache_size} episodes\")\n    for i in tqdm(range(cache_size)):\n        sX, sy, qX, qy = generator.get_episode()\n        \n        # Async GPU transfer with pinned memory\n        sX = sX.pin_memory().to(device, non_blocking=True)\n        sy = sy.pin_memory().to(device, non_blocking=True)\n        qX = qX.pin_memory().to(device, non_blocking=True)\n        qy = qy.pin_memory().to(device, non_blocking=True)\n        \n        episode_cache.append((sX, sy, qX, qy))\n\n        # Clear CUDA cache every 25 episodes to prevent OOM\n        if i % 25 == 0 and i > 0:\n            torch.cuda.empty_cache()\n            \n        # Force garbage collection every 50 episodes\n        if i % 50 == 0 and i > 0:\n            gc.collect()\n            torch.cuda.empty_cache()\n            \n        # Print memory usage every 100 episodes\n        if i % 100 == 0 and i > 0:\n            allocated = torch.cuda.memory_allocated() / 1e9\n            reserved = torch.cuda.memory_reserved() / 1e9\n            print(f\"GPU Memory: {allocated:.2f}GB allocated, {reserved:.2f}GB reserved\")\n    \n    # Final cleanup\n    gc.collect()\n    torch.cuda.empty_cache()\n    \n    return episode_cache\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:52:08.324596Z","iopub.execute_input":"2025-06-04T13:52:08.324870Z","iopub.status.idle":"2025-06-04T13:52:08.331553Z","shell.execute_reply.started":"2025-06-04T13:52:08.324851Z","shell.execute_reply":"2025-06-04T13:52:08.330922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PrototypicalNetworks(nn.Module):\n    def __init__(self, backbone: nn.Module):\n        super().__init__()  # Required for PyTorch modules\n        self.backbone = backbone  # Essential for parameter storage\n\n    def forward(self, sX, sy, qX):\n        z_support = self.backbone(sX)\n        z_query = self.backbone(qX)\n        \n        prototypes = torch.stack([\n            z_support[sy == cls].mean(0) \n            for cls in torch.unique(sy)\n        ])\n        \n        return -torch.cdist(z_query, prototypes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:52:10.269121Z","iopub.execute_input":"2025-06-04T13:52:10.269659Z","iopub.status.idle":"2025-06-04T13:52:10.274289Z","shell.execute_reply.started":"2025-06-04T13:52:10.269635Z","shell.execute_reply":"2025-06-04T13:52:10.273451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run once before training\ntrain_dataset = OptimizedBrainDataset(train_csv, BASE_DIR, activity_mapping)\ntest_dataset =  OptimizedBrainDataset(test_csv, BASE_DIR, activity_mapping)\nval_dataset =  OptimizedBrainDataset(val_csv, BASE_DIR, activity_mapping)\nprint(\"Data Processing done\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:52:58.776471Z","iopub.execute_input":"2025-06-03T11:52:58.776973Z","iopub.status.idle":"2025-06-03T11:57:37.145310Z","shell.execute_reply.started":"2025-06-03T11:52:58.776950Z","shell.execute_reply":"2025-06-03T11:57:37.144534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize backbone\n# efficientnet = models.efficientnet_b3(pretrained=True)\nefficientnet = models.efficientnet_v2_s(pretrained=True)\nefficientnet.classifier = nn.Identity()  # Keep only feature extractor\nmodel = PrototypicalNetworks(efficientnet).cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T13:52:18.656570Z","iopub.execute_input":"2025-06-04T13:52:18.656843Z","iopub.status.idle":"2025-06-04T13:52:19.854947Z","shell.execute_reply.started":"2025-06-04T13:52:18.656822Z","shell.execute_reply":"2025-06-04T13:52:19.854181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from collections import defaultdict\n# import random\n\ndef get_image_episode(dataset, n_way=6, k_shot=3, q_queries=3, return_original=False):\n    class_to_indices = defaultdict(list)\n    for idx, (_, label) in enumerate(dataset):\n        class_idx = label.item() if torch.is_tensor(label) else label\n        class_to_indices[class_idx].append(idx)\n\n    selected_classes = random.sample(list(class_to_indices.keys()), n_way)\n    support_images, query_images = [], []\n    support_labels, query_labels = [], []\n    label_map = {cls: i for i, cls in enumerate(selected_classes)}\n\n    for cls in selected_classes:\n        indices = class_to_indices[cls]\n        selected = random.sample(indices, k_shot + q_queries)\n        s_idx, q_idx = selected[:k_shot], selected[k_shot:]\n\n        for idx in s_idx:\n            img, _ = dataset[idx]\n            support_images.append(img)\n            support_labels.append(label_map[cls])\n        for idx in q_idx:\n            img, _ = dataset[idx]\n            query_images.append(img)\n            query_labels.append(label_map[cls])\n\n    if return_original:\n        return (torch.stack(support_images),\n                torch.tensor(support_labels),\n                torch.stack(query_images),\n                torch.tensor(query_labels),\n                selected_classes)  # Return original class IDs\n    else:\n        return (torch.stack(support_images),\n                torch.tensor(support_labels),\n                torch.stack(query_images),\n                torch.tensor(query_labels))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:57:38.356837Z","iopub.execute_input":"2025-06-03T11:57:38.357464Z","iopub.status.idle":"2025-06-03T11:57:38.364379Z","shell.execute_reply.started":"2025-06-03T11:57:38.357436Z","shell.execute_reply":"2025-06-03T11:57:38.363750Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Task sampling configuration\nN_WAY = 6\nN_SHOT = 10\nN_QUERY = 10\n\n# Initial evaluation before training\nmodel.eval()\nsX_test, sy_test, qX_test, qy_test = get_image_episode(test_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\nsX_test, sy_test = sX_test.to(device), sy_test.to(device)\nqX_test, qy_test = qX_test.to(device), qy_test.to(device)\n\nwith torch.no_grad():\n    scores = model(sX_test, sy_test, qX_test)\n    class_ids, _ = torch.unique(sy_test, sorted=True, return_inverse=True)\n    class_to_index = {cls.item(): idx for idx, cls in enumerate(class_ids)}\n    qy_indices = torch.tensor([class_to_index[c.item()] for c in qy_test], device=device)\n    preds = torch.argmax(scores, dim=1)\n    initial_acc = (preds == qy_indices).float().mean().item()\n\nprint(f\"Initial Test Accuracy: {initial_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:58:15.716642Z","iopub.execute_input":"2025-06-03T11:58:15.716897Z","iopub.status.idle":"2025-06-03T11:58:24.792653Z","shell.execute_reply.started":"2025-06-03T11:58:15.716880Z","shell.execute_reply":"2025-06-03T11:58:24.791861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training loop\nembedding_model = model  # Use your PrototypicalNetworks model\noptimizer = torch.optim.Adam(embedding_model.parameters(), lr=1e-4)\ncriterion = nn.CrossEntropyLoss()\n\nn_episodes = 1000  # or whatever you prefer\npatience = 10\n# patience_counter = 0\n# best_val_acc = 0.0\n\n# Pre-generate episode cache\n# EPISODE_CACHE_SIZE = 1000  # Adjust based on memory constraints\nEPISODE_CACHE_SIZE = min(300, len(train_dataset)//(N_WAY*(N_SHOT+N_QUERY)))\nprint(f\"episode cache size: {EPISODE_CACHE_SIZE}\")\nepisode_cache = []\n\nprint(\"==> Pre-generating training episodes...\")\nepisode_cache = pregenerate_episodes(train_dataset, EPISODE_CACHE_SIZE)\nprint(\"Done\")\n\n# Modified training loop with cached episodes\nbest_val_acc = 0.0\npatience_counter = 0\n\nfor episode in range(n_episodes):\n    # Get pre-generated episode\n    sX, sy, qX, qy = episode_cache[episode % EPISODE_CACHE_SIZE]\n    \n    optimizer.zero_grad()\n    \n    # Forward pass\n    scores = embedding_model(sX, sy, qX)\n    \n    # Create episode-specific class mapping\n    class_ids, _ = torch.unique(sy, sorted=True, return_inverse=True)\n    class_to_index = {cls.item(): idx for idx, cls in enumerate(class_ids)}\n    qy_indices = torch.tensor([class_to_index[c.item()] for c in qy], device=device)\n    \n    # Calculate loss\n    loss = criterion(scores, qy_indices)\n    loss.backward()\n    optimizer.step()\n\n    # Training metrics\n    preds = torch.argmax(scores, dim=1)\n    acc = (preds == qy_indices).float().mean().item()\n    \n    if episode % 10 == 0:\n        print(f\"[Episode {episode}] Loss: {loss.item():.4f} | Train Acc: {acc:.4f}\")\n\n    # Memory management during training\n    if episode % 50 == 0 and episode > 0:\n        torch.cuda.empty_cache()\n\n    # Validation (unchanged)\n    if episode % 50 == 0:\n        embedding_model.eval()\n        with torch.no_grad():\n            val_sX, val_sy, val_qX, val_qy = get_image_episode(val_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\n            val_sX, val_sy = val_sX.to(device), val_sy.to(device)\n            val_qX, val_qy = val_qX.to(device), val_qy.to(device)\n            \n            val_scores = embedding_model(val_sX, val_sy, val_qX)\n            val_class_ids, _ = torch.unique(val_sy, sorted=True, return_inverse=True)\n            val_class_to_index = {cls.item(): idx for idx, cls in enumerate(val_class_ids)}\n            val_qy_indices = torch.tensor([val_class_to_index[c.item()] for c in val_qy], device=device)\n            \n            val_preds = torch.argmax(val_scores, dim=1)\n            val_acc = (val_preds == val_qy_indices).float().mean().item()\n            \n            print(f\"--> [Validation] Episode {episode} | Val Acc: {val_acc:.4f}\")\n            if val_acc > best_val_acc:\n                best_val_acc = val_acc\n                torch.save(embedding_model.state_dict(), \"best_embedding_model.pt\")\n                patience_counter = 0\n            else:\n                patience_counter += 1\n                print(f\" No improvement. Patience: {patience_counter}/{patience}\")\n                if patience_counter >= patience:\n                    print(\" Early stopping triggered.\")\n                    break\n            # Clear validation tensors to free memory\n            del val_sX, val_sy, val_qX, val_qy, val_scores\n            torch.cuda.empty_cache()\n            \n        embedding_model.train()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:58:28.106680Z","iopub.execute_input":"2025-06-03T11:58:28.107217Z","iopub.status.idle":"2025-06-03T12:07:12.307135Z","shell.execute_reply.started":"2025-06-03T11:58:28.107171Z","shell.execute_reply":"2025-06-03T12:07:12.306358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Classification metrics\nfrom sklearn.metrics import classification_report, confusion_matrix, accuracy_score\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# ... [your existing evaluation code] ...\n# Final evaluation on test set with comprehensive metrics\nprint(\"\\n=== Final Evaluation ===\")\nembedding_model.load_state_dict(torch.load(\"best_embedding_model.pt\"))\nembedding_model.eval()\n\ntest_accuracies = []\nall_true = []\nall_preds = []\nn_test_episodes = 100  # Standard practice for few-shot evaluation\n\nwith torch.no_grad():\n    for episode_idx in tqdm(range(n_test_episodes), desc=\"Test Episodes\"):\n        # Generate episode ensuring valid class distribution\n        while True:\n            try:\n                sX_test, sy_test, qX_test, qy_test = get_image_episode(\n                    test_dataset, \n                    n_way=N_WAY, \n                    k_shot=N_SHOT, \n                    q_queries=N_QUERY\n                )\n                # Validate episode structure\n                assert len(torch.unique(sy_test)) == N_WAY\n                break\n            except (ValueError, AssertionError):\n                continue\n\n        # Device transfer\n        sX_test, sy_test = sX_test.to(device), sy_test.to(device)\n        qX_test, qy_test = qX_test.to(device), qy_test.to(device)\n\n        # Forward pass\n        scores = embedding_model(sX_test, sy_test, qX_test)\n\n        # Create episode-specific mapping\n        class_ids, _ = torch.unique(sy_test, sorted=True, return_inverse=True)\n        class_to_index = {cls.item(): idx for idx, cls in enumerate(class_ids)}\n        \n        # Convert query labels using episode mapping\n        qy_indices = torch.tensor([class_to_index[c.item()] for c in qy_test], device=device)\n        \n        # Calculate metrics\n        preds = torch.argmax(scores, dim=1)\n        \n        # Store original class labels for comprehensive metrics\n        true_labels = [class_ids[idx].item() for idx in qy_indices]\n        pred_labels = [class_ids[idx].item() for idx in preds]\n        \n        all_true.extend(true_labels)\n        all_preds.extend(pred_labels)\n\n        # Calculate episode accuracy\n        episode_acc = (preds == qy_indices).float().mean().item()\n        test_accuracies.append(episode_acc)\n\n# Calculate final statistics\nmean_acc = np.mean(test_accuracies)\nstd_acc = np.std(test_accuracies)\nconfidence_interval = 1.96 * std_acc / np.sqrt(n_test_episodes)\n\n# Classification metrics\nfrom sklearn.metrics import classification_report, confusion_matrix\nimport pandas as pd\n\nprint(\"\\n=== Comprehensive Metrics ===\")\nprint(f\"Final Test Accuracy over {n_test_episodes} episodes:\")\nprint(f\"Mean Accuracy: {mean_acc:.4f} ± {confidence_interval:.4f} (95% CI)\")\nprint(f\"Standard Deviation: {std_acc:.4f}\")\n\n# Generate classification report\nreport = classification_report(\n    all_true, \n    all_preds, \n    labels=list(range(len(brain_activities))),\n    target_names=brain_activities,\n    zero_division=0,\n    output_dict=True\n)\n\n# Calculate overall accuracy separately\noverall_accuracy = accuracy_score(all_true, all_preds)\n\n# Macro-averaged metrics\nprint(\"\\nMacro-Averaged Scores:\")\nprint(f\"Precision: {report['macro avg']['precision']:.4f}\")\nprint(f\"Recall: {report['macro avg']['recall']:.4f}\")\nprint(f\"F1-Score: {report['macro avg']['f1-score']:.4f}\")\n\n# Confusion matrix with heatmap visualization\ncm = confusion_matrix(all_true, all_preds, labels=list(range(len(brain_activities))))\n\n# Create heatmap\nplt.figure(figsize=(10, 8))\nsns.set(font_scale=1.2)\n\n# Create DataFrame for better labeling\ncm_df = pd.DataFrame(\n    cm,\n    index=[f\"Actual {c}\" for c in brain_activities],\n    columns=[f\"Predicted {c}\" for c in brain_activities]\n)\n\n# Plot heatmap\nsns.heatmap(cm_df, \n           annot=True,           # Show numbers in cells\n           fmt='d',              # Format as integers\n           cmap='Blues',         # Color scheme\n           square=True,          # Square cells\n           linewidths=0.5,       # Add grid lines\n           cbar_kws={'label': 'Number of Samples'})\n\nplt.title('Confusion Matrix - Brain Activity Classification', fontsize=16, pad=20)\nplt.xlabel('Predicted Labels', fontsize=14)\nplt.ylabel('Actual Labels', fontsize=14)\nplt.savefig('confusion_matrix.png', dpi=300, bbox_inches='tight')\nplt.tight_layout()\nplt.show()\n\n# Also print the numerical matrix for reference\nprint(\"\\nConfusion Matrix (Numerical):\")\nprint(cm_df)\n\n# Per-class metrics\nprint(\"\\nDetailed Class Performance:\")\nfor class_name in brain_activities:\n    if class_name in report:\n        print(f\"\\n{class_name}:\")\n        print(f\"  Precision: {report[class_name]['precision']:.4f}\")\n        print(f\"  Recall:    {report[class_name]['recall']:.4f}\")\n        print(f\"  F1-Score:  {report[class_name]['f1-score']:.4f}\")\n        print(f\"  Support:   {report[class_name]['support']}\")\n\n# Additional metrics\nprint(\"\\nAdditional Metrics:\")\nprint(f\"Weighted Avg F1: {report['weighted avg']['f1-score']:.4f}\")\nprint(f\"Overall Accuracy: {overall_accuracy:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T12:10:25.427813Z","iopub.execute_input":"2025-06-03T12:10:25.428104Z","iopub.status.idle":"2025-06-03T12:23:32.749307Z","shell.execute_reply.started":"2025-06-03T12:10:25.428084Z","shell.execute_reply":"2025-06-03T12:23:32.748604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}