{"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"},{"sourceId":407358,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":332854,"modelId":353781}],"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-05-30T04:29:46.322361Z","iopub.execute_input":"2025-05-30T04:29:46.322547Z","iopub.status.idle":"2025-05-30T04:30:09.790415Z","shell.execute_reply.started":"2025-05-30T04:29:46.32253Z","shell.execute_reply":"2025-05-30T04:30:09.789841Z"}},"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/\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:34:06.925875Z","iopub.execute_input":"2025-05-30T04:34:06.926521Z","iopub.status.idle":"2025-05-30T04:34:06.930089Z","shell.execute_reply.started":"2025-05-30T04:34:06.926495Z","shell.execute_reply":"2025-05-30T04:34:06.929266Z"}},"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-05-30T04:34:08.739644Z","iopub.execute_input":"2025-05-30T04:34:08.739916Z","iopub.status.idle":"2025-05-30T04:34:09.088183Z","shell.execute_reply.started":"2025-05-30T04:34:08.739894Z","shell.execute_reply":"2025-05-30T04:34:09.087425Z"}},"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_subset, 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-05-30T04:34:12.784145Z","iopub.execute_input":"2025-05-30T04:34:12.784433Z","iopub.status.idle":"2025-05-30T04:34:12.870482Z","shell.execute_reply.started":"2025-05-30T04:34:12.784412Z","shell.execute_reply":"2025-05-30T04:34:12.869818Z"}},"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-05-30T04:34:14.500134Z","iopub.execute_input":"2025-05-30T04:34:14.500691Z","iopub.status.idle":"2025-05-30T04:34:14.565705Z","shell.execute_reply.started":"2025-05-30T04:34:14.50067Z","shell.execute_reply":"2025-05-30T04:34:14.565097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ChunkedBrainActivityDataset(Dataset):\n    def __init__(self, csv_file, base_dir, activity_mapping):\n        self.df = pd.read_csv(csv_file)\n        self.base_dir = base_dir\n        self.activity_mapping = activity_mapping\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        spect_id, label, offset = self.df.iloc[idx][[\"spectrogram_id\", \"expert_consensus\", \"spectrogram_label_offset_seconds\"]]\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        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        temp_arr_uint8 = np.uint8(255 * temp_arr)\n        rgb_image = cv2.applyColorMap(temp_arr_uint8, cv2.COLORMAP_JET)\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        # y_tensor = torch.tensor(y, dtype=torch.long)\n        y_tensor = torch.tensor(y, dtype=torch.long)\n\n\n        return rgb_image_tensor, y_tensor\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:34:16.379646Z","iopub.execute_input":"2025-05-30T04:34:16.379939Z","iopub.status.idle":"2025-05-30T04:34:16.387153Z","shell.execute_reply.started":"2025-05-30T04:34:16.379917Z","shell.execute_reply":"2025-05-30T04:34:16.386519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PrototypicalNetworks(nn.Module):\n    def __init__(self, backbone: nn.Module):\n        super(PrototypicalNetworks, self).__init__()\n        self.backbone = backbone\n\n    def forward(\n        self,\n        support_images: torch.Tensor,\n        support_labels: torch.Tensor,\n        query_images: torch.Tensor,\n    ) -> torch.Tensor:\n        # Feature extraction\n        z_support = self.backbone(support_images)\n        z_query = self.backbone(query_images)\n\n        # Prototype computation using index_add\n        class_ids, sy_indices = torch.unique(support_labels, return_inverse=True)\n        prototypes = torch.zeros((len(class_ids), z_support.size(1)), device=z_support.device)\n        prototypes = prototypes.index_add(0, sy_indices, z_support)\n        counts = torch.bincount(sy_indices, minlength=len(class_ids)).unsqueeze(1)\n        prototypes = prototypes / counts\n\n        # Distance calculation\n        dists = torch.cdist(z_query, prototypes)\n        return -dists","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:34:18.424993Z","iopub.execute_input":"2025-05-30T04:34:18.425275Z","iopub.status.idle":"2025-05-30T04:34:18.430884Z","shell.execute_reply.started":"2025-05-30T04:34:18.425255Z","shell.execute_reply":"2025-05-30T04:34:18.430138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# 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-05-30T04:34:29.674829Z","iopub.execute_input":"2025-05-30T04:34:29.675551Z","iopub.status.idle":"2025-05-30T04:34:31.476761Z","shell.execute_reply.started":"2025-05-30T04:34:29.675525Z","shell.execute_reply":"2025-05-30T04:34:31.476204Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = ChunkedBrainActivityDataset(train_csv, BASE_DIR, activity_mapping)\ntest_dataset = ChunkedBrainActivityDataset(test_csv, BASE_DIR, activity_mapping)\nval_dataset = ChunkedBrainActivityDataset(val_csv, BASE_DIR, activity_mapping)\nprint(\"Data Processing done\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:34:34.114614Z","iopub.execute_input":"2025-05-30T04:34:34.115115Z","iopub.status.idle":"2025-05-30T04:34:34.14004Z","shell.execute_reply.started":"2025-05-30T04:34:34.115092Z","shell.execute_reply":"2025-05-30T04:34:34.139454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import defaultdict\nimport random\n\ndef get_image_episode(dataset, n_way=3, k_shot=5, q_queries=5):\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    return (\n        torch.stack(support_images),\n        torch.tensor(support_labels),\n        torch.stack(query_images),\n        torch.tensor(query_labels)\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:34:42.313914Z","iopub.execute_input":"2025-05-30T04:34:42.314725Z","iopub.status.idle":"2025-05-30T04:34:42.323843Z","shell.execute_reply.started":"2025-05-30T04:34:42.314695Z","shell.execute_reply":"2025-05-30T04:34:42.323016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Task sampling configuration\nN_WAY = 3\nN_SHOT = 5\nN_QUERY = 5\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}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:34:44.413712Z","iopub.execute_input":"2025-05-30T04:34:44.414032Z","iopub.status.idle":"2025-05-30T04:36:28.216341Z","shell.execute_reply.started":"2025-05-30T04:34:44.414007Z","shell.execute_reply":"2025-05-30T04:36:28.21553Z"}},"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\npatience_counter = 0\nbest_val_acc = 0.0\n\nfor episode in range(n_episodes):\n    sX, sy, qX, qy = get_image_episode(train_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\n    sX, sy = sX.to(device), sy.to(device)\n    qX, qy = qX.to(device), qy.to(device)\n\n    optimizer.zero_grad()\n    scores = embedding_model(sX, sy, qX)\n    # Map query labels to current episode's class indices\n    class_ids, sy_indices = 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    loss = criterion(scores, qy_indices)\n    loss.backward()\n    optimizer.step()\n\n    preds = torch.argmax(scores, dim=1)\n    acc = (preds == qy_indices).float().mean().item()\n    if episode % 10 == 0:\n        print(f\"[Episode {episode}] Loss: {loss.item():.4f} | Train Acc: {acc:.4f}\")\n\n    # Validation\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            val_scores = embedding_model(val_sX, val_sy, val_qX)\n            val_class_ids, val_sy_indices = 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            val_preds = torch.argmax(val_scores, dim=1)\n            val_acc = (val_preds == val_qy_indices).float().mean().item()\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        embedding_model.train()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T04:36:28.217597Z","iopub.execute_input":"2025-05-30T04:36:28.218054Z","execution_failed":"2025-05-30T08:43:08.505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Final evaluation on test set\nprint(\"\\n=== Final Evaluation ===\")\nembedding_model.load_state_dict(torch.load(\"best_embedding_model.pt\"))\nembedding_model.eval()\n\ntest_accuracies = []\nn_test_episodes = 100  # Standard practice for few-shot evaluation\n\nwith torch.no_grad():\n    for _ in range(n_test_episodes):\n        sX_test, sy_test, qX_test, qy_test = get_image_episode(test_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\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        scores = embedding_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        \n        preds = torch.argmax(scores, dim=1)\n        test_acc = (preds == qy_indices).float().mean().item()\n        test_accuracies.append(test_acc)\n\nfinal_test_acc = sum(test_accuracies)/len(test_accuracies)\nprint(f\"\\nFinal Test Accuracy over {n_test_episodes} episodes: {final_test_acc:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training loop with optimizations\nembedding_model = model  # Use your PrototypicalNetworks model\n\n# 1. Enhanced Optimizer Configuration\noptimizer = torch.optim.Adam(embedding_model.parameters(), lr=3e-5, weight_decay=1e-5)\ncriterion = nn.CrossEntropyLoss()\n\n# 2. Learning Rate Scheduler\nscheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=3, verbose=True)\n\n# 3. Training State Initialization\ncheckpoint_path = \"/kaggle/input/effnet/pytorch/default/1/best_embedding_model.pt\"\noutput_path = \" \"\nstart_episode = 0\nbest_val_acc = 0.0\npatience_counter = 0\nn_episodes = 1000\npatience = 10\naccumulation_steps = 4  # For gradient accumulation\nscaler = torch.cuda.amp.GradScaler()  # For mixed precision\n\n# 4. Checkpoint Loading\nif os.path.exists(checkpoint_path):\n    print(\"Loading existing checkpoint...\")\n    checkpoint = torch.load(checkpoint_path, map_location=device)\n    embedding_model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n    start_episode = checkpoint['episode'] + 1\n    best_val_acc = checkpoint['best_val_acc']\n    patience_counter = checkpoint['patience_counter']\n    print(f\"Resuming from episode {start_episode} | Best val: {best_val_acc:.4f}\")\n\n# 5. Optimized Training Loop\nfor episode in range(start_episode, n_episodes):\n    # --- Training Phase ---\n    embedding_model.train()\n    sX, sy, qX, qy = get_image_episode(train_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\n    sX, sy, qX, qy = sX.to(device), sy.to(device), qX.to(device), qY.to(device)\n\n    with torch.cuda.amp.autocast():  # Mixed precision\n        scores = embedding_model(sX, sy, qX)\n        # Class mapping\n        class_ids, sy_indices = torch.unique(sy, 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        loss = criterion(scores, qy_indices) / accumulation_steps  # Scale loss\n    \n    # Gradient accumulation\n    scaler.scale(loss).backward()\n    \n    if (episode + 1) % accumulation_steps == 0:\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n\n    # --- Progress Monitoring ---\n    if episode % 10 == 0:\n        preds = torch.argmax(scores.detach(), dim=1)\n        acc = (preds == qy_indices).float().mean().item()\n        print(f\"[Episode {episode}] Loss: {loss.item()*accumulation_steps:.4f} | Acc: {acc:.4f}\")\n\n    # --- Validation Phase ---\n    if episode % 50 == 0:\n        embedding_model.eval()\n        val_accs = []\n        with torch.no_grad():\n            for _ in range(5):  # Validate on 5 episodes\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)\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_accs.append((val_preds == val_qy_indices).float().mean().item())\n        \n        mean_val_acc = np.mean(val_accs)\n        scheduler.step(mean_val_acc)  # Update learning rate\n        \n        print(f\"Val Acc: {mean_val_acc:.4f} | LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        # Checkpoint saving\n        if mean_val_acc > best_val_acc:\n            best_val_acc = mean_val_acc\n            torch.save({\n                'episode': episode,\n                'model_state_dict': embedding_model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'best_val_acc': best_val_acc,\n                'patience_counter': patience_counter\n            },  \"best_embedding_model.pt\")\n            patience_counter = 0\n        else:\n            patience_counter += 1\n            if patience_counter >= patience:\n                print(f\"Early stopping at episode {episode}\")\n                break\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Final Evaluation\nprint(\"\\n=== Final Evaluation ===\")\nembedding_model.load_state_dict(torch.load(checkpoint_path)['model_state_dict'])\ntest_accs = []\nwith torch.no_grad():\n    for _ in range(100):  # Standard 100 test episodes\n        sX, sy, qX, qy = get_image_episode(test_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\n        sX, sy = sX.to(device), sy.to(device)\n        qX, qy = qX.to(device), qy.to(device)\n        \n        scores = embedding_model(sX, sy, qX)\n        class_ids = torch.unique(sy)\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        preds = torch.argmax(scores, dim=1)\n        test_accs.append((preds == qy_indices).float().mean().item())\n\nprint(f\"Test Accuracy: {np.mean(test_accs):.4f} ± {np.std(test_accs):.4f}\")\n\n\nwith torch.no_grad():\n    for _ in range(n_test_episodes):\n        sX_test, sy_test, qX_test, qy_test = get_image_episode(test_dataset, n_way=N_WAY, k_shot=N_SHOT, q_queries=N_QUERY)\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        scores = embedding_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        \n        preds = torch.argmax(scores, dim=1)\n        test_acc = (preds == qy_indices).float().mean().item()\n        test_accuracies.append(test_acc)\n\nfinal_test_acc = sum(test_accuracies)/len(test_accuracies)\nprint(f\"\\nFinal Test Accuracy over {n_test_episodes} episodes: {final_test_acc:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}