{"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":"none","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"datasetVersion","sourceId":7691678,"datasetId":4421634,"databundleVersionId":7789892},{"sourceType":"datasetVersion","sourceId":7564328,"datasetId":4404591,"databundleVersionId":7658781}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-25T22:09:44.808940Z","iopub.execute_input":"2024-02-25T22:09:44.809441Z","iopub.status.idle":"2024-02-25T22:10:29.296130Z","shell.execute_reply.started":"2024-02-25T22:09:44.809401Z","shell.execute_reply":"2024-02-25T22:10:29.294922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_pretrained = True\nload_embs = True\ntest_eval = False\nfind_threshold = False\nget_dists = False","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.298069Z","iopub.execute_input":"2024-02-25T22:10:29.298810Z","iopub.status.idle":"2024-02-25T22:10:29.304285Z","shell.execute_reply.started":"2024-02-25T22:10:29.298775Z","shell.execute_reply":"2024-02-25T22:10:29.302424Z"},"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-02-25T22:10:29.306217Z","iopub.execute_input":"2024-02-25T22:10:29.306944Z","iopub.status.idle":"2024-02-25T22:10:29.652605Z","shell.execute_reply.started":"2024-02-25T22:10:29.306899Z","shell.execute_reply":"2024-02-25T22:10:29.651264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.655330Z","iopub.execute_input":"2024-02-25T22:10:29.655743Z","iopub.status.idle":"2024-02-25T22:10:29.692840Z","shell.execute_reply.started":"2024-02-25T22:10:29.655711Z","shell.execute_reply":"2024-02-25T22:10:29.691659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.iloc[:,-6:] = train_df.iloc[:,-6:].values / train_df.iloc[:,-6:].sum(axis=1).values.reshape((-1, 1))","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.694468Z","iopub.execute_input":"2024-02-25T22:10:29.695667Z","iopub.status.idle":"2024-02-25T22:10:29.750447Z","shell.execute_reply.started":"2024-02-25T22:10:29.695610Z","shell.execute_reply":"2024-02-25T22:10:29.749424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cols = ['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)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.752001Z","iopub.execute_input":"2024-02-25T22:10:29.752797Z","iopub.status.idle":"2024-02-25T22:10:29.821458Z","shell.execute_reply.started":"2024-02-25T22:10:29.752754Z","shell.execute_reply":"2024-02-25T22:10:29.820641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Загружаю все спеки разом\nif not load_embs:\n    dct = np.load('/kaggle/input/default-specs/spectograms.npy', allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.822730Z","iopub.execute_input":"2024-02-25T22:10:29.823231Z","iopub.status.idle":"2024-02-25T22:10:29.827341Z","shell.execute_reply.started":"2024-02-25T22:10:29.823202Z","shell.execute_reply":"2024-02-25T22:10:29.826593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Делю по пациентам на трейн, валидацию и тест\npatients = 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-02-25T22:10:29.828557Z","iopub.execute_input":"2024-02-25T22:10:29.829035Z","iopub.status.idle":"2024-02-25T22:10:29.844291Z","shell.execute_reply.started":"2024-02-25T22:10:29.829007Z","shell.execute_reply":"2024-02-25T22:10:29.843154Z"},"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-02-25T22:10:29.845961Z","iopub.execute_input":"2024-02-25T22:10:29.846308Z","iopub.status.idle":"2024-02-25T22:10:29.898872Z","shell.execute_reply.started":"2024-02-25T22:10:29.846281Z","shell.execute_reply":"2024-02-25T22:10:29.897645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Функция, чтобы считать среднее и стандартное отклонение в цикле, потому мтодами numpy памяти не хватает\ndef online_mean_std(data):\n    n = 0\n    mean = 0\n    M2 = 0\n\n    for x in tqdm(data):\n        n = n + 1\n        x = np.nan_to_num(x)\n        delta = x - mean\n        mean = mean + delta/n\n        M2 = M2 + delta*(x - mean)\n\n    variance = M2/(n - 1)\n    return np.sqrt(variance.mean()), mean.mean()","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.904455Z","iopub.execute_input":"2024-02-25T22:10:29.904864Z","iopub.status.idle":"2024-02-25T22:10:29.912830Z","shell.execute_reply.started":"2024-02-25T22:10:29.904834Z","shell.execute_reply":"2024-02-25T22:10:29.911534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Считаю средние и ст.отклонения для каждого типа спектограмм\nif not load_pretrained:\n    means = []\n    stds = []\n    for el in ['LL', 'RL', 'LP', 'RP']:\n        res = np.concatenate([dct[sid][el][None, :, int(slos)//2: int(slos)//2+300] for sid, slos in zip(train_patients.spectrogram_id, train_patients.spectrogram_label_offset_seconds)], axis=0)\n        std, mean = online_mean_std(res)\n        means.append(mean)\n        stds.append(std)\n        del res\n        gc.collect()\n\n    norm_mean = np.array(means).reshape((4, 1, 1))\n    norm_std = np.array(stds).reshape((4, 1, 1))","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.915215Z","iopub.execute_input":"2024-02-25T22:10:29.915570Z","iopub.status.idle":"2024-02-25T22:10:29.926229Z","shell.execute_reply.started":"2024-02-25T22:10:29.915527Z","shell.execute_reply":"2024-02-25T22:10:29.924500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### [None, :, int(slos)//2: int(slos)//2+300] Нужно, чтобы по 10 минут из спек вырезать. Новую ось создаю, чтобы по ней конкатить спеки","metadata":{}},{"cell_type":"code","source":"def normalize(x):\n    '''[c, h, w]'''\n    return (x - norm_mean) / norm_std\n\nclass SpecDataset(Dataset):\n# В данных размер спеки 99 x 300    \n    def __init__(self, df, dct, img_size=(99, 300)):\n        self.df = df\n        self.dct = dct\n        self.image_size = img_size\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        spec = self.dct[self.df.iloc[index].spectrogram_id]\n        shift = self.df.iloc[index].spectrogram_label_offset_seconds\n        ll, rl, lp, rp = spec['LL'], spec['RL'], spec['LP'], spec['RP']\n        x = np.concatenate([ll[None, :, int(shift)//2: int(shift)//2+300], rl[None, :, int(shift)//2: int(shift)//2+300], lp[None, :, int(shift)//2: int(shift)//2+300], rp[None, :, int(shift)//2: int(shift)//2+300]], axis=0)\n        x = torch.from_numpy(normalize(x)).float()\n        x = torch.nan_to_num(x, 0)\n        transforms = Resize([self.image_size[0], self.image_size[1]])\n        x = transforms(x)        \n        return x","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.927827Z","iopub.execute_input":"2024-02-25T22:10:29.928181Z","iopub.status.idle":"2024-02-25T22:10:29.942347Z","shell.execute_reply.started":"2024-02-25T22:10:29.928152Z","shell.execute_reply":"2024-02-25T22:10:29.941192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Сеть представляет из себя просто ResNet блоки, которые уменьшают/увличивают ширину и высоту в два раза и увличивают/уменьшают число каналов в два раза","metadata":{}},{"cell_type":"code","source":"class ResNetBlock(nn.Module):\n    def __init__(self, in_channels, kernel_size, modify=False, bn=True):\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 = nn.ConvTranspose2d(in_channels=in_channels, out_channels=in_channels//2, stride=2, kernel_size=kernel_size, output_padding=1, 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=1, 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        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        x = self.conv(x)\n        x = self.rnb1(x)\n        x = self.rnb2(x)\n        x = self.rnb3(x)\n        x = self.rnb4(x)\n        x = self.rnb5(x)\n        x = self.rnb6(x)\n        x = self.rnb7(x)\n        x = self.rnb8(x)\n        return x\n    \nclass Decoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.rnb1 = ResNetBlock(4096, 3, modify='upsample')\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.conv = nn.Conv2d(16, 4, 3, 1, 3//2)\n\n    def forward(self, x):\n        x = self.rnb1(x)\n        x = self.rnb2(x)\n        x = self.rnb3(x)\n        x = self.rnb4(x)\n        x = self.rnb5(x)\n        x = self.rnb6(x)\n        x = self.rnb7(x)\n        x = self.rnb8(x)\n        x = self.conv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:29.944226Z","iopub.execute_input":"2024-02-25T22:10:29.944720Z","iopub.status.idle":"2024-02-25T22:10:29.979399Z","shell.execute_reply.started":"2024-02-25T22:10:29.944678Z","shell.execute_reply":"2024-02-25T22:10:29.978228Z"},"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-02-25T22:10:29.980980Z","iopub.execute_input":"2024-02-25T22:10:29.981399Z","iopub.status.idle":"2024-02-25T22:10:29.988039Z","shell.execute_reply.started":"2024-02-25T22:10:29.981352Z","shell.execute_reply":"2024-02-25T22:10:29.986858Z"},"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        x = batch.to(device)\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-02-25T22:10:29.989439Z","iopub.execute_input":"2024-02-25T22:10:29.990360Z","iopub.status.idle":"2024-02-25T22:10:29.999806Z","shell.execute_reply.started":"2024-02-25T22:10:29.990253Z","shell.execute_reply":"2024-02-25T22:10:29.998663Z"},"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.to(device)\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-02-25T22:10:30.002285Z","iopub.execute_input":"2024-02-25T22:10:30.002824Z","iopub.status.idle":"2024-02-25T22:10:30.014091Z","shell.execute_reply.started":"2024-02-25T22:10:30.002777Z","shell.execute_reply":"2024-02-25T22:10:30.012967Z"},"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-02-25T22:10:30.016386Z","iopub.execute_input":"2024-02-25T22:10:30.016786Z","iopub.status.idle":"2024-02-25T22:10:30.030119Z","shell.execute_reply.started":"2024-02-25T22:10:30.016753Z","shell.execute_reply":"2024-02-25T22:10:30.029008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    dataset_train = SpecDataset(train_patients, dct, img_size=(256, 256))\n    dataset_val = SpecDataset(val_patients, dct, img_size=(256, 256))\n    dataset_test = SpecDataset(test_patients, dct, img_size=(256, 256))\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-02-25T22:10:30.031645Z","iopub.execute_input":"2024-02-25T22:10:30.032022Z","iopub.status.idle":"2024-02-25T22:10:30.044879Z","shell.execute_reply.started":"2024-02-25T22:10:30.031991Z","shell.execute_reply":"2024-02-25T22:10:30.043639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# small, _ = torch.utils.data.random_split(dataset_train, [256, len(dataset_train) - 256])\n# small_val, _ = torch.utils.data.random_split(dataset_train, [256, len(dataset_train) - 256])\n# small_dataloader_train = DataLoader(\n#     dataset=small,\n#     batch_size=128,\n#     shuffle=True,\n#     drop_last=True\n# )\n\n# small_dataloader_val = DataLoader(\n#     dataset=small_val,\n#     batch_size=128,\n#     shuffle=False,\n#     drop_last=False\n# )","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:30.046225Z","iopub.execute_input":"2024-02-25T22:10:30.046612Z","iopub.status.idle":"2024-02-25T22:10:30.059170Z","shell.execute_reply.started":"2024-02-25T22:10:30.046548Z","shell.execute_reply":"2024-02-25T22:10:30.058056Z"},"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 = 100","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:30.061178Z","iopub.execute_input":"2024-02-25T22:10:30.061534Z","iopub.status.idle":"2024-02-25T22:10:37.953613Z","shell.execute_reply.started":"2024-02-25T22:10:30.061504Z","shell.execute_reply":"2024-02-25T22:10:37.952229Z"},"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-02-25T22:10:37.955120Z","iopub.execute_input":"2024-02-25T22:10:37.955607Z","iopub.status.idle":"2024-02-25T22:10:37.962394Z","shell.execute_reply.started":"2024-02-25T22:10:37.955537Z","shell.execute_reply":"2024-02-25T22:10:37.961062Z"},"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-02-25T22:10:37.964206Z","iopub.execute_input":"2024-02-25T22:10:37.964680Z","iopub.status.idle":"2024-02-25T22:10:37.973379Z","shell.execute_reply.started":"2024-02-25T22:10:37.964640Z","shell.execute_reply":"2024-02-25T22:10:37.972187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def nan_hook(self, inp, output):\n#     if not isinstance(output, tuple):\n#         outputs = [output]\n#     else:\n#         outputs = output\n\n#     for i, out in enumerate(outputs):\n#         nan_mask = torch.isnan(out)\n#         if nan_mask.any():\n#             print(\"In\", self.__class__.__name__)\n#             raise RuntimeError(f\"Found NAN in output {i} at indices: \", nan_mask.nonzero(), \"where:\", out[nan_mask.nonzero()[:, 0].unique(sorted=True)])\n\n# for submodule in model.modules():\n#     submodule.register_forward_hook(nan_hook)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:37.975012Z","iopub.execute_input":"2024-02-25T22:10:37.975474Z","iopub.status.idle":"2024-02-25T22:10:37.991857Z","shell.execute_reply.started":"2024-02-25T22:10:37.975421Z","shell.execute_reply":"2024-02-25T22:10:37.990230Z"},"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-02-25T22:10:37.993780Z","iopub.execute_input":"2024-02-25T22:10:37.994224Z","iopub.status.idle":"2024-02-25T22:10:38.002048Z","shell.execute_reply.started":"2024-02-25T22:10:37.994183Z","shell.execute_reply":"2024-02-25T22:10:38.000656Z"},"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-02-25T22:10:38.004193Z","iopub.execute_input":"2024-02-25T22:10:38.004838Z","iopub.status.idle":"2024-02-25T22:10:38.012587Z","shell.execute_reply.started":"2024-02-25T22:10:38.004780Z","shell.execute_reply":"2024-02-25T22:10:38.011318Z"},"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-02-25T22:10:38.014308Z","iopub.execute_input":"2024-02-25T22:10:38.014817Z","iopub.status.idle":"2024-02-25T22:10:38.023441Z","shell.execute_reply.started":"2024-02-25T22:10:38.014767Z","shell.execute_reply":"2024-02-25T22:10:38.022336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Что-то типа Аркфейса","metadata":{}},{"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":{"iopub.status.busy":"2024-02-25T22:10:38.025843Z","iopub.execute_input":"2024-02-25T22:10:38.026407Z","iopub.status.idle":"2024-02-25T22:10:38.038303Z","shell.execute_reply.started":"2024-02-25T22:10:38.026362Z","shell.execute_reply":"2024-02-25T22:10:38.036896Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:38.051404Z","iopub.execute_input":"2024-02-25T22:10:38.051824Z","iopub.status.idle":"2024-02-25T22:10:38.063796Z","shell.execute_reply.started":"2024-02-25T22:10:38.051792Z","shell.execute_reply":"2024-02-25T22:10:38.062187Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:38.066192Z","iopub.execute_input":"2024-02-25T22:10:38.067296Z","iopub.status.idle":"2024-02-25T22:10:49.134190Z","shell.execute_reply.started":"2024-02-25T22:10:38.067250Z","shell.execute_reply":"2024-02-25T22:10:49.132929Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.135840Z","iopub.execute_input":"2024-02-25T22:10:49.136283Z","iopub.status.idle":"2024-02-25T22:10:49.144583Z","shell.execute_reply.started":"2024-02-25T22:10:49.136244Z","shell.execute_reply":"2024-02-25T22:10:49.142982Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.146909Z","iopub.execute_input":"2024-02-25T22:10:49.149027Z","iopub.status.idle":"2024-02-25T22:10:49.194958Z","shell.execute_reply.started":"2024-02-25T22:10:49.148981Z","shell.execute_reply":"2024-02-25T22:10:49.193882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def get_acc_fpr_tpr(threshold, df, dists_matrix):\n#     preds = np.less(dists_matrix, threshold)\n#     targets = []\n    \n#     for patient_id in df.patient_id:\n#         targets.append((df.patient_id == patient_id).values.reshape((1,-1)))\n#     targets = np.concatenate(targets, axis=0)\n    \n#     tp = np.sum(np.logical_and(preds, targets), axis=1)\n#     fp = np.sum(np.logical_and(preds, np.logical_not(targets)), axis=1)\n#     tn = np.sum(np.logical_and(np.logical_not(preds), np.logical_not(targets)), axis=1)\n#     fn = np.sum(np.logical_and(np.logical_not(preds), targets), axis=1)\n\n#     tpr = np.where(tp + fn==0, 0, tp / (tp + fn))\n#     fpr = np.where(fp + tn==0, 0, fp / (fp + tn))\n#     acc = float(tp + tn) / dist.size\n#     return tpr.mean(), fpr.mean(), acc.mean()","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.197265Z","iopub.execute_input":"2024-02-25T22:10:49.197721Z","iopub.status.idle":"2024-02-25T22:10:49.205096Z","shell.execute_reply.started":"2024-02-25T22:10:49.197678Z","shell.execute_reply":"2024-02-25T22:10:49.204008Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.206550Z","iopub.execute_input":"2024-02-25T22:10:49.207317Z","iopub.status.idle":"2024-02-25T22:10:49.219664Z","shell.execute_reply.started":"2024-02-25T22:10:49.207277Z","shell.execute_reply":"2024-02-25T22:10:49.218535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def find_opt_threshold(df, dists_matrix, thresholds):\n#     for th in tqdm(thresholds):\n#         tpr_mean = []\n#         fpr_mean = []\n#         acc_mean = []\n#         for i in range(len(df)):\n#             tpr, fpr, acc = get_acc_fpr_tpr(th, df, dists_matrix)\n#             tpr_mean.append(tpr)\n#             fpr_mean.append(fpr)\n#             acc_mean.append(acc)\n#     return tpr_mean, fpr_mean, acc_mean","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.220985Z","iopub.execute_input":"2024-02-25T22:10:49.221650Z","iopub.status.idle":"2024-02-25T22:10:49.229393Z","shell.execute_reply.started":"2024-02-25T22:10:49.221606Z","shell.execute_reply":"2024-02-25T22:10:49.228296Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.230892Z","iopub.execute_input":"2024-02-25T22:10:49.231710Z","iopub.status.idle":"2024-02-25T22:10:49.240989Z","shell.execute_reply.started":"2024-02-25T22:10:49.231678Z","shell.execute_reply":"2024-02-25T22:10:49.239478Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.242845Z","iopub.execute_input":"2024-02-25T22:10:49.243527Z","iopub.status.idle":"2024-02-25T22:10:49.266613Z","shell.execute_reply.started":"2024-02-25T22:10:49.243495Z","shell.execute_reply":"2024-02-25T22:10:49.265345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 30\nthresholds_train[i], tpr_train[i], fpr_train[i], acc_train[i]","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.267972Z","iopub.execute_input":"2024-02-25T22:10:49.268308Z","iopub.status.idle":"2024-02-25T22:10:49.276204Z","shell.execute_reply.started":"2024-02-25T22:10:49.268279Z","shell.execute_reply":"2024-02-25T22:10:49.274989Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.278059Z","iopub.execute_input":"2024-02-25T22:10:49.278483Z","iopub.status.idle":"2024-02-25T22:10:49.877173Z","shell.execute_reply.started":"2024-02-25T22:10:49.278452Z","shell.execute_reply":"2024-02-25T22:10:49.875981Z"},"trusted":true},"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":{"iopub.status.busy":"2024-02-25T22:10:49.878438Z","iopub.execute_input":"2024-02-25T22:10:49.878791Z","iopub.status.idle":"2024-02-25T22:10:49.885252Z","shell.execute_reply.started":"2024-02-25T22:10:49.878762Z","shell.execute_reply":"2024-02-25T22:10:49.883840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:\n    tpr_val, fpr_val, acc_val","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.887022Z","iopub.execute_input":"2024-02-25T22:10:49.887479Z","iopub.status.idle":"2024-02-25T22:10:49.894569Z","shell.execute_reply.started":"2024-02-25T22:10:49.887445Z","shell.execute_reply":"2024-02-25T22:10:49.893507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if get_dists:    \n    tpr_test, fpr_test, acc_test","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.896413Z","iopub.execute_input":"2024-02-25T22:10:49.896864Z","iopub.status.idle":"2024-02-25T22:10:49.905250Z","shell.execute_reply.started":"2024-02-25T22:10:49.896825Z","shell.execute_reply":"2024-02-25T22:10:49.904132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Поиск аномальных наблюдений или пациентов (для которых loss большой)","metadata":{}},{"cell_type":"code","source":"train_patients['loss'] = losses_train\nval_patients['loss'] = losses_val\ntest_patients['loss'] = losses_test","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.906783Z","iopub.execute_input":"2024-02-25T22:10:49.907260Z","iopub.status.idle":"2024-02-25T22:10:49.918238Z","shell.execute_reply.started":"2024-02-25T22:10:49.907228Z","shell.execute_reply":"2024-02-25T22:10:49.916840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_author_classes(df, high_level_of_agreement=0.6, edge_std=0.1, proto_std=0.1):\n    '''\n    Размечаем наш датасет по 'непонятности' пациентов по тексту авторов:\n    We call segments where there are high levels of agreement “idealized” patterns.\n    Cases where ~1/2 of experts give a label as “other” and ~1/2 give one of the remaining five labels, we call “proto patterns”.\n    Cases where experts are approximately split between 2 of the 5 named patterns, we call “edge cases”.\n    \n    df — Наш датафрейм с колонками таргетов;\n    high_level_of_agreement — Уверенность в диагнозе, после которой мы относим наблюдение к классу idealized.\n    edge_std — Стандартное отклонение от 0.5, которое будет использоваться для классификации наблюдения как edge.\n                То есть две уверенности из 5 должны быть в интервале [0.5 - edge_std; 0.5 + edge_std]\n    proto_std — Стандартное отклонение от 0.5, которое будет использоваться для классификации наблюдения как proto.\n                То есть уверенность класса 'other_vote' и какого-то еще одного должны быть в интервале [0.5 - proto_std; 0.5 + proto_std]\n                \n    Все, что не попадает в эти интервалы, будет классифицироваться как класс other\n    '''\n    df['authors_class'] = 'other'\n    # Убрал other из idealized кейса, хотя формально по тексту описания подходит, но 'idealized other' как-то стремно звучит.\n    idealized = (df.iloc[:,-7:-2] > high_level_of_agreement).any(axis=1)\n    proto = ((df.iloc[:, -7:-2] > 0.5 - proto_std) & (df.iloc[:, -7:-2] < 0.5 + proto_std)).any(axis=1) & (df['other_vote'] > 0.5 - proto_std) & (df['other_vote'] < 0.5 + proto_std)\n    edge = ((df.iloc[:, -7:-2] > 0.5 - edge_std) & (df.iloc[:, -7:-2] < 0.5 + edge_std)).sum(axis=1) == 2\n    df.loc[idealized, 'authors_class'] = 'idealized'\n    df.loc[proto, 'authors_class'] = 'proto'\n    df.loc[edge, 'authors_class'] = 'edge'\n    return df","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.920235Z","iopub.execute_input":"2024-02-25T22:10:49.920664Z","iopub.status.idle":"2024-02-25T22:10:49.932344Z","shell.execute_reply.started":"2024-02-25T22:10:49.920621Z","shell.execute_reply":"2024-02-25T22:10:49.931127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_patients = add_author_classes(train_patients)\nval_patients = add_author_classes(val_patients)\ntest_patients = add_author_classes(test_patients)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:49.933842Z","iopub.execute_input":"2024-02-25T22:10:49.934474Z","iopub.status.idle":"2024-02-25T22:10:50.001537Z","shell.execute_reply.started":"2024-02-25T22:10:49.934440Z","shell.execute_reply":"2024-02-25T22:10:50.000141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def group(df):\n    by_patient = df.groupby('patient_id').mean('loss')['loss'].reset_index()\n    by_spec = df.groupby('spectrogram_id').mean('loss')['loss'].reset_index()\n    by_illness = df.groupby('expert_consensus').mean('loss')['loss'].reset_index()\n    by_authors = df.groupby('authors_class').mean('loss')['loss'].reset_index()\n    return by_patient, by_spec, by_illness, by_authors","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.002946Z","iopub.execute_input":"2024-02-25T22:10:50.003282Z","iopub.status.idle":"2024-02-25T22:10:50.009921Z","shell.execute_reply.started":"2024-02-25T22:10:50.003252Z","shell.execute_reply":"2024-02-25T22:10:50.009075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_patient_train, by_spec_train, by_illness_train, by_authors_train = group(train_patients)\nby_patient_val, by_spec_val, by_illness_val, by_authors_val = group(val_patients)\nby_patient_test, by_spec_test, by_illness_test, by_authors_test = group(test_patients)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.011099Z","iopub.execute_input":"2024-02-25T22:10:50.012126Z","iopub.status.idle":"2024-02-25T22:10:50.132125Z","shell.execute_reply.started":"2024-02-25T22:10:50.012084Z","shell.execute_reply":"2024-02-25T22:10:50.130681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_illness_train","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.133707Z","iopub.execute_input":"2024-02-25T22:10:50.134084Z","iopub.status.idle":"2024-02-25T22:10:50.145694Z","shell.execute_reply.started":"2024-02-25T22:10:50.134052Z","shell.execute_reply":"2024-02-25T22:10:50.144500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_illness_val","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.147475Z","iopub.execute_input":"2024-02-25T22:10:50.147852Z","iopub.status.idle":"2024-02-25T22:10:50.160121Z","shell.execute_reply.started":"2024-02-25T22:10:50.147821Z","shell.execute_reply":"2024-02-25T22:10:50.158855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_illness_test","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.161535Z","iopub.execute_input":"2024-02-25T22:10:50.162218Z","iopub.status.idle":"2024-02-25T22:10:50.175685Z","shell.execute_reply.started":"2024-02-25T22:10:50.162185Z","shell.execute_reply":"2024-02-25T22:10:50.174383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_authors_train","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.177136Z","iopub.execute_input":"2024-02-25T22:10:50.178239Z","iopub.status.idle":"2024-02-25T22:10:50.190947Z","shell.execute_reply.started":"2024-02-25T22:10:50.178196Z","shell.execute_reply":"2024-02-25T22:10:50.189658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_authors_val","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.193480Z","iopub.execute_input":"2024-02-25T22:10:50.193930Z","iopub.status.idle":"2024-02-25T22:10:50.205907Z","shell.execute_reply.started":"2024-02-25T22:10:50.193895Z","shell.execute_reply":"2024-02-25T22:10:50.204777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"by_authors_test","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.207385Z","iopub.execute_input":"2024-02-25T22:10:50.207756Z","iopub.status.idle":"2024-02-25T22:10:50.222367Z","shell.execute_reply.started":"2024-02-25T22:10:50.207726Z","shell.execute_reply":"2024-02-25T22:10:50.220899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_patients['entropy'].hist(bins=20)\ndef entropy_mapper(entropy):\n    if entropy == 0:\n        return '0'\n    elif entropy <= 0.75:\n        return '(0: 0.75]'\n    elif entropy <= 1.25:\n        return '(0.75: 1.25]'\n    else:\n        return '(1.25: )'\ntrain_patients['entropy_group'] = train_patients['entropy'].map(lambda x: entropy_mapper(x))\nval_patients['entropy_group'] = val_patients['entropy'].map(lambda x: entropy_mapper(x))\ntest_patients['entropy_group'] = test_patients['entropy'].map(lambda x: entropy_mapper(x))","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.224243Z","iopub.execute_input":"2024-02-25T22:10:50.225337Z","iopub.status.idle":"2024-02-25T22:10:50.588162Z","shell.execute_reply.started":"2024-02-25T22:10:50.225300Z","shell.execute_reply":"2024-02-25T22:10:50.586836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_patients['loss'].quantile([0.65, 0.75, 0.85, 0.9, 0.95])","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:24:25.079895Z","iopub.execute_input":"2024-02-25T22:24:25.080426Z","iopub.status.idle":"2024-02-25T22:24:25.096835Z","shell.execute_reply.started":"2024-02-25T22:24:25.080384Z","shell.execute_reply":"2024-02-25T22:24:25.095469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def loss_mapper(loss):\n    if loss <= 0.000014:\n        return '<0.65'\n    elif loss <= 0.00002:\n        return '0.65<=x<0.75'\n    elif loss <= 0.0022:\n        return '0.75<=x<0.85'\n    elif loss <= 0.064518:\n        return '0.85<=x<0.9'\n    elif loss <= 0.935622:\n        return '0.9<=x<0.95'\n    else:\n        return '0.95<=x'\ntrain_patients['loss_group'] = train_patients['loss'].map(lambda x: loss_mapper(x))\nval_patients['loss_group'] = val_patients['loss'].map(lambda x: loss_mapper(x))\ntest_patients['loss_group'] = test_patients['loss'].map(lambda x: loss_mapper(x))","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:35:15.360776Z","iopub.execute_input":"2024-02-25T22:35:15.361298Z","iopub.status.idle":"2024-02-25T22:35:15.434936Z","shell.execute_reply.started":"2024-02-25T22:35:15.361263Z","shell.execute_reply":"2024-02-25T22:35:15.433818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### UMAP","metadata":{}},{"cell_type":"code","source":"umap_obj = umap.UMAP(n_components=2, random_state=42)\numap_data_train = umap_obj.fit_transform(embs_train)\numap_data_val = umap_obj.transform(embs_val)\numap_data_test = umap_obj.transform(embs_test)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:10:50.589552Z","iopub.execute_input":"2024-02-25T22:10:50.590065Z","iopub.status.idle":"2024-02-25T22:14:54.491094Z","shell.execute_reply.started":"2024-02-25T22:10:50.590032Z","shell.execute_reply":"2024-02-25T22:14:54.489615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_group(umap_data_train, umap_data_val, umap_data_test, df_train, df_val, df_test, group_col, hue_order):\n    fig, axes = plt.subplots(1, 3, sharex=True, figsize=(16,8))\n    sns.scatterplot(ax=axes[0], x=umap_data_train[:,0], y=umap_data_train[:,1], hue=df_train[group_col], hue_order=hue_order, s=7)\n    axes[0].set_title('Train')\n\n    sns.scatterplot(ax=axes[1], x=umap_data_val[:,0], y=umap_data_val[:,1], hue=df_val[group_col], hue_order=hue_order, s=7)\n    axes[1].set_title('Val')\n\n    sns.scatterplot(ax=axes[2], x=umap_data_test[:,0], y=umap_data_test[:,1], hue=df_test[group_col], hue_order=hue_order, s=7)\n    axes[2].set_title('Test')\n    \n    plt.setp(axes[0].get_legend().get_texts(), fontsize=7)\n    plt.setp(axes[1].get_legend().get_texts(), fontsize=7)\n    plt.setp(axes[2].get_legend().get_texts(), fontsize=7)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:14:54.493077Z","iopub.execute_input":"2024-02-25T22:14:54.493488Z","iopub.status.idle":"2024-02-25T22:14:54.507162Z","shell.execute_reply.started":"2024-02-25T22:14:54.493443Z","shell.execute_reply":"2024-02-25T22:14:54.506118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cons = ['GPD', 'LRDA', 'Seizure', 'Other', 'GRDA', 'LPD']\nplot_group(umap_data_train, umap_data_val, umap_data_test, train_patients, val_patients, test_patients, group_col='expert_consensus', hue_order=hue_order_cons)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:14:54.508797Z","iopub.execute_input":"2024-02-25T22:14:54.509410Z","iopub.status.idle":"2024-02-25T22:15:02.931254Z","shell.execute_reply.started":"2024-02-25T22:14:54.509378Z","shell.execute_reply":"2024-02-25T22:15:02.930341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cls = ['idealized', 'proto', 'edge', 'other'] \nplot_group(umap_data_train, umap_data_val, umap_data_test, train_patients, val_patients, test_patients, group_col='authors_class', hue_order=hue_order_cls)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:15:02.932621Z","iopub.execute_input":"2024-02-25T22:15:02.933470Z","iopub.status.idle":"2024-02-25T22:15:06.831779Z","shell.execute_reply.started":"2024-02-25T22:15:02.933436Z","shell.execute_reply":"2024-02-25T22:15:06.830707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_patients['entropy'].hist(bins=20)\ndef entropy_mapper(entropy):\n    if entropy == 0:\n        return '0'\n    elif entropy <= 0.75:\n        return '(0: 0.75]'\n    elif entropy <= 1.25:\n        return '(0.75: 1.25]'\n    else:\n        return '(1.25: )'\ntrain_patients['entropy_group'] = train_patients['entropy'].map(lambda x: entropy_mapper(x))\nval_patients['entropy_group'] = val_patients['entropy'].map(lambda x: entropy_mapper(x))\ntest_patients['entropy_group'] = test_patients['entropy'].map(lambda x: entropy_mapper(x))","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:15:06.833421Z","iopub.execute_input":"2024-02-25T22:15:06.833815Z","iopub.status.idle":"2024-02-25T22:15:07.184108Z","shell.execute_reply.started":"2024-02-25T22:15:06.833781Z","shell.execute_reply":"2024-02-25T22:15:07.182756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_ent = ['0', '(0: 0.75]', '(0.75: 1.25]', '(1.25: )'] \nplot_group(umap_data_train, umap_data_val, umap_data_test, train_patients, val_patients, test_patients, group_col='entropy_group', hue_order=hue_order_ent)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:15:07.185689Z","iopub.execute_input":"2024-02-25T22:15:07.186153Z","iopub.status.idle":"2024-02-25T22:15:11.085909Z","shell.execute_reply.started":"2024-02-25T22:15:07.186110Z","shell.execute_reply":"2024-02-25T22:15:11.084779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_loss = ['<0.65', '0.65<=x<0.75', '0.75<=x<0.85', '0.85<=x<0.9', '0.9<=x<0.95', '0.95<=x'] \nplot_group(umap_data_train, umap_data_val, umap_data_test, train_patients, val_patients, test_patients, group_col='loss_group', hue_order=hue_order_loss)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:45:21.904337Z","iopub.execute_input":"2024-02-25T22:45:21.905801Z","iopub.status.idle":"2024-02-25T22:45:30.330182Z","shell.execute_reply.started":"2024-02-25T22:45:21.905759Z","shell.execute_reply":"2024-02-25T22:45:30.328899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Remove Other label","metadata":{}},{"cell_type":"code","source":"train_patients_no_other = train_patients[train_patients['authors_class'] != 'other'].copy()\nval_patients_no_other = val_patients[val_patients['authors_class'] != 'other'].copy()\ntest_patients_no_other = test_patients[test_patients['authors_class'] != 'other'].copy()","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:38:17.341545Z","iopub.execute_input":"2024-02-25T22:38:17.342022Z","iopub.status.idle":"2024-02-25T22:38:17.399443Z","shell.execute_reply.started":"2024-02-25T22:38:17.341981Z","shell.execute_reply":"2024-02-25T22:38:17.398207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"umap_obj = umap.UMAP(n_components=2, random_state=42)\numap_data_train_no_other = umap_obj.fit_transform(embs_train[train_patients_no_other.index])\numap_data_val_no_other = umap_obj.transform(embs_val[val_patients_no_other.index])\numap_data_test_no_other = umap_obj.transform(embs_test[test_patients_no_other.index])","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:15:11.156709Z","iopub.execute_input":"2024-02-25T22:15:11.157289Z","iopub.status.idle":"2024-02-25T22:17:25.020246Z","shell.execute_reply.started":"2024-02-25T22:15:11.157257Z","shell.execute_reply":"2024-02-25T22:17:25.018878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cons = ['GPD', 'LRDA', 'Seizure', 'Other', 'GRDA', 'LPD']\nplot_group(umap_data_train_no_other, umap_data_val_no_other, umap_data_test_no_other, train_patients_no_other, val_patients_no_other, test_patients_no_other, group_col='expert_consensus', hue_order=hue_order_cons)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:17:25.022364Z","iopub.execute_input":"2024-02-25T22:17:25.022861Z","iopub.status.idle":"2024-02-25T22:17:31.709183Z","shell.execute_reply.started":"2024-02-25T22:17:25.022817Z","shell.execute_reply":"2024-02-25T22:17:31.708184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cls = ['idealized', 'proto', 'edge'] \nplot_group(umap_data_train_no_other, umap_data_val_no_other, umap_data_test_no_other, train_patients_no_other, val_patients_no_other, test_patients_no_other, group_col='authors_class', hue_order=hue_order_cls)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:17:31.710685Z","iopub.execute_input":"2024-02-25T22:17:31.711301Z","iopub.status.idle":"2024-02-25T22:17:35.228749Z","shell.execute_reply.started":"2024-02-25T22:17:31.711267Z","shell.execute_reply":"2024-02-25T22:17:35.227649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_ent = ['0', '(0: 0.75]', '(0.75: 1.25]', '(1.25: )'] \nplot_group(umap_data_train_no_other, umap_data_val_no_other, umap_data_test_no_other, train_patients_no_other, val_patients_no_other, test_patients_no_other, group_col='entropy_group', hue_order=hue_order_ent)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:17:35.230428Z","iopub.execute_input":"2024-02-25T22:17:35.231158Z","iopub.status.idle":"2024-02-25T22:17:38.762970Z","shell.execute_reply.started":"2024-02-25T22:17:35.231105Z","shell.execute_reply":"2024-02-25T22:17:38.761592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_loss = ['<0.65', '0.65<=x<0.75', '0.75<=x<0.85', '0.85<=x<0.9', '0.9<=x<0.95', '0.95<=x'] \nplot_group(umap_data_train_no_other, umap_data_val_no_other, umap_data_test_no_other, train_patients_no_other, val_patients_no_other, test_patients_no_other, group_col='loss_group', hue_order=hue_order_loss)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:44:26.635829Z","iopub.execute_input":"2024-02-25T22:44:26.637199Z","iopub.status.idle":"2024-02-25T22:44:33.577625Z","shell.execute_reply.started":"2024-02-25T22:44:26.637140Z","shell.execute_reply":"2024-02-25T22:44:33.576652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Remove high entropy","metadata":{}},{"cell_type":"code","source":"train_patients_no_ent = train_patients[train_patients['entropy_group'] != '(1.25: )'].copy()\nval_patients_no_ent = val_patients[val_patients['entropy_group'] != '(1.25: )'].copy()\ntest_patients_no_ent = test_patients[test_patients['entropy_group'] != '(1.25: )'].copy()","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:38:23.105793Z","iopub.execute_input":"2024-02-25T22:38:23.106371Z","iopub.status.idle":"2024-02-25T22:38:23.182367Z","shell.execute_reply.started":"2024-02-25T22:38:23.106323Z","shell.execute_reply":"2024-02-25T22:38:23.180805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"umap_obj = umap.UMAP(n_components=2, random_state=42)\numap_data_train_no_ent = umap_obj.fit_transform(embs_train[train_patients_no_ent.index])\numap_data_val_no_ent = umap_obj.transform(embs_val[val_patients_no_ent.index])\numap_data_test_no_ent = umap_obj.transform(embs_test[test_patients_no_ent.index])","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:17:38.827158Z","iopub.execute_input":"2024-02-25T22:17:38.827860Z","iopub.status.idle":"2024-02-25T22:20:46.378568Z","shell.execute_reply.started":"2024-02-25T22:17:38.827817Z","shell.execute_reply":"2024-02-25T22:20:46.377255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cons = ['GPD', 'LRDA', 'Seizure', 'Other', 'GRDA', 'LPD']\nplot_group(umap_data_train_no_ent, umap_data_val_no_ent, umap_data_test_no_ent, train_patients_no_ent, val_patients_no_ent, test_patients_no_ent, group_col='expert_consensus', hue_order=hue_order_cons)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:20:46.380258Z","iopub.execute_input":"2024-02-25T22:20:46.380630Z","iopub.status.idle":"2024-02-25T22:20:54.884529Z","shell.execute_reply.started":"2024-02-25T22:20:46.380598Z","shell.execute_reply":"2024-02-25T22:20:54.883399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cls = ['idealized', 'proto', 'edge'] \nplot_group(umap_data_train_no_ent, umap_data_val_no_ent, umap_data_test_no_ent, train_patients_no_ent, val_patients_no_ent, test_patients_no_ent, group_col='authors_class', hue_order=hue_order_cls)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:20:54.885838Z","iopub.execute_input":"2024-02-25T22:20:54.886189Z","iopub.status.idle":"2024-02-25T22:20:59.294339Z","shell.execute_reply.started":"2024-02-25T22:20:54.886158Z","shell.execute_reply":"2024-02-25T22:20:59.293124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_ent = ['0', '(0: 0.75]', '(0.75: 1.25]'] \nplot_group(umap_data_train_no_ent, umap_data_val_no_ent, umap_data_test_no_ent, train_patients_no_ent, val_patients_no_ent, test_patients_no_ent, group_col='entropy_group', hue_order=hue_order_ent)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:20:59.295863Z","iopub.execute_input":"2024-02-25T22:20:59.296279Z","iopub.status.idle":"2024-02-25T22:21:03.492676Z","shell.execute_reply.started":"2024-02-25T22:20:59.296243Z","shell.execute_reply":"2024-02-25T22:21:03.490963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_loss = ['<0.65', '0.65<=x<0.75', '0.75<=x<0.85', '0.85<=x<0.9', '0.9<=x<0.95', '0.95<=x'] \nplot_group(umap_data_train_no_ent, umap_data_val_no_ent, umap_data_test_no_ent, train_patients_no_ent, val_patients_no_ent, test_patients_no_ent, group_col='loss_group', hue_order=hue_order_loss)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T22:38:26.631796Z","iopub.execute_input":"2024-02-25T22:38:26.632547Z","iopub.status.idle":"2024-02-25T22:38:34.968429Z","shell.execute_reply.started":"2024-02-25T22:38:26.632511Z","shell.execute_reply":"2024-02-25T22:38:34.967258Z"},"trusted":true},"execution_count":null,"outputs":[]}]}