{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11417007,"sourceType":"datasetVersion","datasetId":7150373}],"dockerImageVersionId":31011,"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\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport warnings\nimport time\nimport cv2\nwarnings.filterwarnings('ignore')\nfrom sklearn.model_selection import train_test_split","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:35.145485Z","iopub.execute_input":"2025-04-15T10:55:35.146103Z","iopub.status.idle":"2025-04-15T10:55:35.150427Z","shell.execute_reply.started":"2025-04-15T10:55:35.146079Z","shell.execute_reply":"2025-04-15T10:55:35.149750Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\n\nstart_time = time.time()\nbrain_activities = ['Seizure', 'GPD', 'LRDA', 'Other', 'GRDA', 'LPD']\nactivity_mapping = {activity: idx for idx, activity in enumerate(brain_activities)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:36.671274Z","iopub.execute_input":"2025-04-15T10:55:36.671965Z","iopub.status.idle":"2025-04-15T10:55:36.675547Z","shell.execute_reply.started":"2025-04-15T10:55:36.671942Z","shell.execute_reply":"2025-04-15T10:55:36.674733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EarlyStopping:\n    def __init__(self, patience=3):\n        \"\"\"\n        Stops training if validation accuracy doesn't improve after 'patience' epochs.\n        \"\"\"\n        self.patience = patience\n        self.best_acc = 0\n        self.counter = 0\n\n    def __call__(self, val_acc):\n        if val_acc > self.best_acc:\n            self.best_acc = val_acc\n            self.counter = 0 # Reset counter if accuracy improves\n        else:\n            self.counter += 1\n            if self.counter >= self.patience:\n                print(f\"Early stopping triggered! Best Val Accuracy: {self.best_acc:.4f}\")\n                return True # Stop training\n        return False # Continue training\n\nclass 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\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:37.876612Z","iopub.execute_input":"2025-04-15T10:55:37.877301Z","iopub.status.idle":"2025-04-15T10:55:37.885103Z","shell.execute_reply.started":"2025-04-15T10:55:37.877280Z","shell.execute_reply":"2025-04-15T10:55:37.884377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# df= pd.read_csv(f\"{BASE_DIR}train.csv\")\n# df.head()\n\n# df_org = pd.read_csv(f\"{BASE_DIR}train.csv\")\n# # Print the total number of rows in the dataset\n# print(f\"Total rows in the dataset: {len(df)}\")\n\n# Randomly select 10,000 rows for a quick training check\n# df_subset = df_org.sample(n=14286, random_state=42)\n# df_subset = df_org.sample(n=1000, random_state=42)\n\n# # # Display the first few rows of the sampled dataframe\n# df_subset.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:38.463589Z","iopub.execute_input":"2025-04-15T10:55:38.464283Z","iopub.status.idle":"2025-04-15T10:55:38.467260Z","shell.execute_reply.started":"2025-04-15T10:55:38.464257Z","shell.execute_reply":"2025-04-15T10:55:38.466702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# #code to split df in train,test val 70, 15,15 \n# train_df, temp_df = train_test_split(df_subset, test_size=0.30, random_state=42)\n# test_df, val_df = train_test_split(temp_df, test_size=0.50, random_state=42)\n\n# print(f\"Training set size: {len(train_df)}\")\n# print(f\"Validation set size: {len(val_df)}\")\n# print(f\"Test set size: {len(test_df)}\")\n\n# # Save the datasets\n# train_csv = \"/kaggle/working/train10K_70.csv\"\n# val_csv = \"/kaggle/working/val10K_15.csv\"\n# test_csv = \"/kaggle/working/test10K_15.csv\"\n\n# train_df.to_csv(train_csv, index=False)\n# val_df.to_csv(val_csv, index=False)\n# test_df.to_csv(test_csv, index=False)\n\n# print(f\"Train CSV saved to: {train_csv}\")\n# print(f\"Validation CSV saved to: {val_csv}\")\n# print(f\"Test CSV saved to: {test_csv}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:39.031095Z","iopub.execute_input":"2025-04-15T10:55:39.031333Z","iopub.status.idle":"2025-04-15T10:55:39.034831Z","shell.execute_reply.started":"2025-04-15T10:55:39.031314Z","shell.execute_reply":"2025-04-15T10:55:39.034259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_csv = \"/kaggle/input/hms-data/train_80.csv\"\nval_csv = \"/kaggle/input/hms-data/val_10.csv\"\ntest_csv = \"/kaggle/input/hms-data/test_10.csv\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:41.714187Z","iopub.execute_input":"2025-04-15T10:55:41.714451Z","iopub.status.idle":"2025-04-15T10:55:41.718312Z","shell.execute_reply.started":"2025-04-15T10:55:41.714430Z","shell.execute_reply":"2025-04-15T10:55:41.717634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load CSVs\ntrain_df = pd.read_csv(train_csv)\nval_df = pd.read_csv(val_csv)\ntest_df = pd.read_csv(test_csv)\n\n# Print number of rows\nprint(f\"Number of rows in train set: {len(train_df)}\")\nprint(f\"Number of rows in validation set: {len(val_df)}\")\nprint(f\"Number of rows in test set: {len(test_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:43.420837Z","iopub.execute_input":"2025-04-15T10:55:43.421155Z","iopub.status.idle":"2025-04-15T10:55:43.683880Z","shell.execute_reply.started":"2025-04-15T10:55:43.421132Z","shell.execute_reply":"2025-04-15T10:55:43.683296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#call  ChunkedBrainActivityDataset() vy train,val,test\ntrain_dataset = ChunkedBrainActivityDataset(csv_file=train_csv, base_dir=BASE_DIR, activity_mapping=activity_mapping)\nval_dataset = ChunkedBrainActivityDataset(csv_file=val_csv, base_dir=BASE_DIR, activity_mapping=activity_mapping)\ntest_dataset = ChunkedBrainActivityDataset(csv_file=test_csv, base_dir=BASE_DIR, activity_mapping=activity_mapping)\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=12, pin_memory=True, prefetch_factor=2)\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=False, num_workers=12, pin_memory=True, prefetch_factor=2)\ntest_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers=12, pin_memory=True, prefetch_factor=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:45.407073Z","iopub.execute_input":"2025-04-15T10:55:45.407362Z","iopub.status.idle":"2025-04-15T10:55:45.559425Z","shell.execute_reply.started":"2025-04-15T10:55:45.407343Z","shell.execute_reply":"2025-04-15T10:55:45.558885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = models.efficientnet_v2_s(pretrained=True)\nnum_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(num_features, 6)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\nmodel.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:55:54.769903Z","iopub.execute_input":"2025-04-15T10:55:54.770415Z","iopub.status.idle":"2025-04-15T10:55:56.178145Z","shell.execute_reply.started":"2025-04-15T10:55:54.770393Z","shell.execute_reply":"2025-04-15T10:55:56.177382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 101\npatience = 15\nbest_val_loss = float('inf')\npatience_counter = 0\n\ntrain_losses, val_losses, train_accuracies, val_accuracies = [], [], [], []\nearly_stopping = EarlyStopping(patience=3)\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    # tr_cntr=0\n    for data in tqdm(train_loader):\n        inputs, labels = data\n        inputs, labels = inputs.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        _, labels_ = torch.max(labels.data, 1)\n        correct += (predicted == labels_).sum().item()\n                    \n        # batch_accuracy = 100 * correct / total\n        # print(f\"Train Batch: {tr_cntr}/585, Loss: {loss.item():.4f}, Accuracy: {batch_accuracy:.2f}%\")\n        # tr_cntr+=1\n        \n        torch.cuda.empty_cache()\n\n    train_loss = running_loss / len(train_loader)\n    train_accuracy = 100 * correct / total\n    train_losses.append(train_loss)\n    train_accuracies.append(train_accuracy)\n\n    model.eval()\n    val_loss = 0.0\n    correct = 0\n    total = 0\n    # val_cntr=0\n    with torch.no_grad():\n        for data in tqdm(val_loader):\n            inputs, labels = data\n            inputs, labels = inputs.to(device), labels.to(device)\n            # print(\"Data for some epoch loaded!\")\n\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            val_loss += loss.item()\n\n            _, predicted = torch.max(outputs.data, 1)\n            total += labels.size(0)\n            _, labels_ = torch.max(labels.data, 1)\n            correct += (predicted == labels_).sum().item()\n            # batch_accuracy = 100 * correct / total\n            # print(f\"Val Batch: {val_cntr}/126, Loss: {loss.item():.4f}, Accuracy: {batch_accuracy:.2f}%\")\n            # val_cntr+=1\n            torch.cuda.empty_cache()\n\n    val_loss /= len(val_loader)\n    val_accuracy = 100 * correct / total\n    val_losses.append(val_loss)\n    val_accuracies.append(val_accuracy)\n\n    elapsed_time = time.time() - start_time\n    print(f\"Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.2f}%, Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.2f}%, , Time: {elapsed_time}\")\n    scheduler.step(val_loss)\n    \n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        patience_counter = 0\n        torch.save(model.state_dict(), 'HMS_model_v1_resnet50.pth')\n    else:\n        patience_counter += 1\n\n    if patience_counter >= patience:\n        print(\"Early stopping triggered\")\n        break\n\n    if early_stopping(val_accuracy): # Stop if no improvement  \n        print(\"Early stopping triggered\")\n        break\n    \n\n    torch.cuda.empty_cache()\n\n# Testing\nif(os.path.isfile('HMS_model_v1_resnet50.pth')):\n    model.load_state_dict(torch.load('HMS_model_v1_resnet50.pth'))\nmodel.eval()\n\ntest_loss = 0.0\ncorrect = 0\ntotal = 0\n# test_cntr=0\nwith torch.no_grad():\n    for data in test_loader:\n        inputs, labels = data\n        inputs, labels = inputs.to(device), labels.to(device)\n\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        test_loss += loss.item()\n\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        _, labels_ = torch.max(labels.data, 1)\n        correct += (predicted == labels_).sum().item()\n        # batch_accuracy = 100 * correct / total\n        # print(f\"Test Batch: {test_cntr}/126, Loss: {loss.item():.4f}, Accuracy: {batch_accuracy:.2f}%\")\n        # test_cntr+=1\n        \n        torch.cuda.empty_cache()\n\ntest_loss /= len(test_loader)\ntest_accuracy = 100 * correct / total\n\nprint(f\"Test Accuracy: {test_accuracy:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T10:56:04.509658Z","iopub.execute_input":"2025-04-15T10:56:04.509946Z","execution_failed":"2025-04-15T10:57:58.756Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig = plt.figure(figsize=(12, 6))\n\nplt.subplot(1, 2, 1)\nplt.plot(train_accuracies)\nplt.plot(val_accuracies)\nplt.title('Model accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.legend(['Train', 'Validation'], loc='upper left')\n\nplt.subplot(1, 2, 2)\nplt.plot(train_losses)\nplt.plot(val_losses)\nplt.title('Model loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend(['Train', 'Validation'], loc='upper left')\n\nplt.tight_layout()\nplt.show()\nfig.savefig(\"EffNetB3_v15_plots.png\", dpi=fig.dpi)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T08:49:49.724294Z","iopub.execute_input":"2025-04-15T08:49:49.724581Z","iopub.status.idle":"2025-04-15T08:49:50.340857Z","shell.execute_reply.started":"2025-04-15T08:49:49.724559Z","shell.execute_reply":"2025-04-15T08:49:50.340233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\n\n# Collecting ground truth and predictions\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for data in test_loader:\n        inputs, labels = data\n        inputs, labels = inputs.to(device), labels.to(device)\n\n        outputs = model(inputs)\n        _, predicted = torch.max(outputs, 1)  # Get predicted class index\n        _, labels_ = torch.max(labels, 1)  # Get true class index\n        \n        all_preds.extend(predicted.cpu().numpy())\n        all_labels.extend(labels_.cpu().numpy())\n\n# Compute the confusion matrix\ncm = confusion_matrix(all_labels, all_preds)\n\n# Display the confusion matrix\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=brain_activities)\nfig, ax = plt.subplots(figsize=(8, 8))\ndisp.plot(ax=ax, cmap='Blues', values_format='d')\n\nplt.title(\"Confusion Matrix for HMS Model\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T08:49:54.786558Z","iopub.execute_input":"2025-04-15T08:49:54.787131Z","iopub.status.idle":"2025-04-15T08:49:58.750334Z","shell.execute_reply.started":"2025-04-15T08:49:54.787107Z","shell.execute_reply":"2025-04-15T08:49:58.749444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}