{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport pywt\n\nimport torch\nfrom torch import nn\nfrom torch.optim import Adam\nimport torchvision.models as models\nimport torch.nn.functional as f\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchaudio import transforms\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, confusion_matrix\n\nimport glob\nfrom tqdm.notebook import tqdm, trange","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-11-18T09:14:56.605047Z","iopub.execute_input":"2024-11-18T09:14:56.605523Z","iopub.status.idle":"2024-11-18T09:15:03.113144Z","shell.execute_reply.started":"2024-11-18T09:14:56.605471Z","shell.execute_reply":"2024-11-18T09:15:03.112130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metadata = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\n\neeg_paths = glob.glob(\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/*\")\neeg_path_df = pd.DataFrame(eeg_paths, columns=[\"path\"])\neeg_path_df[\"eeg_id\"] = eeg_path_df[\"path\"].apply(lambda x: int(x.split(\"/\")[-1].replace(\".parquet\", \"\")))","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:03.114597Z","iopub.execute_input":"2024-11-18T09:15:03.115034Z","iopub.status.idle":"2024-11-18T09:15:03.654689Z","shell.execute_reply.started":"2024-11-18T09:15:03.115001Z","shell.execute_reply":"2024-11-18T09:15:03.653877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating non-overlapping eegs","metadata":{}},{"cell_type":"code","source":"df = metadata[[\"eeg_id\", \"spectrogram_id\", \"label_id\", \"patient_id\", \"expert_consensus\"]].groupby(\"eeg_id\").agg([\"first\"])\ndf.columns = [\"spectrogram_id\", \"label_id\", \"patient_id\", \"expert_consensus\"]\n\ndf[metadata.columns[-6:]] = metadata[[\"eeg_id\"] + list(metadata.columns[-6:])].groupby(\"eeg_id\").agg([\"sum\"])\ntemp = df[metadata.columns[-6:]].values.astype(np.float64)\ntemp /= temp.sum(axis=1, keepdims=True)\ndf[metadata.columns[-6:]] = temp\n\ndf = df.reset_index()","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:11.055896Z","iopub.execute_input":"2024-11-18T09:15:11.056769Z","iopub.status.idle":"2024-11-18T09:15:11.126962Z","shell.execute_reply.started":"2024-11-18T09:15:11.056729Z","shell.execute_reply":"2024-11-18T09:15:11.126213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:13.455149Z","iopub.execute_input":"2024-11-18T09:15:13.456029Z","iopub.status.idle":"2024-11-18T09:15:13.485469Z","shell.execute_reply.started":"2024-11-18T09:15:13.455984Z","shell.execute_reply":"2024-11-18T09:15:13.484536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.merge(eeg_path_df, on=\"eeg_id\")","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:13.893502Z","iopub.execute_input":"2024-11-18T09:15:13.893821Z","iopub.status.idle":"2024-11-18T09:15:13.909459Z","shell.execute_reply.started":"2024-11-18T09:15:13.893790Z","shell.execute_reply":"2024-11-18T09:15:13.908665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessor Utilities","metadata":{}},{"cell_type":"code","source":"class Preprocess:\n    \n    @staticmethod\n    def read_eeg(df, idx):\n        path = df.loc[idx, \"path\"]\n        eeg_seg = pd.read_parquet(path)\n        \n        eeg_seg = eeg_seg.T\n        eeg_seg = eeg_seg.values\n        \n        return eeg_seg\n    \n    @staticmethod\n    def to_tensor(eeg_seg):\n        for i in range(eeg_seg.shape[0]):\n            eeg_mean = np.nanmean(eeg_seg[i], axis=0)\n            \n            if (np.isnan(eeg_seg[i]).mean() == 1):\n                eeg_seg[i] = np.zeros_like(eeg_seg[i])\n            else:\n                eeg_seg[i] = np.nan_to_num(eeg_seg[i],nan=eeg_mean)\n            \n        \n        return torch.tensor(eeg_seg)\n    \n    @staticmethod\n    def to_spec(signal, sr=200, length = 50, n_mels=128, n_fft=1024, width = 256, top_db = 160):\n        if signal.shape[1] < sr*length:\n            pad_len = sr*length - signal.shape[1]\n            pad_sig = torch.zeros((signal.shape[0], pad_len))\n            signal = torch.cat([pad_len, pad_signal], dim=1)\n            \n        else:\n            mid = (signal.shape[1] - sr*length) // 2\n            signal = signal[:, mid:mid + sr*length]\n            \n        specs = torch.zeros(signal.shape[0], n_mels, width)    \n        for i in range(signal.shape[0]):\n            specs[i,:,:] = transforms.MelSpectrogram(sr, n_fft=n_fft, hop_length=sr*length//width, n_mels=n_mels, win_length=128)(signal[i])[:, :width]\n            specs[i,:,:] = transforms.AmplitudeToDB(top_db=None)(specs[i,:,:])\n            specs[i,:,:] = (specs[i,:,:] + top_db) / top_db\n            \n            \n        return specs     ","metadata":{"execution":{"iopub.status.busy":"2024-11-18T12:31:03.652024Z","iopub.execute_input":"2024-11-18T12:31:03.652424Z","iopub.status.idle":"2024-11-18T12:31:03.665428Z","shell.execute_reply.started":"2024-11-18T12:31:03.652386Z","shell.execute_reply":"2024-11-18T12:31:03.664464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pp = Preprocess()\nsignal = pp.read_eeg(train, 10)\n# signal = pp.denoise(signal)\nsignal = pp.to_tensor(signal)\nspecs = pp.to_spec(signal)\n\nprint(signal.shape)\nprint(specs.shape)\n\nplt.subplots(10, 2, figsize=(10, 50))\nfor i in range(20):\n    plt.subplot(10,2,i+1)\n    plt.imshow(specs[:,:,i],aspect='auto',origin='lower')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-18T12:31:03.960601Z","iopub.execute_input":"2024-11-18T12:31:03.961456Z","iopub.status.idle":"2024-11-18T12:31:08.531328Z","shell.execute_reply.started":"2024-11-18T12:31:03.961414Z","shell.execute_reply":"2024-11-18T12:31:08.530048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DataLoader","metadata":{}},{"cell_type":"code","source":"class SpecDataset(Dataset):\n    def __init__(self, df):\n        super().__init__()\n        self.df = df\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        pp = Preprocess()\n        \n        signal = pp.read_eeg(self.df, idx)\n        signal = pp.to_tensor(signal)\n        spectrograms = pp.to_spec(signal,top_db = 160)\n#         spectrograms = torch.cat([spectrograms, spectrograms])\n        \n        label = self.df.loc[idx, self.df.columns[-7:-1]].values.astype(np.float64)\n        label = torch.tensor(label)\n        \n        return spectrograms, label\n    \n# ds = SpecDataset(df)    \n# dl = DataLoader(ds, 64)","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:17.203484Z","iopub.execute_input":"2024-11-18T09:15:17.204103Z","iopub.status.idle":"2024-11-18T09:15:17.211245Z","shell.execute_reply.started":"2024-11-18T09:15:17.204063Z","shell.execute_reply":"2024-11-18T09:15:17.210237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class PreprocessModel(nn.Module):\n    def __init__(self, num_channels, out_features, concat=False):\n        super().__init__()\n        self.concat = concat\n        self.depthwise_conv = nn.Conv2d(num_channels, num_channels, 3, 1, 1, groups= num_channels)\n        self.pointwise_conv = nn.Conv2d(num_channels, out_features, 1)\n        \n    def forward(self, x):\n        if self.concat:\n            x = torch.cat([x, x], dim=2)\n            \n        x = self.depthwise_conv(x)\n        x = self.pointwise_conv(x)\n        \n        return x\n    \n# model = PreprocessModel(20, 3)\n# model(torch.zeros(2, 20, 128, 256)).shape","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:18.694257Z","iopub.execute_input":"2024-11-18T09:15:18.694972Z","iopub.status.idle":"2024-11-18T09:15:18.701459Z","shell.execute_reply.started":"2024-11-18T09:15:18.694932Z","shell.execute_reply":"2024-11-18T09:15:18.700448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClassifierModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.preprocessor = self.preprocessor_block()\n        self.base_model = models.efficientnet_b0()\n        self.base_model = nn.Sequential(*list(self.base_model.children())[:-1])\n        \n        self.avg_pool = nn.AvgPool2d(7)\n        self.classifier = self.classifier_block()\n        self.sigmoid = nn.Sigmoid()\n        \n    def preprocessor_block(self):\n        return nn.Sequential(*[\n            PreprocessModel(20, 10, concat=True),\n            nn.BatchNorm2d(10),\n            nn.ReLU(),\n            \n            PreprocessModel(10, 3),\n            nn.BatchNorm2d(3),\n            nn.ReLU(),            \n        ])\n    \n    def classifier_block(self):\n        return nn.Sequential(*[\n            nn.Linear(1280, 640),\n            nn.BatchNorm1d(640),\n            nn.ReLU(),\n            \n            nn.Linear(640, 160),\n            nn.BatchNorm1d(160),\n            nn.ReLU(),\n            \n            nn.Linear(160, 6),\n            nn.BatchNorm1d(6),\n        ])\n        \n    def forward(self, x):\n        x = self.preprocessor(x)\n        x = self.base_model(x)\n#         x = self.avg_pool(x)\n        x = x.flatten(start_dim=1)\n        \n        x = self.classifier(x)\n        x = self.sigmoid(x)\n        \n        return x\n    \n# model = ClassifierModel()\n# model(torch.zeros(64, 20, 128, 256)).shape\n    ","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:19.844470Z","iopub.execute_input":"2024-11-18T09:15:19.845330Z","iopub.status.idle":"2024-11-18T09:15:19.856302Z","shell.execute_reply.started":"2024-11-18T09:15:19.845277Z","shell.execute_reply":"2024-11-18T09:15:19.855161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device1 = \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\ndevice2 = \"cuda:1\" if torch.cuda.is_available() else \"cpu\"\n\ntrain, val = train_test_split(df, test_size=0.20, shuffle=False)\ntrain, val = train.reset_index(drop=False), val.reset_index(drop=False)\n\ntrain_ds = SpecDataset(train)    \ntrain_dl = DataLoader(train_ds, 64)\n\nval_ds = SpecDataset(val)    \nval_dl = DataLoader(val_ds, 64)\n\nmodel = ClassifierModel().to(device1)\nmodel = model.train()\n\nepochs = 10\n\nlr = 1e-3\noptimizer = Adam(model.parameters(), lr)\n\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:20.593797Z","iopub.execute_input":"2024-11-18T09:15:20.594405Z","iopub.status.idle":"2024-11-18T09:15:21.021778Z","shell.execute_reply.started":"2024-11-18T09:15:20.594365Z","shell.execute_reply":"2024-11-18T09:15:21.020927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:21.594652Z","iopub.execute_input":"2024-11-18T09:15:21.595351Z","iopub.status.idle":"2024-11-18T09:15:21.599800Z","shell.execute_reply.started":"2024-11-18T09:15:21.595294Z","shell.execute_reply":"2024-11-18T09:15:21.598554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"losses = []\naccuracies = []","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:15:21.946165Z","iopub.execute_input":"2024-11-18T09:15:21.946952Z","iopub.status.idle":"2024-11-18T09:15:21.950839Z","shell.execute_reply.started":"2024-11-18T09:15:21.946914Z","shell.execute_reply":"2024-11-18T09:15:21.949788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train Loop","metadata":{}},{"cell_type":"code","source":"for epoch in range(epochs):\n    model = model.train()\n    model.to(device1)\n    \n    torch.cuda.empty_cache()\n    \n    for i, (X, y) in enumerate(tqdm(train_dl)):\n        optimizer.zero_grad()\n#         print(X.shape)\n        \n        X = X.to(device1)\n        y = y.to(device1)\n        \n        preds = model(X)\n        loss = criterion(preds, y)\n        \n        loss.backward()\n        optimizer.step()\n        \n    avg_loss = 0\n    y_labels = []\n    preds_labels = []\n    \n    model = model.eval()\n    model.to(device2)\n    \n    torch.cuda.empty_cache()\n    for j, (X, y) in enumerate(tqdm(val_dl)):\n        with torch.no_grad():\n            X = X.to(device2)\n            y = y.to(device2)\n\n            preds = model(X)\n            loss = criterion(preds, y)\n            avg_loss += loss / len(val_dl)\n\n            preds_label = torch.argmax(preds, dim=1).detach().cpu().numpy()\n            preds_labels.extend(preds_label)\n            y_label = torch.argmax(y, dim=1).detach().cpu().numpy()\n            y_labels.extend(y_label)\n        \n    acc = accuracy_score(y_labels, preds_labels)\n    losses.append(avg_loss.item())\n    accuracies.append(acc)\n    \n    \n    print(f\"epoch = {epoch}, loss = {avg_loss.item()}, accuracy = {acc}\")","metadata":{"execution":{"iopub.status.busy":"2024-11-18T09:18:14.750114Z","iopub.execute_input":"2024-11-18T09:18:14.750554Z","iopub.status.idle":"2024-11-18T12:27:43.366829Z","shell.execute_reply.started":"2024-11-18T09:18:14.750514Z","shell.execute_reply":"2024-11-18T12:27:43.365776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(losses)","metadata":{"execution":{"iopub.status.busy":"2024-11-18T12:28:10.835054Z","iopub.execute_input":"2024-11-18T12:28:10.835925Z","iopub.status.idle":"2024-11-18T12:28:11.117153Z","shell.execute_reply.started":"2024-11-18T12:28:10.835881Z","shell.execute_reply":"2024-11-18T12:28:11.116212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(accuracies)","metadata":{"execution":{"iopub.status.busy":"2024-11-18T12:28:12.867038Z","iopub.execute_input":"2024-11-18T12:28:12.867437Z","iopub.status.idle":"2024-11-18T12:28:13.099544Z","shell.execute_reply.started":"2024-11-18T12:28:12.867398Z","shell.execute_reply":"2024-11-18T12:28:13.098601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/working/model.h5\"\ntorch.save(model.state_dict(), MODEL_PATH)","metadata":{"execution":{"iopub.status.busy":"2024-11-18T12:28:23.397603Z","iopub.execute_input":"2024-11-18T12:28:23.398002Z","iopub.status.idle":"2024-11-18T12:28:23.477519Z","shell.execute_reply.started":"2024-11-18T12:28:23.397962Z","shell.execute_reply":"2024-11-18T12:28:23.476502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load(MODEL_PATH))","metadata":{"execution":{"iopub.status.busy":"2024-11-18T12:29:15.431830Z","iopub.execute_input":"2024-11-18T12:29:15.432624Z","iopub.status.idle":"2024-11-18T12:29:15.534654Z","shell.execute_reply.started":"2024-11-18T12:29:15.432570Z","shell.execute_reply":"2024-11-18T12:29:15.533599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}