{"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":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7564328,"sourceType":"datasetVersion","datasetId":4404591},{"sourceId":7673559,"sourceType":"datasetVersion","datasetId":4454469}],"dockerImageVersionId":30646,"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\n# import timm\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":{"execution":{"iopub.status.busy":"2024-02-21T20:19:12.498314Z","iopub.execute_input":"2024-02-21T20:19:12.498933Z","iopub.status.idle":"2024-02-21T20:20:03.427805Z","shell.execute_reply.started":"2024-02-21T20:19:12.498868Z","shell.execute_reply":"2024-02-21T20:20:03.426133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_pretrained = True\nload_embs = True\ntest_eval = False\nget_dists = False\nfind_threshold = False","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:03.431237Z","iopub.execute_input":"2024-02-21T20:20:03.432395Z","iopub.status.idle":"2024-02-21T20:20:03.439314Z","shell.execute_reply.started":"2024-02-21T20:20:03.432340Z","shell.execute_reply":"2024-02-21T20:20:03.437547Z"},"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-21T20:20:03.441647Z","iopub.execute_input":"2024-02-21T20:20:03.442227Z","iopub.status.idle":"2024-02-21T20:20:03.462045Z","shell.execute_reply.started":"2024-02-21T20:20:03.442178Z","shell.execute_reply":"2024-02-21T20:20:03.460501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\ntrain_df.iloc[:,-6:] = train_df.iloc[:,-6:].values / train_df.iloc[:,-6:].sum(axis=1).values.reshape((-1, 1))\ntrain_df = add_author_classes(train_df)","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:36:26.160999Z","iopub.execute_input":"2024-02-21T20:36:26.162692Z","iopub.status.idle":"2024-02-21T20:36:26.522018Z","shell.execute_reply.started":"2024-02-21T20:36:26.162628Z","shell.execute_reply":"2024-02-21T20:36:26.520279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    dct = np.load('/kaggle/input/default-specs/spectograms.npy', allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:03.912051Z","iopub.execute_input":"2024-02-21T20:20:03.912465Z","iopub.status.idle":"2024-02-21T20:20:03.919655Z","shell.execute_reply.started":"2024-02-21T20:20:03.912433Z","shell.execute_reply":"2024-02-21T20:20:03.917623Z"},"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-02-21T20:20:03.921479Z","iopub.execute_input":"2024-02-21T20:20:03.922275Z","iopub.status.idle":"2024-02-21T20:20:03.939028Z","shell.execute_reply.started":"2024-02-21T20:20:03.922233Z","shell.execute_reply":"2024-02-21T20:20:03.937507Z"},"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-21T20:37:00.952445Z","iopub.execute_input":"2024-02-21T20:37:00.952912Z","iopub.status.idle":"2024-02-21T20:37:01.012032Z","shell.execute_reply.started":"2024-02-21T20:37:00.952881Z","shell.execute_reply":"2024-02-21T20:37:01.010545Z"},"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-21T20:20:03.994194Z","iopub.execute_input":"2024-02-21T20:20:03.994607Z","iopub.status.idle":"2024-02-21T20:20:04.005008Z","shell.execute_reply.started":"2024-02-21T20:20:03.994578Z","shell.execute_reply":"2024-02-21T20:20:04.003015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Считаю средние и ст.отклонения для каждого типа спектограмм\nif not load_pretrained or not load_embs:\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-21T20:20:04.007562Z","iopub.execute_input":"2024-02-21T20:20:04.008138Z","iopub.status.idle":"2024-02-21T20:20:04.028888Z","shell.execute_reply.started":"2024-02-21T20:20:04.008102Z","shell.execute_reply":"2024-02-21T20:20:04.027484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patients_to_index_mapper = {i:j for i, j in zip(patients_train, np.arange(len(patients_train)))}\nindex_to_patients_mapper = {j:i for i, j in zip(patients_train, np.arange(len(patients_train)))}","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.036275Z","iopub.execute_input":"2024-02-21T20:20:04.037032Z","iopub.status.idle":"2024-02-21T20:20:04.051020Z","shell.execute_reply.started":"2024-02-21T20:20:04.036982Z","shell.execute_reply":"2024-02-21T20:20:04.049263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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), evaluate=False):\n        self.df = df\n        self.dct = dct\n        self.image_size = img_size\n        self.evaluate = evaluate\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        if not self.evaluate:\n            patient_index = patients_to_index_mapper[self.df.iloc[index].patient_id]\n            y = torch.tensor(patient_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        if not self.evaluate:\n            return x, y\n        else:\n            return x","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.053383Z","iopub.execute_input":"2024-02-21T20:20:04.054304Z","iopub.status.idle":"2024-02-21T20:20:04.076681Z","shell.execute_reply.started":"2024-02-21T20:20:04.054248Z","shell.execute_reply":"2024-02-21T20:20:04.073979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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_classes=len(patients_train)):\n        super().__init__()\n        self.n_classes = n_classes\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        self.linear = nn.Linear(4096, self.n_classes)\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        x = self.linear(x.squeeze())\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.079188Z","iopub.execute_input":"2024-02-21T20:20:04.080366Z","iopub.status.idle":"2024-02-21T20:20:04.117649Z","shell.execute_reply.started":"2024-02-21T20:20:04.080313Z","shell.execute_reply":"2024-02-21T20:20:04.115209Z"},"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, y = batch\n        x = x.to(device)\n        y = y.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            out = model(x)\n            loss = loss_fn(out, y)\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-21T20:20:04.120452Z","iopub.execute_input":"2024-02-21T20:20:04.121967Z","iopub.status.idle":"2024-02-21T20:20:04.136713Z","shell.execute_reply.started":"2024-02-21T20:20:04.121907Z","shell.execute_reply":"2024-02-21T20:20:04.135008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_list(lst, name):\n    with open(f'{name}.txt', 'w') as fp:\n        for item in lst:\n            fp.write(\"%s\\n\" % item)\n            \ndef load_list(name):\n    lst = []\n    with open(name, 'r') as fp:\n        for line in fp:\n            x = line[:-1]\n            lst.append(x)\n    return lst","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.138462Z","iopub.execute_input":"2024-02-21T20:20:04.138915Z","iopub.status.idle":"2024-02-21T20:20:04.157527Z","shell.execute_reply.started":"2024-02-21T20:20:04.138878Z","shell.execute_reply":"2024-02-21T20:20:04.156233Z"},"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=4, scaler=scaler):\n        losses_train = []\n#         losses_val = []\n        best_loss_train = 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 epoch >= 50 and epoch % 10 == 0:\n                torch.save(model.state_dict(), f'best_model_{epoch}.pth')\n                torch.save(optimizer, 'optimizer.pth')\n#             if best_loss_train > loss_train:\n#                 torch.save(model.state_dict(), 'best_model.pth')\n#                 torch.save(optimizer, 'optimizer.pth')\n#                 best_loss_train = loss_train\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            save_list(losses_train, 'losses_train')\n#             save_list(losses_val, 'losses_val')\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-21T20:20:04.160160Z","iopub.execute_input":"2024-02-21T20:20:04.160740Z","iopub.status.idle":"2024-02-21T20:20:04.176476Z","shell.execute_reply.started":"2024-02-21T20:20:04.160697Z","shell.execute_reply":"2024-02-21T20:20:04.175004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    dataset_train = SpecDataset(train_patients, dct, img_size=(99, 300))\n    dataset_val = SpecDataset(val_patients, dct, img_size=(99, 300), evaluate=True)\n    dataset_test = SpecDataset(test_patients, dct, img_size=(99, 300), evaluate=True)\n\n    dataloader_train = DataLoader(\n        dataset=dataset_train,\n        batch_size=256,\n        shuffle=True,\n        drop_last=True\n    )\n\n    dataloader_val = DataLoader(\n        dataset=dataset_val,\n        batch_size=256,\n        shuffle=False,\n        drop_last=False\n    )\n\n    dataloader_test = DataLoader(\n        dataset=dataset_test,\n        batch_size=256,\n        shuffle=False,\n        drop_last=False\n    )","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.178224Z","iopub.execute_input":"2024-02-21T20:20:04.179701Z","iopub.status.idle":"2024-02-21T20:20:04.196765Z","shell.execute_reply.started":"2024-02-21T20:20:04.179631Z","shell.execute_reply":"2024-02-21T20:20:04.194970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    lr = 3e-4\n    model = Encoder()\n#     model = timm.create_model(\n#         'wide_resnet101_2',\n#         pretrained=True,\n#         num_classes=len(patients_train)\n#     )\n#     model.conv1 = nn.Conv2d(4, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n    model= nn.DataParallel(model)\n    loss_fn = nn.CrossEntropyLoss()\nif not load_pretrained:\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)\n    num_epochs = 101","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.198537Z","iopub.execute_input":"2024-02-21T20:20:04.198942Z","iopub.status.idle":"2024-02-21T20:20:04.218783Z","shell.execute_reply.started":"2024-02-21T20:20:04.198910Z","shell.execute_reply":"2024-02-21T20:20:04.217536Z"},"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-21T20:20:04.220118Z","iopub.execute_input":"2024-02-21T20:20:04.220627Z","iopub.status.idle":"2024-02-21T20:20:04.233345Z","shell.execute_reply.started":"2024-02-21T20:20:04.220585Z","shell.execute_reply":"2024-02-21T20:20:04.231994Z"},"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/arcface-weights/best_model.pth', map_location=torch.device(device)))","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.234655Z","iopub.execute_input":"2024-02-21T20:20:04.235005Z","iopub.status.idle":"2024-02-21T20:20:04.252003Z","shell.execute_reply.started":"2024-02-21T20:20:04.234977Z","shell.execute_reply":"2024-02-21T20:20:04.249960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_pretrained:\n    losses_train = run_experiment(model, dataloader_train, dataloader_val, loss_fn, optimizer, num_epochs, device, stop_after=15)","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.254219Z","iopub.execute_input":"2024-02-21T20:20:04.254653Z","iopub.status.idle":"2024-02-21T20:20:04.266814Z","shell.execute_reply.started":"2024-02-21T20:20:04.254620Z","shell.execute_reply":"2024-02-21T20:20:04.265122Z"},"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-21T20:20:04.268961Z","iopub.execute_input":"2024-02-21T20:20:04.269425Z","iopub.status.idle":"2024-02-21T20:20:04.289728Z","shell.execute_reply.started":"2024-02-21T20:20:04.269376Z","shell.execute_reply":"2024-02-21T20:20:04.288169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_embs(model, dataloader, device):\n    model = model.to(device)\n    save_linear = model.module.fc\n    model.module.fc = nn.Identity()\n    embs = []\n    with torch.no_grad():\n        model.eval()\n        for batch in tqdm(dataloader, total=len(dataloader)):\n            x = batch\n            x = x.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                out = model(x).squeeze()\n            embs.append(out)\n        embs = torch.cat(embs, dim=0).detach().cpu().numpy()\n    model.module.fc = save_linear\n    del save_linear\n    gc.collect()\n    return embs","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.292250Z","iopub.execute_input":"2024-02-21T20:20:04.292899Z","iopub.status.idle":"2024-02-21T20:20:04.306221Z","shell.execute_reply.started":"2024-02-21T20:20:04.292845Z","shell.execute_reply":"2024-02-21T20:20:04.304772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load_embs:\n    \n    dataset_train_eval = SpecDataset(train_patients, dct, img_size=(99, 300), evaluate=True)\n    \n    dataloader_train_eval = DataLoader(\n        dataset=dataset_train_eval,\n        batch_size=256,\n        shuffle=False,\n        drop_last=False\n    )\n    \n    embs_test = get_embs(model, dataloader_test, device)\n    embs_val = get_embs(model, dataloader_val, device)\n    embs_train = get_embs(model, dataloader_train_eval, device)\n    \n    np.save('specs_arcface_test_embs.npy', embs_test)\n    np.save('specs_arcface_val_embs.npy', embs_val)\n    np.save('specs_arcface_train_embs.npy', embs_train)","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.307757Z","iopub.execute_input":"2024-02-21T20:20:04.308275Z","iopub.status.idle":"2024-02-21T20:20:04.322658Z","shell.execute_reply.started":"2024-02-21T20:20:04.308229Z","shell.execute_reply":"2024-02-21T20:20:04.320956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if load_embs:\n    embs_test = np.load('/kaggle/input/arcface-weights/specs_autoencoder_test_embs.npy')\n    embs_val = np.load('/kaggle/input/arcface-weights/specs_autoencoder_val_embs.npy')\n    embs_train = np.load('/kaggle/input/arcface-weights/specs_autoencoder_train_embs.npy')","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:04.324213Z","iopub.execute_input":"2024-02-21T20:20:04.324762Z","iopub.status.idle":"2024-02-21T20:20:12.438009Z","shell.execute_reply.started":"2024-02-21T20:20:04.324715Z","shell.execute_reply":"2024-02-21T20:20:12.436541Z"},"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).astype(np.float32)\n    embs_val = sc.transform(embs_val).astype(np.float32)\n    embs_test = sc.transform(embs_test).astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:12.439986Z","iopub.execute_input":"2024-02-21T20:20:12.440432Z","iopub.status.idle":"2024-02-21T20:20:12.446925Z","shell.execute_reply.started":"2024-02-21T20:20:12.440399Z","shell.execute_reply":"2024-02-21T20:20:12.445939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from sklearn.metrics.pairwise import euclidean_distances\n# def get_pairwise_dists(embs):\n#     dists = []\n#     for i in tqdm(range(len(embs))):\n#         dists_i = []\n#         if i + 1 == len(embs):\n#             break\n#         for j in range(i+1, len(embs)):\n#             dist = euclidean_distances(embs[i].reshape(1, -1), embs[j].reshape(1, -1))[0, 0]\n#             dists_i.append(dist)\n#         dists.append(dists_i)\n#     return dists\n# get_pairwise_dists(embs_train)","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:12.448581Z","iopub.execute_input":"2024-02-21T20:20:12.449231Z","iopub.status.idle":"2024-02-21T20:20:12.492825Z","shell.execute_reply.started":"2024-02-21T20:20:12.449197Z","shell.execute_reply":"2024-02-21T20:20:12.491098Z"},"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)\n    dists_val = pairwise_distances(embs_val)\n    dists_test = pairwise_distances(embs_test)","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:12.495106Z","iopub.execute_input":"2024-02-21T20:20:12.495938Z","iopub.status.idle":"2024-02-21T20:20:12.508344Z","shell.execute_reply.started":"2024-02-21T20:20:12.495890Z","shell.execute_reply":"2024-02-21T20:20:12.506409Z"},"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-21T20:20:12.516489Z","iopub.execute_input":"2024-02-21T20:20:12.517142Z","iopub.status.idle":"2024-02-21T20:20:12.529207Z","shell.execute_reply.started":"2024-02-21T20:20:12.516871Z","shell.execute_reply":"2024-02-21T20:20:12.527634Z"},"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-21T20:20:12.531262Z","iopub.execute_input":"2024-02-21T20:20:12.532661Z","iopub.status.idle":"2024-02-21T20:20:12.545313Z","shell.execute_reply.started":"2024-02-21T20:20:12.532608Z","shell.execute_reply":"2024-02-21T20:20:12.543318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"thresholds_train = np.arange(0, 100, 0.1)\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/arcface-weights/tpr_train.npy')\n    fpr_train = np.load('/kaggle/input/arcface-weights/fpr_train.npy')\n    acc_train = np.load('/kaggle/input/arcface-weights/acc_train.npy')","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:20:12.546950Z","iopub.execute_input":"2024-02-21T20:20:12.548615Z","iopub.status.idle":"2024-02-21T20:20:12.579870Z","shell.execute_reply.started":"2024-02-21T20:20:12.548559Z","shell.execute_reply":"2024-02-21T20:20:12.578341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Походу что-то не то сохраняется\nfig, 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-21T20:20:12.581659Z","iopub.execute_input":"2024-02-21T20:20:12.582678Z","iopub.status.idle":"2024-02-21T20:20:13.132919Z","shell.execute_reply.started":"2024-02-21T20:20:12.582636Z","shell.execute_reply":"2024-02-21T20:20:13.131363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### UMAP","metadata":{}},{"cell_type":"code","source":"umap_obj = umap.UMAP(n_components=2)\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-21T20:20:13.134875Z","iopub.execute_input":"2024-02-21T20:20:13.135443Z","iopub.status.idle":"2024-02-21T20:25:12.511033Z","shell.execute_reply.started":"2024-02-21T20:20:13.135395Z","shell.execute_reply":"2024-02-21T20:25:12.509420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cons = ['GPD', 'LRDA', 'Seizure', 'Other', 'GRDA', 'LPD']\nax=sns.scatterplot(x=umap_data_train[:,0], y=umap_data_train[:,1], hue=train_patients['expert_consensus'], hue_order=hue_order_cons)\nplt.setp(ax.get_legend().get_texts(), fontsize='10')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:39:28.539706Z","iopub.execute_input":"2024-02-21T20:39:28.540301Z","iopub.status.idle":"2024-02-21T20:39:34.702414Z","shell.execute_reply.started":"2024-02-21T20:39:28.540261Z","shell.execute_reply":"2024-02-21T20:39:34.701046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax=sns.scatterplot(x=umap_data_val[:,0], y=umap_data_val[:,1], hue=val_patients['expert_consensus'], hue_order=hue_order_cons)\nplt.setp(ax.get_legend().get_texts(), fontsize='10')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:39:38.017703Z","iopub.execute_input":"2024-02-21T20:39:38.018232Z","iopub.status.idle":"2024-02-21T20:39:39.922619Z","shell.execute_reply.started":"2024-02-21T20:39:38.018186Z","shell.execute_reply":"2024-02-21T20:39:39.920194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax=sns.scatterplot(x=umap_data_test[:,0], y=umap_data_test[:,1], hue=test_patients['expert_consensus'], hue_order=hue_order_cons)\nplt.setp(ax.get_legend().get_texts(), fontsize='10')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:39:42.593929Z","iopub.execute_input":"2024-02-21T20:39:42.594384Z","iopub.status.idle":"2024-02-21T20:39:44.564718Z","shell.execute_reply.started":"2024-02-21T20:39:42.594353Z","shell.execute_reply":"2024-02-21T20:39:44.562961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hue_order_cls = ['idealized', 'proto', 'edge', 'other'] \nax=sns.scatterplot(x=umap_data_train[:,0], y=umap_data_train[:,1], hue=train_patients['authors_class'], hue_order=hue_order_cls)\nplt.setp(ax.get_legend().get_texts(), fontsize='10')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:40:42.417588Z","iopub.execute_input":"2024-02-21T20:40:42.418289Z","iopub.status.idle":"2024-02-21T20:40:48.463650Z","shell.execute_reply.started":"2024-02-21T20:40:42.418238Z","shell.execute_reply":"2024-02-21T20:40:48.462323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax=sns.scatterplot(x=umap_data_val[:,0], y=umap_data_val[:,1], hue=val_patients['authors_class'], hue_order=hue_order_cls)\nplt.setp(ax.get_legend().get_texts(), fontsize='10')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:40:52.330644Z","iopub.execute_input":"2024-02-21T20:40:52.331169Z","iopub.status.idle":"2024-02-21T20:40:54.005523Z","shell.execute_reply.started":"2024-02-21T20:40:52.331129Z","shell.execute_reply":"2024-02-21T20:40:54.004415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax=sns.scatterplot(x=umap_data_test[:,0], y=umap_data_test[:,1], hue=test_patients['authors_class'], hue_order=hue_order_cls)\nplt.setp(ax.get_legend().get_texts(), fontsize='10')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-21T20:40:54.485420Z","iopub.execute_input":"2024-02-21T20:40:54.486490Z","iopub.status.idle":"2024-02-21T20:40:56.298739Z","shell.execute_reply.started":"2024-02-21T20:40:54.486444Z","shell.execute_reply":"2024-02-21T20:40:56.297088Z"},"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-21T19:14:32.228924Z","iopub.status.idle":"2024-02-21T19:14:32.229322Z","shell.execute_reply.started":"2024-02-21T19:14:32.229104Z","shell.execute_reply":"2024-02-21T19:14:32.229119Z"},"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-21T19:14:32.231365Z","iopub.status.idle":"2024-02-21T19:14:32.231914Z","shell.execute_reply.started":"2024-02-21T19:14:32.231618Z","shell.execute_reply":"2024-02-21T19:14:32.231639Z"},"trusted":true},"execution_count":null,"outputs":[]}]}