{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7414022,"sourceType":"datasetVersion","datasetId":4312784},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995}],"dockerImageVersionId":30664,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport time\n\nfrom IPython.display import clear_output\nfrom tqdm import tqdm\nfrom tqdm.contrib import tzip\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport umap\n\nfrom sklearn.metrics.pairwise import euclidean_distances\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.transforms import Resize\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\n\n\n# import albumentations as A\nimport albumentations as albu\nfrom albumentations.pytorch import ToTensorV2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-04T23:13:04.34628Z","iopub.execute_input":"2024-03-04T23:13:04.346539Z","iopub.status.idle":"2024-03-04T23:13:36.720391Z","shell.execute_reply.started":"2024-03-04T23:13:04.346511Z","shell.execute_reply":"2024-03-04T23:13:36.719505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Specs","metadata":{}},{"cell_type":"code","source":"load_pretrained = False\nload_embs = False\ntest_eval = False\nfind_threshold = False\nget_dists = False","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:13:36.721888Z","iopub.execute_input":"2024-03-04T23:13:36.722432Z","iopub.status.idle":"2024-03-04T23:13:36.727383Z","shell.execute_reply.started":"2024-03-04T23:13:36.722406Z","shell.execute_reply":"2024-03-04T23:13:36.726167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    dct = np.load('/kaggle/input/brain-eeg-spectrograms/eeg_specs.npy', allow_pickle=True).item()\n# 1.0 - /kaggle/input/brain-spectrograms\n# 2.0 - /kaggle/input/eeg-spectrograms\n# 3.0 - /kaggle/input/brain-eeg-spectrograms","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:13:36.728457Z","iopub.execute_input":"2024-03-04T23:13:36.728742Z","iopub.status.idle":"2024-03-04T23:14:48.254061Z","shell.execute_reply.started":"2024-03-04T23:13:36.728719Z","shell.execute_reply":"2024-03-04T23:14:48.253217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.257033Z","iopub.execute_input":"2024-03-04T23:14:48.257453Z","iopub.status.idle":"2024-03-04T23:14:48.557903Z","shell.execute_reply.started":"2024-03-04T23:14:48.257416Z","shell.execute_reply":"2024-03-04T23:14:48.556869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\nTARGETS = train_df.columns[-6:]","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.559162Z","iopub.execute_input":"2024-03-04T23:14:48.559505Z","iopub.status.idle":"2024-03-04T23:14:48.751286Z","shell.execute_reply.started":"2024-03-04T23:14:48.559476Z","shell.execute_reply":"2024-03-04T23:14:48.750474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Preprocess 1","metadata":{}},{"cell_type":"code","source":"train_df.iloc[:,-6:] = train_df.iloc[:,-6:].values / train_df.iloc[:,-6:].sum(axis=1).values.reshape((-1, 1))\ncols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\ntrain_df['entropy'] = -(train_df[cols] * np.log(train_df[cols])).sum(axis=1)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.75242Z","iopub.execute_input":"2024-03-04T23:14:48.752763Z","iopub.status.idle":"2024-03-04T23:14:48.844755Z","shell.execute_reply.started":"2024-03-04T23:14:48.752735Z","shell.execute_reply":"2024-03-04T23:14:48.843848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.845958Z","iopub.execute_input":"2024-03-04T23:14:48.846265Z","iopub.status.idle":"2024-03-04T23:14:48.881163Z","shell.execute_reply.started":"2024-03-04T23:14:48.84624Z","shell.execute_reply":"2024-03-04T23:14:48.880262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patients = train_df.patient_id.unique()\npatients_train, patients_val_test, _, _ = train_test_split(patients, np.arange(len(patients)), test_size=0.3, random_state=123)\npatients_val, patients_test, _, _ = train_test_split(patients_val_test, np.arange(len(patients_val_test)), test_size=0.5, random_state=123)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.882238Z","iopub.execute_input":"2024-03-04T23:14:48.882514Z","iopub.status.idle":"2024-03-04T23:14:48.894339Z","shell.execute_reply.started":"2024-03-04T23:14:48.882489Z","shell.execute_reply":"2024-03-04T23:14:48.8935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_patients = train_df.loc[train_df.patient_id.isin(patients_train)].copy().reset_index(drop=True)\nval_patients = train_df.loc[train_df.patient_id.isin(patients_val)].copy().reset_index(drop=True)\ntest_patients = train_df.loc[train_df.patient_id.isin(patients_test)].copy().reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.895593Z","iopub.execute_input":"2024-03-04T23:14:48.896221Z","iopub.status.idle":"2024-03-04T23:14:48.933161Z","shell.execute_reply.started":"2024-03-04T23:14:48.896187Z","shell.execute_reply":"2024-03-04T23:14:48.932351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Preprocess 2","metadata":{}},{"cell_type":"code","source":"all_eegs = dct","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.93734Z","iopub.execute_input":"2024-03-04T23:14:48.937607Z","iopub.status.idle":"2024-03-04T23:14:48.941512Z","shell.execute_reply.started":"2024-03-04T23:14:48.937584Z","shell.execute_reply":"2024-03-04T23:14:48.940706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARS = {'Seizure':0, 'LPD':1, 'GPD':2, 'LRDA':3, 'GRDA':4, 'Other':5}\nTARS2 = {x: y for y, x in TARS.items()}\n\n\nclass EEGDataset(Dataset):\n    \n#     def __init__(self, data, augment=False, mode='train', specs=spectrograms, eeg_specs=all_eegs): \n    def __init__(self, data, augment=False, mode='train', eeg_specs=all_eegs): \n        self.data = data\n        self.augment = augment\n        self.mode = mode\n#         self.specs = specs\n        self.eeg_specs = eeg_specs\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, index):\n        return self.__getitems__([index])\n    \n    def __getitems__(self, indices):\n        X, y = self._generate_data(indices)\n        if self.augment:\n            X = self.__augment(X) \n        if self.mode == 'train':\n            return list(zip(X, y))\n        else:\n            return X\n    \n#     def _generate_data(self, indexes):\n#         X = np.zeros((len(indexes), 128, 256, 4),dtype='float32')\n#         y = np.zeros((len(indexes), 6),dtype='float32')\n#         img = np.ones((128, 256),dtype='float32')\n        \n#         for j, i in enumerate(indexes):\n#             row = self.data.iloc[i]\n        \n#             # EEG SPECTROGRAMS\n#             img = self.eeg_specs[row.eeg_id]\n#             X[j, :, :, :] = img\n                \n#             if self.mode != 'test':\n#                 y[j,] = row[TARGETS]\n            \n#         return X, y\n    def _generate_data(self, indexes):\n        X = np.zeros((len(indexes), 4, 128, 256), dtype='float32')  # Adjusted the shape here\n        y = np.zeros((len(indexes), 6), dtype='float32')\n        \n        for j, i in enumerate(indexes):\n            row = self.data.iloc[i]\n        \n            # EEG SPECTROGRAMS\n            img = self.eeg_specs[row.eeg_id]\n            img = np.transpose(img, (2, 0, 1))  # Adjusted the order here to match (channels, height, width)\n            X[j] = img\n                \n            if self.mode != 'test':\n                # Assuming you have a correct way to set y[j,] based on your labels\n                y[j,] = row[TARGETS]  # Make sure TARGETS is correctly defined elsewhere\n            \n#         return torch.tensor(X), torch.tensor(y)\n        return X, y\n    \n    def _random_transform(self, img):\n        composition = albu.Compose([\n            albu.HorizontalFlip(p=0.5),\n            # albu.CoarseDropout(max_holes=8,max_height=32,max_width=32,fill_value=0,p=0.5),\n        ])\n        return composition(image=img)['image']\n            \n    def __augment(self, img_batch):\n        for i in range(img_batch.shape[0]):\n            img_batch[i,] = self._random_transform(img_batch[i,])\n        return img_batch","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.943079Z","iopub.execute_input":"2024-03-04T23:14:48.943512Z","iopub.status.idle":"2024-03-04T23:14:48.960499Z","shell.execute_reply.started":"2024-03-04T23:14:48.943478Z","shell.execute_reply":"2024-03-04T23:14:48.959675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = EEGDataset(train_patients)\ndataloader = DataLoader(dataset, batch_size=32, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.961637Z","iopub.execute_input":"2024-03-04T23:14:48.962352Z","iopub.status.idle":"2024-03-04T23:14:48.975428Z","shell.execute_reply.started":"2024-03-04T23:14:48.962325Z","shell.execute_reply":"2024-03-04T23:14:48.974593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport torch\n\nROWS = 2\nCOLS = 3\nBATCHES = 2\n\nfor i, (x, y) in enumerate(dataloader):\n    plt.figure(figsize=(20, 8))\n    batch_size = x.size(0)\n    for j in range(ROWS):\n        for k in range(COLS):\n            idx = j*COLS + k\n            if idx < batch_size:  # Make sure not to exceed the batch size\n                plt.subplot(ROWS, COLS, idx + 1)\n                t = y[idx]\n                img = x[idx, 0, :, :]  # Assuming the channel dimension is second, not last\n                mn, mx = img.min(), img.max()\n                img = (img - mn) / (mx - mn)  # Normalize the image\n                plt.imshow(img, aspect='auto', origin='lower')  # Set origin to 'lower' to correct orientation\n                tars = ', '.join([f'{val:.2f}' for val in t.tolist()])\n                eeg_id = train_patients.eeg_id.values[i*batch_size + idx]\n                plt.title(f'EEG = {eeg_id}\\nTarget = [{tars}]', size=12)\n                plt.yticks([])\n                plt.ylabel('Frequencies (Hz)', size=14)\n                plt.xlabel('Time (sec)', size=16)\n    plt.tight_layout()  # Adjust the layout\n    plt.show()\n    if i == BATCHES - 1:\n        break\n","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:48.976597Z","iopub.execute_input":"2024-03-04T23:14:48.977045Z","iopub.status.idle":"2024-03-04T23:14:51.727007Z","shell.execute_reply.started":"2024-03-04T23:14:48.977001Z","shell.execute_reply":"2024-03-04T23:14:51.726025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ROWS = 2\n# COLS = 3\n# BATCHES = 2\n\n# for i, (x, y) in enumerate(dataloader):\n#     plt.figure(figsize=(20, 8))\n#     for j in range(ROWS):\n#         for k in range(COLS):\n#             plt.subplot(ROWS, COLS, j*COLS + k + 1)\n#             t = y[j*COLS + k]\n#             img = torch.flip(x[j*COLS+k, :, :, 0], (0,))\n#             mn = img.flatten().min()\n#             mx = img.flatten().max()\n#             img = (img-mn)/(mx-mn)\n#             plt.imshow(img)\n#             tars = f'[{t[0]:0.2f}]'\n#             for s in t[1:]:\n#                 tars += f', {s:0.2f}'\n#             eeg = train_patients.eeg_id.values[i*32+j*COLS+k]\n#             plt.title(f'EEG = {eeg}\\nTarget = {tars}',size=12)\n#             plt.yticks([])\n#             plt.ylabel('Frequencies (Hz)',size=14)\n#             plt.xlabel('Time (sec)',size=16)\n#     plt.show()\n#     if i == BATCHES-1:\n#         break","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:51.728716Z","iopub.execute_input":"2024-03-04T23:14:51.729066Z","iopub.status.idle":"2024-03-04T23:14:51.734516Z","shell.execute_reply.started":"2024-03-04T23:14:51.729037Z","shell.execute_reply":"2024-03-04T23:14:51.7336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomUpsample(nn.Module):\n    def __init__(self, scale_factor=(2, 2)):\n        super().__init__()\n        self.scale_factor = scale_factor\n\n    def forward(self, x):\n        return F.interpolate(x, scale_factor=self.scale_factor, mode='nearest')\n\n\nclass ResNetBlock(nn.Module):\n    def __init__(self, in_channels, kernel_size, modify=False, bn=True, scale_factor=(1, 1)):\n        super().__init__()\n        self.modify = modify\n        if modify=='downsample':\n            self.conv1 = nn.Conv2d(in_channels=in_channels, out_channels=in_channels*2, stride=2, kernel_size=kernel_size, padding=kernel_size//2, bias=False)\n            self.conv2 = nn.Conv2d(in_channels=in_channels*2, out_channels=in_channels*2, kernel_size=kernel_size, padding=kernel_size//2,bias=False)\n            if bn:\n                self.bn1 = nn.BatchNorm2d(in_channels*2)\n                self.bn2 = nn.BatchNorm2d(in_channels*2)\n            else:\n                self.bn1 = nn.Identity()\n                self.bn2 = nn.Identity()\n                \n        elif modify=='upsample':\n#             self.conv1 = CustomUpsample(scale_factor=scale_factor)\n            self.conv1 = nn.ConvTranspose2d(in_channels=in_channels, out_channels=in_channels//2, stride=2, kernel_size=kernel_size, output_padding=scale_factor, padding=kernel_size//2, bias=False)\n            self.conv2 = nn.Conv2d(in_channels=in_channels//2, out_channels=in_channels//2, kernel_size=kernel_size, padding=kernel_size//2, bias=False)\n            self.bn1 = nn.BatchNorm2d(in_channels//2)\n            self.bn2 = nn.BatchNorm2d(in_channels//2)\n        else:\n            self.conv1 = nn.Conv2d(in_channels=in_channels, out_channels=in_channels, kernel_size=kernel_size, padding=kernel_size//2)\n            self.conv2 = nn.Conv2d(in_channels=in_channels, out_channels=in_channels, kernel_size=kernel_size, padding=kernel_size//2)\n            self.bn1 = nn.BatchNorm2d(in_channels)\n            self.bn2 = nn.BatchNorm2d(in_channels)\n        self.act = nn.ReLU()\n        \n        if modify=='downsample':\n            self.proj = nn.Conv2d(in_channels=in_channels, out_channels=in_channels*2, stride=2, kernel_size=kernel_size, padding=kernel_size//2)\n        if modify=='upsample':\n            self.proj = nn.ConvTranspose2d(in_channels=in_channels, out_channels=in_channels//2, stride=2, kernel_size=kernel_size, output_padding=scale_factor, padding=kernel_size//2)\n\n\n    def forward(self, x):\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.act(out)\n        out = self.conv2(out)\n        out = self.bn2(out)\n        if self.modify:\n            x = self.proj(x)\n        out = x + out\n        out = self.act(out)\n        return out\n\n    \nclass Encoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv = nn.Conv2d(4, 16, 7, 1, 7//2)\n#         2x specs\n#         self.conv = nn.Conv2d(8, 16, 7, 1, 7//2)\n        self.rnb1 = ResNetBlock(16, 3, modify='downsample')\n        self.rnb2 = ResNetBlock(32, 3, modify='downsample')\n        self.rnb3 = ResNetBlock(64, 3, modify='downsample')\n        self.rnb4 = ResNetBlock(128, 3, modify='downsample')\n        self.rnb5 = ResNetBlock(256, 3, modify='downsample')\n        self.rnb6 = ResNetBlock(512, 3, modify='downsample')\n        self.rnb7 = ResNetBlock(1024, 3, modify='downsample')\n        self.rnb8 = ResNetBlock(2048, 3, modify='downsample')\n        \n    def forward(self, x):\n#         print(f\"Input to Encoder: {x.size()}\")\n        x = self.conv(x)\n#         print(f\"After conv: {x.size()}\")\n        x = self.rnb1(x)\n#         print(f\"After rnb1: {x.size()}\")\n        x = self.rnb2(x)\n#         print(f\"After rnb2: {x.size()}\")\n        x = self.rnb3(x)\n#         print(f\"After rnb3: {x.size()}\")\n        x = self.rnb4(x)\n#         print(f\"After rnb4: {x.size()}\")\n        x = self.rnb5(x)\n#         print(f\"After rnb5: {x.size()}\")\n        x = self.rnb6(x)\n#         print(f\"After rnb6: {x.size()}\")\n        x = self.rnb7(x)\n#         print(f\"After rnb7: {x.size()}\")\n        x = self.rnb8(x)\n#         print(f\"After rnb8: {x.size()}\")\n        return x\n    \nclass Decoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.rnb1 = ResNetBlock(4096, 3, modify='upsample', scale_factor=(0, 1))\n        self.rnb2 = ResNetBlock(2048, 3, modify='upsample')\n        self.rnb3 = ResNetBlock(1024, 3, modify='upsample')\n        self.rnb4 = ResNetBlock(512, 3, modify='upsample')\n        self.rnb5 = ResNetBlock(256, 3, modify='upsample')\n        self.rnb6 = ResNetBlock(128, 3, modify='upsample')\n        self.rnb7 = ResNetBlock(64, 3, modify='upsample')\n        self.rnb8 = ResNetBlock(32, 3, modify='upsample')\n#         self.rnb9 = ResNetBlock(16, 3, modify='upsample')\n        self.conv = nn.Conv2d(16, 4, 3, 1, 3//2)\n\n    def forward(self, x):\n#         print(f\"Input to Decoder: {x.size()}\")\n        x = self.rnb1(x)\n#         print(f\"After rnb1: {x.size()}\")\n        x = self.rnb2(x)\n#         print(f\"After rnb2: {x.size()}\")\n        x = self.rnb3(x)\n#         print(f\"After rnb3: {x.size()}\")\n        x = self.rnb4(x)\n#         print(f\"After rnb4: {x.size()}\")\n        x = self.rnb5(x)\n#         print(f\"After rnb5: {x.size()}\")\n        x = self.rnb6(x)\n#         print(f\"After rnb6: {x.size()}\")\n        x = self.rnb7(x)\n#         print(f\"After rnb7: {x.size()}\")\n        x = self.rnb8(x)\n#         print(f\"After rnb8: {x.size()}\")\n#         x = self.rnb9(x)\n        x = self.conv(x)\n#         print(f\"After final conv: {x.size()}\")\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:51.736091Z","iopub.execute_input":"2024-03-04T23:14:51.73637Z","iopub.status.idle":"2024-03-04T23:14:51.76667Z","shell.execute_reply.started":"2024-03-04T23:14:51.736346Z","shell.execute_reply":"2024-03-04T23:14:51.765738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SimpleAE(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.net = nn.Sequential(\n            Encoder(),\n            Decoder()\n        )\n        \n    def forward(self, x):\n        return self.net(x)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:51.76792Z","iopub.execute_input":"2024-03-04T23:14:51.768283Z","iopub.status.idle":"2024-03-04T23:14:51.779836Z","shell.execute_reply.started":"2024-03-04T23:14:51.768249Z","shell.execute_reply":"2024-03-04T23:14:51.77893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ae = SimpleAE()\nx = torch.rand(128, 4, 128, 256)  # Replace with your actual input dimensions\nx_recon = ae(x)\nassert x.size() == x_recon.size(), f\"Input {x.size()}, Output {x_recon.size()}\"","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:14:51.780962Z","iopub.execute_input":"2024-03-04T23:14:51.781248Z","iopub.status.idle":"2024-03-04T23:15:10.795111Z","shell.execute_reply.started":"2024-03-04T23:14:51.781225Z","shell.execute_reply":"2024-03-04T23:15:10.794277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_epoch(model, dataloader, loss_fn, optimizer, epoch, device, scaler):\n    model = model.to(device)\n    model.train()\n    losses = []\n    for batch in tqdm(dataloader, total=len(dataloader)):\n#         print(batch[0])\n        x = batch[0].to(device)\n        y = batch[1]\n        \n#         with torch.autocast(device_type='cuda' if device=='cuda' else 'cpu', dtype=torch.float16 if device=='cuda' else torch.bfloat16):\n        x_recon = model(x)\n        loss = loss_fn(x, x_recon)\n\n        loss.backward()\n        optimizer.step()\n#         scaler.scale(loss).backward()\n#         scaler.step(optimizer)\n#         scaler.update()\n        \n        optimizer.zero_grad()\n                \n        losses.append(loss.detach().cpu().item())\n#     print(f'Не нан значений во время train: {np.count_nonzero(~np.isnan(losses))}')\n    return np.nanmean(losses)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:10.796356Z","iopub.execute_input":"2024-03-04T23:15:10.796639Z","iopub.status.idle":"2024-03-04T23:15:10.804211Z","shell.execute_reply.started":"2024-03-04T23:15:10.796615Z","shell.execute_reply":"2024-03-04T23:15:10.803045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate(model, dataloader, loss_fn, device, scaler):\n    model = model.to(device)\n    losses = []\n    with torch.no_grad():\n        model.eval()\n        for batch in tqdm(dataloader, total=len(dataloader)):\n            x = batch[0].to(device)\n            y = batch[1]\n\n#             with torch.autocast(device_type='cuda' if device=='cuda' else 'cpu', dtype=torch.float16 if device=='cuda' else torch.bfloat16):\n            x_recon = model(x)\n            loss = loss_fn(x, x_recon)\n#             scaler.scale(loss)\n            losses.append(loss.detach().cpu().item())\n#     print(f'Не нан значений во время eval: {np.count_nonzero(~np.isnan(losses))}')\n    return np.nanmean(losses)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:10.805555Z","iopub.execute_input":"2024-03-04T23:15:10.806451Z","iopub.status.idle":"2024-03-04T23:15:10.819714Z","shell.execute_reply.started":"2024-03-04T23:15:10.806425Z","shell.execute_reply":"2024-03-04T23:15:10.81889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_pretrained:\n    scaler = torch.cuda.amp.GradScaler()\n    def run_experiment(model, dataloader_train, dataloader_val, loss_fn, optimizer, num_epochs, device, stop_after=5, scaler=scaler):\n        losses_train = []\n        losses_val = []\n        best_loss_val = np.inf\n        c = 0\n        total_runtime = 0\n        for epoch in range(num_epochs):\n            start = time.time()\n\n            if c == stop_after:\n                print(f'Обучение остановлено, так как лосс на валидации не падал {stop_after} эпох')\n                break\n\n            loss_train = run_epoch(model, dataloader_train, loss_fn, optimizer, epoch, device, scaler)\n            loss_val = evaluate(model, dataloader_val, loss_fn, device, scaler)\n            losses_train.append(loss_train)\n            losses_val.append(loss_val)\n            clear_output()\n            if best_loss_val > loss_val:\n                torch.save(model.state_dict(), 'best_model.pth')\n                torch.save(optimizer, 'optimizer.pth')\n                best_loss_val = loss_val\n                c = 0\n            else:\n                c += 1\n\n            print(f\"epoch: {str(epoch).zfill(3)} | loss_train: {loss_train:5.5f} | loss_val: {loss_val:5.5f} | best_loss: {best_loss_val:5.5f}\")\n\n            plt.plot(losses_train, label='Loss train')\n            plt.plot(losses_val, label='Loss val')\n            plt.legend()\n            plt.show()\n\n            stop = time.time()\n            runtime = stop - start\n            total_runtime += runtime\n            if 12*60*60 - 600 - total_runtime < runtime:\n                break\n\n        return losses_train, losses_val, model","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:10.821033Z","iopub.execute_input":"2024-03-04T23:15:10.821393Z","iopub.status.idle":"2024-03-04T23:15:10.888022Z","shell.execute_reply.started":"2024-03-04T23:15:10.821363Z","shell.execute_reply":"2024-03-04T23:15:10.887033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    dataset_train = EEGDataset(train_patients, dct)\n    dataset_val = EEGDataset(val_patients, dct)\n    dataset_test = EEGDataset(test_patients, dct)\n\n    dataloader_train = DataLoader(\n        dataset=dataset_train,\n        batch_size=128,\n        shuffle=True,\n        drop_last=True\n    )\n\n    dataloader_val = DataLoader(\n        dataset=dataset_val,\n        batch_size=128,\n        shuffle=False,\n        drop_last=False\n    )\n\n    dataloader_test = DataLoader(\n        dataset=dataset_test,\n        batch_size=128,\n        shuffle=False,\n        drop_last=False\n    )","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:10.889592Z","iopub.execute_input":"2024-03-04T23:15:10.890053Z","iopub.status.idle":"2024-03-04T23:15:10.901935Z","shell.execute_reply.started":"2024-03-04T23:15:10.890019Z","shell.execute_reply":"2024-03-04T23:15:10.901177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nlr = 3e-4\nmodel = SimpleAE()\nmodel= nn.DataParallel(model)\nloss_fn = nn.MSELoss()\nif not load_pretrained:\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)\n    num_epochs = 20","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:10.902981Z","iopub.execute_input":"2024-03-04T23:15:10.903314Z","iopub.status.idle":"2024-03-04T23:15:18.252566Z","shell.execute_reply.started":"2024-03-04T23:15:10.903236Z","shell.execute_reply":"2024-03-04T23:15:18.251795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_weights(w):\n    if isinstance(w, nn.Linear) or isinstance(w, nn.Conv2d) or isinstance(w, nn.ConvTranspose2d):\n        nn.init.xavier_uniform_(w.weight)\nif not load_pretrained:   \n    model.apply(init_weights);","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:18.253565Z","iopub.execute_input":"2024-03-04T23:15:18.253854Z","iopub.status.idle":"2024-03-04T23:15:18.271776Z","shell.execute_reply.started":"2024-03-04T23:15:18.253825Z","shell.execute_reply":"2024-03-04T23:15:18.27091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if load_pretrained and not load_embs:\n    model.load_state_dict(torch.load('/kaggle/input/autoencoder-weights/best_model.pth', map_location=torch.device(device)))\n# optimizer = torch.load('/kaggle/input/autoencoder-weights/optimizer.pth', map_location=torch.device(device))","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:18.272824Z","iopub.execute_input":"2024-03-04T23:15:18.273061Z","iopub.status.idle":"2024-03-04T23:15:18.277641Z","shell.execute_reply.started":"2024-03-04T23:15:18.27304Z","shell.execute_reply":"2024-03-04T23:15:18.276802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_pretrained:\n    losses_train, losses_val, model = run_experiment(model, dataloader_train, dataloader_val, loss_fn, optimizer, num_epochs, device, stop_after=15)","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:15:18.278847Z","iopub.execute_input":"2024-03-04T23:15:18.279136Z","iopub.status.idle":"2024-03-04T23:44:24.115693Z","shell.execute_reply.started":"2024-03-04T23:15:18.279113Z","shell.execute_reply":"2024-03-04T23:44:24.114202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_pretrained:\n    plt.plot(losses_train, label='Loss train')\n    plt.plot(losses_val, label='Loss val')\n    plt.legend()","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:52:11.408691Z","iopub.execute_input":"2024-03-04T23:52:11.409119Z","iopub.status.idle":"2024-03-04T23:52:11.460566Z","shell.execute_reply.started":"2024-03-04T23:52:11.409085Z","shell.execute_reply":"2024-03-04T23:52:11.459193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if test_eval:\n    print('Test loss:', evaluate(model, dataloader_test, loss_fn, device, scaler))","metadata":{"execution":{"iopub.status.busy":"2024-03-04T23:10:16.993241Z","iopub.status.idle":"2024-03-04T23:10:16.993656Z","shell.execute_reply.started":"2024-03-04T23:10:16.993465Z","shell.execute_reply":"2024-03-04T23:10:16.993482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_embs(model, dataloader):\n    model = model.to(device)\n    losses = []\n    embs = []\n    encoder = nn.DataParallel(model.module.net[0])\n    decoder = nn.DataParallel(model.module.net[1])\n    with torch.no_grad():\n        encoder.eval()\n        decoder.eval()\n        for batch in tqdm(dataloader, total=len(dataloader)):\n            x = batch.to(device)\n            \n            emb = encoder(x)\n            x_recon = decoder(emb)\n            embs.append(emb.squeeze(2, 3))\n            for i in range(len(x)):\n                loss = loss_fn(x[i], x_recon[i])\n                losses.append(loss.detach().cpu().item())\n    return losses, torch.cat(embs, dim=0).detach().cpu().numpy()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    dataloader_train = DataLoader(\n        dataset=dataset_train,\n        batch_size=128,\n        shuffle=False,\n        drop_last=False\n    )\n    losses_test, embs_test = get_embs(model, dataloader_test)\n    losses_val, embs_val = get_embs(model, dataloader_val)\n    losses_train, embs_train = get_embs(model, dataloader_train)\n    \n    with open(r'losses_test.txt', 'w') as fp:\n        for item in losses_test:\n            fp.write(\"%s\\n\" % item)\n        print('Done')\n        \n    with open(r'losses_val.txt', 'w') as fp:\n        for item in losses_val:\n            fp.write(\"%s\\n\" % item)\n        print('Done')\n        \n    with open(r'losses_train.txt', 'w') as fp:\n        for item in losses_train:\n            fp.write(\"%s\\n\" % item)\n        print('Done')\n        \n    losses_test = np.array(losses_test, dtype=np.float64)\n    losses_val = np.array(losses_val, dtype=np.float64)\n    losses_train = np.array(losses_train, dtype=np.float64)\n    \n    np.save('specs_autoencoder_test_embs.npy', embs_test)\n    np.save('specs_autoencoder_val_embs.npy', embs_val)\n    np.save('specs_autoencoder_train_embs.npy', embs_train)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if load_embs:\n    losses_test = []\n    with open('/kaggle/input/autoencoder-weights/losses_test.txt', 'r') as fp:\n        for line in fp:\n            x = line[:-1]\n            losses_test.append(x)\n    losses_test = np.array(losses_test, dtype=np.float64)\n    embs_test = np.load('/kaggle/input/autoencoder-weights/specs_autoencoder_test_embs.npy')\n    \n    losses_val = []\n    with open('/kaggle/input/autoencoder-weights/losses_val.txt', 'r') as fp:\n        for line in fp:\n            x = line[:-1]\n            losses_val.append(x)\n    losses_val = np.array(losses_val, dtype=np.float64)\n    embs_val = np.load('/kaggle/input/autoencoder-weights/specs_autoencoder_val_embs.npy')\n    \n    losses_train = []\n    with open('/kaggle/input/autoencoder-weights/losses_train.txt', 'r') as fp:\n        for line in fp:\n            x = line[:-1]\n            losses_train.append(x)\n    losses_train = np.array(losses_train, dtype=np.float64)\n    embs_train = np.load('/kaggle/input/autoencoder-weights/specs_autoencoder_train_embs.npy')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:\n    from sklearn.preprocessing import StandardScaler\n    sc = StandardScaler().fit(embs_train)\n    embs_train = sc.transform(embs_train)\n    embs_val = sc.transform(embs_val)\n    embs_test = sc.transform(embs_test)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:\n    from sklearn.metrics import pairwise_distances\n    if find_threshold:\n        dists_train = pairwise_distances(embs_train).astype(np.float32)\n    dists_val = pairwise_distances(embs_val).astype(np.float32)\n    dists_test = pairwise_distances(embs_test).astype(np.float32)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_acc_fpr_tpr(index, patient_id, threshold, df, dists_matrix):\n    dist = dists_matrix[index]\n    preds = np.less(dist, threshold)\n    targets = (df.patient_id == patient_id).values\n    tp = np.sum(np.logical_and(preds, targets))\n    fp = np.sum(np.logical_and(preds, np.logical_not(targets)))\n    tn = np.sum(np.logical_and(np.logical_not(preds), np.logical_not(targets)))\n    fn = np.sum(np.logical_and(np.logical_not(preds), targets))\n\n    tpr = 0 if (tp + fn == 0) else float(tp) / float(tp + fn)\n    fpr = 0 if (fp + tn == 0) else float(fp) / float(fp + tn)\n    acc = float(tp + tn) / dist.size\n    return tpr, fpr, acc","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def find_opt_threshold(df, dists_matrix, thresholds):\n    tpr_mean = []\n    fpr_mean = []\n    acc_mean = []\n    for th in tqdm(thresholds):\n        tpr_ = []\n        fpr_ = []\n        acc_ = []\n        for i in range(len(df)):\n            p_id = df.patient_id[i]\n            tpr, fpr, acc = get_acc_fpr_tpr(i, p_id, th, df, dists_matrix)\n            tpr_.append(tpr)\n            fpr_.append(fpr)\n            acc_.append(acc)\n        tpr_mean.append(np.mean(tpr_))\n        fpr_mean.append(np.mean(fpr_))\n        acc_mean.append(np.mean(acc_))\n    return tpr_mean, fpr_mean, acc_mean","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"thresholds_train = np.arange(0, 1.001, 0.001)\nif find_threshold:\n    tpr_train, fpr_train, acc_train = find_opt_threshold(train_patients, dists_train, thresholds_train)\n    np.save('tpr_train', np.array(tpr_train))\n    np.save('fpr_train', np.array(fpr_train))\n    np.save('acc_train', np.array(acc_train))\nelse:\n    tpr_train = np.load('/kaggle/input/autoencoder-weights/tpr_train.npy')\n    fpr_train = np.load('/kaggle/input/autoencoder-weights/fpr_train.npy')\n    acc_train = np.load('/kaggle/input/autoencoder-weights/acc_train.npy')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 30\nthresholds_train[i], tpr_train[i], fpr_train[i], acc_train[i]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axs = plt.subplots(1, 3, figsize=(20, 5))\naxs[0].plot(thresholds_train, tpr_train)\naxs[1].plot(thresholds_train, fpr_train)\naxs[2].plot(fpr_train, tpr_train)\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:\n    tpr_val, fpr_val, acc_val = find_opt_threshold(val_patients, dists_val, [thresholds_train[i]])\n    tpr_test, fpr_test, acc_test = find_opt_threshold(test_patients, dists_test, [thresholds_train[i]])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:\n    tpr_val, fpr_val, acc_val","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:    \n    tpr_test, fpr_test, acc_test","metadata":{},"execution_count":null,"outputs":[]}]}