{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Detecting Continuous Gravitational Waves","metadata":{}},{"cell_type":"code","source":"%pip install skorch timm --user","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:20.858252Z","iopub.execute_input":"2022-12-19T09:28:20.859329Z","iopub.status.idle":"2022-12-19T09:28:31.717816Z","shell.execute_reply.started":"2022-12-19T09:28:20.859264Z","shell.execute_reply":"2022-12-19T09:28:31.716364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\nimport torch\nfrom torch import nn\nfrom torch.optim import Adam, AdamW\nfrom torch.utils.data import Dataset, random_split\n\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom skorch import NeuralNetClassifier\nfrom skorch.callbacks import ProgressBar, Checkpoint, LRScheduler, EpochScoring\nfrom skorch.helper import predefined_split\n\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score, average_precision_score, confusion_matrix, roc_curve, precision_recall_curve, make_scorer, brier_score_loss, recall_score, precision_score\n\nimport h5py\nimport numpy as np\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom os import walk\nfrom multiprocessing import Pool","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-19T09:28:31.721326Z","iopub.execute_input":"2022-12-19T09:28:31.721898Z","iopub.status.idle":"2022-12-19T09:28:31.733163Z","shell.execute_reply.started":"2022-12-19T09:28:31.721843Z","shell.execute_reply":"2022-12-19T09:28:31.731991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline\nplt.style.use('grayscale')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.735024Z","iopub.execute_input":"2022-12-19T09:28:31.73539Z","iopub.status.idle":"2022-12-19T09:28:31.746234Z","shell.execute_reply.started":"2022-12-19T09:28:31.735356Z","shell.execute_reply":"2022-12-19T09:28:31.745116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_WORKERS = 4\nNUM_EPOCHS = 10\nBATCH_SIZE = 2\nLEARNING_RATE = 1e-5\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.749354Z","iopub.execute_input":"2022-12-19T09:28:31.750251Z","iopub.status.idle":"2022-12-19T09:28:31.759411Z","shell.execute_reply.started":"2022-12-19T09:28:31.750216Z","shell.execute_reply":"2022-12-19T09:28:31.758086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"RANDOM_STATE = 177013\ntorch.manual_seed(RANDOM_STATE)\ntorch.cuda.manual_seed(RANDOM_STATE)\nnp.random.seed(RANDOM_STATE)\ntorch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.76138Z","iopub.execute_input":"2022-12-19T09:28:31.762238Z","iopub.status.idle":"2022-12-19T09:28:31.772319Z","shell.execute_reply.started":"2022-12-19T09:28:31.76218Z","shell.execute_reply":"2022-12-19T09:28:31.771368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MAX_LENGTH=4200","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.773981Z","iopub.execute_input":"2022-12-19T09:28:31.775168Z","iopub.status.idle":"2022-12-19T09:28:31.783685Z","shell.execute_reply.started":"2022-12-19T09:28:31.775119Z","shell.execute_reply":"2022-12-19T09:28:31.782474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/g2net-detecting-continuous-gravitational-waves/train_labels.csv')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.785265Z","iopub.execute_input":"2022-12-19T09:28:31.786065Z","iopub.status.idle":"2022-12-19T09:28:31.80175Z","shell.execute_reply.started":"2022-12-19T09:28:31.786027Z","shell.execute_reply":"2022-12-19T09:28:31.800583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = dict(df.values)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.803456Z","iopub.execute_input":"2022-12-19T09:28:31.804134Z","iopub.status.idle":"2022-12-19T09:28:31.811708Z","shell.execute_reply.started":"2022-12-19T09:28:31.804071Z","shell.execute_reply":"2022-12-19T09:28:31.810475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['target'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.813377Z","iopub.execute_input":"2022-12-19T09:28:31.814332Z","iopub.status.idle":"2022-12-19T09:28:31.82657Z","shell.execute_reply.started":"2022-12-19T09:28:31.814286Z","shell.execute_reply":"2022-12-19T09:28:31.825702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The -1 labels are a sort of an inside joke so we'll be omitting those.","metadata":{}},{"cell_type":"code","source":"exclusions = df[df['target'] == -1]['id'].to_list()","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.830089Z","iopub.execute_input":"2022-12-19T09:28:31.830432Z","iopub.status.idle":"2022-12-19T09:28:31.838937Z","shell.execute_reply.started":"2022-12-19T09:28:31.830402Z","shell.execute_reply":"2022-12-19T09:28:31.837892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparing the dataset","metadata":{}},{"cell_type":"code","source":"train_dir = '/kaggle/input/g2net-detecting-continuous-gravitational-waves/train/'\ntest_dir = '/kaggle/input/g2net-detecting-continuous-gravitational-waves/test'","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.84039Z","iopub.execute_input":"2022-12-19T09:28:31.840704Z","iopub.status.idle":"2022-12-19T09:28:31.848319Z","shell.execute_reply.started":"2022-12-19T09:28:31.840675Z","shell.execute_reply":"2022-12-19T09:28:31.847452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def normalize(wave):\n    return wave / 3.488688e-20","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.850545Z","iopub.execute_input":"2022-12-19T09:28:31.850912Z","iopub.status.idle":"2022-12-19T09:28:31.85877Z","shell.execute_reply.started":"2022-12-19T09:28:31.850881Z","shell.execute_reply":"2022-12-19T09:28:31.857759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Transforms:\n    def __init__(self, sequence):\n        self.transforms = sequence\n\n    def __call__(self, img, *args, **kwargs):\n        return self.transforms(image=np.array(img))['image']\n\ntransformations = Transforms(A.Compose([\n\tA.Rotate(10, interpolation=0),\n\tA.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.ColorJitter(),\n\tToTensorV2(),\n]))\n\ninference_transformations = Transforms(A.Compose([\n\tToTensorV2(),\n]))","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.872393Z","iopub.execute_input":"2022-12-19T09:28:31.873198Z","iopub.status.idle":"2022-12-19T09:28:31.883422Z","shell.execute_reply.started":"2022-12-19T09:28:31.873149Z","shell.execute_reply":"2022-12-19T09:28:31.881867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_hdf(path):\n    data = []\n    with h5py.File(path, \"r\") as f:\n        for file_key in f.keys():\n            group = f[file_key]\n            if isinstance(group, h5py._hl.dataset.Dataset):\n                data.append(np.array(group))\n                continue\n            for group_key in group.keys():\n                group2 = group[group_key]\n                if isinstance(group2, h5py._hl.dataset.Dataset):\n                    data.append(np.array(group2))\n                    continue\n                for group_key2 in group2.keys():\n                    group3 = group2[group_key2]\n                    if isinstance(group3, h5py._hl.dataset.Dataset):\n                        data.append(np.array(group3))\n                        continue\n    \n    ch1 = torch.from_numpy(normalize(data[0][:,:MAX_LENGTH]))\n    ch2 = torch.from_numpy(normalize(data[2][:,:MAX_LENGTH]))\n\n    chunk = torch.stack((ch1, ch2))\n    \n    return chunk","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.885033Z","iopub.execute_input":"2022-12-19T09:28:31.885397Z","iopub.status.idle":"2022-12-19T09:28:31.895163Z","shell.execute_reply.started":"2022-12-19T09:28:31.885363Z","shell.execute_reply":"2022-12-19T09:28:31.894006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def powernorm(spec):\n    ch1 = torch.abs(spec[0]) ** 2\n    ch2 = torch.abs(spec[1]) ** 2\n   \n    ch1 = (ch1 - ch1.mean()) / ch1.mean()\n    ch2 = (ch2 - ch2.mean()) / ch2.mean()\n    \n    return ch1, ch2","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.896655Z","iopub.execute_input":"2022-12-19T09:28:31.897274Z","iopub.status.idle":"2022-12-19T09:28:31.907986Z","shell.execute_reply.started":"2022-12-19T09:28:31.897238Z","shell.execute_reply":"2022-12-19T09:28:31.906817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compr(ch, augment=False):\n    result = ch.reshape(360, 56, 75).mean(dim=2)\n    return result","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.909321Z","iopub.execute_input":"2022-12-19T09:28:31.910246Z","iopub.status.idle":"2022-12-19T09:28:31.922239Z","shell.execute_reply.started":"2022-12-19T09:28:31.910198Z","shell.execute_reply":"2022-12-19T09:28:31.921231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HDFHotSwap(torch.utils.data.Dataset):\n    '''\n    Dynamically loads HDF5 files from every subdir.\n    '''\n    @staticmethod\n    def process_file(file):\n        '''\n        Extracts the observation id from filename.\n        '''\n        observation_id = file[:-5]\n        return observation_id\n\n    def __init__(self, folder, labels, augment=False):\n        '''\n        Accepts the folder and labels dict, and whether to use augmentations.\n        '''\n        self.root = folder\n        self.files = []\n        self.augment = augment\n        # build a list of files:\n        pool = Pool(NUM_WORKERS)\n        for subdir, dirs, files in walk(self.root):\n            results = pool.map(HDFHotSwap.process_file, files)\n            if results:\n                self.files.extend(results)\n        self.files = list(set(self.files).difference(set(exclusions)))\n    \n    def __len__(self):\n        return len(self.files)\n    \n    def __getitem__(self, idx):\n        file = self.files[idx]\n        path = self.root + '/' + file + '.hdf5'\n\n        data = read_hdf(path)\n        \n        ch1, ch2 = powernorm(data)\n        \n        avg = torch.stack((compr(ch1, self.augment), compr(ch2, self.augment))).mean(dim=0)\n        result = transformations(avg)[0] if self.augment else inference_transformations(avg)[0]\n        target = labels[file] if file in labels else np.nan\n\n        return result.unsqueeze(0), torch.from_numpy(np.array(target).astype('long')) if file in labels else np.nan","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.923745Z","iopub.execute_input":"2022-12-19T09:28:31.924264Z","iopub.status.idle":"2022-12-19T09:28:31.937433Z","shell.execute_reply.started":"2022-12-19T09:28:31.92423Z","shell.execute_reply":"2022-12-19T09:28:31.936264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main_set = HDFHotSwap(train_dir, labels, augment=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:31.938902Z","iopub.execute_input":"2022-12-19T09:28:31.940345Z","iopub.status.idle":"2022-12-19T09:28:32.190614Z","shell.execute_reply.started":"2022-12-19T09:28:31.940293Z","shell.execute_reply":"2022-12-19T09:28:32.18833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def spectrogram(entry):\n    print(f'Target: {entry[1]}')\n    fig, ax = plt.subplots(figsize=(2,5))\n    sns.heatmap(entry[0][0], ax=ax)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:32.194359Z","iopub.execute_input":"2022-12-19T09:28:32.194895Z","iopub.status.idle":"2022-12-19T09:28:32.202978Z","shell.execute_reply.started":"2022-12-19T09:28:32.194838Z","shell.execute_reply":"2022-12-19T09:28:32.201785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_set = HDFHotSwap(test_dir, labels)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:32.20464Z","iopub.execute_input":"2022-12-19T09:28:32.205644Z","iopub.status.idle":"2022-12-19T09:28:33.727678Z","shell.execute_reply.started":"2022-12-19T09:28:32.205595Z","shell.execute_reply":"2022-12-19T09:28:33.725545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spectrogram(main_set[0])","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:33.73108Z","iopub.execute_input":"2022-12-19T09:28:33.732479Z","iopub.status.idle":"2022-12-19T09:28:35.133081Z","shell.execute_reply.started":"2022-12-19T09:28:33.732431Z","shell.execute_reply":"2022-12-19T09:28:35.131988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_len = int(len(main_set) * 0.8)\nval_len = len(main_set) - train_len\n\ntrain_set, valid_set = random_split(main_set, [train_len, val_len])","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:35.135118Z","iopub.execute_input":"2022-12-19T09:28:35.135568Z","iopub.status.idle":"2022-12-19T09:28:35.142338Z","shell.execute_reply.started":"2022-12-19T09:28:35.135511Z","shell.execute_reply":"2022-12-19T09:28:35.141197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Building models","metadata":{}},{"cell_type":"code","source":"progress = ProgressBar()\nscheduler = LRScheduler(policy='CosineAnnealingLR', T_max=12, eta_min=1e-7)\nroc = EpochScoring(make_scorer(roc_auc_score), name=f'roc_auc', lower_is_better=False, use_caching=True)\nrecall = EpochScoring(make_scorer(recall_score), name=f'recall', lower_is_better=False, use_caching=True)\nprecision = EpochScoring(make_scorer(precision_score), name=f'precision', lower_is_better=False, use_caching=True)\nbrier = EpochScoring(make_scorer(brier_score_loss), name=f'brier', lower_is_better=True, use_caching=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:35.14399Z","iopub.execute_input":"2022-12-19T09:28:35.144352Z","iopub.status.idle":"2022-12-19T09:28:35.154887Z","shell.execute_reply.started":"2022-12-19T09:28:35.14432Z","shell.execute_reply":"2022-12-19T09:28:35.153963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_net(model, val_set=predefined_split(valid_set), n_epochs=5, lr=1e-5, checkpoint_callback=None):\n    net = NeuralNetClassifier(\n                model,\n                warm_start=True,\n                max_epochs = n_epochs,\n                batch_size = BATCH_SIZE,\n                lr = lr,\n                criterion = nn.CrossEntropyLoss(label_smoothing=0.3),\n                optimizer = AdamW,\n                optimizer__weight_decay = 1e-6,\n                device = DEVICE,\n                iterator_train__shuffle = True,\n                iterator_train__num_workers = NUM_WORKERS,\n                iterator_valid__num_workers = NUM_WORKERS,\n                iterator_train__pin_memory = True,\n                iterator_valid__pin_memory = True,\n        \n                train_split = val_set, \n        \n                callbacks = [progress,\n                             checkpoint_callback,\n                             roc, recall, precision, brier,\n                             LRScheduler(policy='CosineAnnealingLR',  T_max=25, eta_min=1e-7),\n                            ],\n    )\n    return net","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:35.156572Z","iopub.execute_input":"2022-12-19T09:28:35.157048Z","iopub.status.idle":"2022-12-19T09:28:35.167508Z","shell.execute_reply.started":"2022-12-19T09:28:35.157012Z","shell.execute_reply":"2022-12-19T09:28:35.166491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model1 = timm.create_model('eca_nfnet_l1', pretrained=False, num_classes=2, in_chans=1)\nmodel2 = timm.create_model('nf_resnet50', pretrained=False, num_classes=2, in_chans=1)\nmodel3 = timm.create_model('nfnet_l0', pretrained=False, num_classes=2, in_chans=1)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:35.169353Z","iopub.execute_input":"2022-12-19T09:28:35.169811Z","iopub.status.idle":"2022-12-19T09:28:36.897123Z","shell.execute_reply.started":"2022-12-19T09:28:35.169767Z","shell.execute_reply":"2022-12-19T09:28:36.896138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cp_net1 = Checkpoint(dirname='net1_checkpoint', load_best=True, monitor='valid_acc_best')\ncp_net2 = Checkpoint(dirname='net2_checkpoint', load_best=True, monitor='valid_acc_best')\ncp_net3 = Checkpoint(dirname='net3_checkpoint', load_best=True, monitor='valid_acc_best')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:36.898532Z","iopub.execute_input":"2022-12-19T09:28:36.898857Z","iopub.status.idle":"2022-12-19T09:28:36.906215Z","shell.execute_reply.started":"2022-12-19T09:28:36.898826Z","shell.execute_reply":"2022-12-19T09:28:36.904981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net1 = build_net(model1, val_set=predefined_split(valid_set), n_epochs=25, checkpoint_callback=cp_net1)\nnet2 = build_net(model2, val_set=predefined_split(valid_set), n_epochs=25, checkpoint_callback=cp_net2)\nnet3 = build_net(model3, val_set=predefined_split(valid_set), n_epochs=25, checkpoint_callback=cp_net3)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:36.912426Z","iopub.execute_input":"2022-12-19T09:28:36.912765Z","iopub.status.idle":"2022-12-19T09:28:36.919217Z","shell.execute_reply.started":"2022-12-19T09:28:36.912735Z","shell.execute_reply":"2022-12-19T09:28:36.91812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Those models were pretrained on ~6000 synthetic samples (NOT the provided train set). Due to the enourmous amount of disk space this takes, only the resulting weights are being loaded:","metadata":{}},{"cell_type":"code","source":"net1.initialize()\nnet1.load_params(f_params='/kaggle/input/gwnet-weights/net1.pkl')\nnet2.initialize()\nnet2.load_params(f_params='/kaggle/input/gwnet-weights/net2.pkl')\nnet3.initialize()\nnet3.load_params(f_params='/kaggle/input/gwnet-weights/net3.pkl')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:36.920496Z","iopub.execute_input":"2022-12-19T09:28:36.920822Z","iopub.status.idle":"2022-12-19T09:28:37.427622Z","shell.execute_reply.started":"2022-12-19T09:28:36.920792Z","shell.execute_reply":"2022-12-19T09:28:37.426462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Synthetic models' evaluation against a provided train set","metadata":{}},{"cell_type":"code","source":"def calculate_metrics(probabilities, target_test):\n    predictions = (probabilities > 0.5)\n    f1 = f1_score(target_test, predictions)\n    roc_auc = roc_auc_score(target_test, probabilities)\n    acc = accuracy_score(target_test, predictions)\n    ap = average_precision_score(target_test, probabilities)\n    cmatrix = confusion_matrix(target_test, predictions)\n    \n    fpr, tpr, _ = roc_curve(target_test, probabilities)\n    precision, recall, thresholds = precision_recall_curve(target_test, probabilities)\n    f1_scores = 2 * recall * precision / (recall + precision)\n    best_f1 = np.max(f1_scores)\n    best_thresh = thresholds[np.argmax(f1_scores)]\n\n    pred_t = (probabilities > best_thresh)\n    best_cmatrix = confusion_matrix(target_test, pred_t)\n    \n    return f1, best_f1, roc_auc, acc, ap, best_thresh, fpr, tpr, recall, precision, cmatrix, best_cmatrix\n\ndef visualize_tests(probabilities, target_test):\n    cmatrices = []\n\n    fig, axes = plt.subplots(1, 2, figsize=(15,6))\n    axes[0].plot([0, 1], linestyle='--')\n    axes[1].plot([0.5, 0.5], linestyle='--')\n\n    print('Processing validation set, please wait warmly...')\n    f1, best_f1, roc_auc, acc, ap, best_thresh, fpr, tpr, recall, precision, cmatrix, best_cmatrix = calculate_metrics (probabilities, target_test)\n    axes[0].plot (fpr, tpr);\n    axes[1].plot (recall, precision);\n    print (f'F1: {f1:.2f} (max: {best_f1:.2f} at {best_thresh:.2f} threshold), ROC_AUC: {roc_auc:.3f}, accuracy: {acc:.0%}, AP (PR_AUC): {ap:.2f}')\n    cmatrices.append(cmatrix)\n    cmatrices.append(best_cmatrix)\n    \n    axes[0].set (xlabel='FPR', ylabel='TPR', title='ROC curve', xlim=(0,1), ylim=(0,1))\n    axes[1].set (xlabel='Recall', ylabel='Precision', title='PR curve', xlim=(0,1), ylim=(0,1))\n    \n    \n    fig, axes = plt.subplots(1, 2, figsize=(13, 5), constrained_layout=True)\n    for cmatrix, ax, title in zip(cmatrices, axes.flat, ['Confusion Maxtrix', 'Confusion Maxtrix (optimal F1 threshold)']):\n        sns.heatmap(cmatrix, ax=ax, annot=True, cmap='Blues', fmt='d').set(title=title, xlabel='Prediction', ylabel='Reality')\n    \n    \n    return best_thresh","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:37.429487Z","iopub.execute_input":"2022-12-19T09:28:37.430185Z","iopub.status.idle":"2022-12-19T09:28:37.448582Z","shell.execute_reply.started":"2022-12-19T09:28:37.430127Z","shell.execute_reply":"2022-12-19T09:28:37.447456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_valid = [l[1] for l in valid_set]","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:28:37.450288Z","iopub.execute_input":"2022-12-19T09:28:37.450961Z","iopub.status.idle":"2022-12-19T09:29:09.716871Z","shell.execute_reply.started":"2022-12-19T09:28:37.450901Z","shell.execute_reply":"2022-12-19T09:29:09.715647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_train = [l[1] for l in train_set]","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:29:09.718506Z","iopub.execute_input":"2022-12-19T09:29:09.719198Z","iopub.status.idle":"2022-12-19T09:31:13.996908Z","shell.execute_reply.started":"2022-12-19T09:29:09.719152Z","shell.execute_reply":"2022-12-19T09:31:13.994592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"probabilities1 = net1.predict_proba(train_set)[:, 1]\nprobabilities2 = net2.predict_proba(train_set)[:, 1]\nprobabilities3 = net3.predict_proba(train_set)[:, 1]","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:31:14.001893Z","iopub.execute_input":"2022-12-19T09:31:14.003769Z","iopub.status.idle":"2022-12-19T09:35:05.151592Z","shell.execute_reply.started":"2022-12-19T09:31:14.003704Z","shell.execute_reply":"2022-12-19T09:35:05.149938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_threshold1 = visualize_tests(probabilities1, target_train)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:35:05.154078Z","iopub.execute_input":"2022-12-19T09:35:05.15603Z","iopub.status.idle":"2022-12-19T09:35:06.410508Z","shell.execute_reply.started":"2022-12-19T09:35:05.155968Z","shell.execute_reply":"2022-12-19T09:35:06.409265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_threshold2 = visualize_tests(probabilities2, target_train)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:35:06.412349Z","iopub.execute_input":"2022-12-19T09:35:06.413467Z","iopub.status.idle":"2022-12-19T09:35:07.592105Z","shell.execute_reply.started":"2022-12-19T09:35:06.413428Z","shell.execute_reply":"2022-12-19T09:35:07.591003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_threshold3 = visualize_tests(probabilities3, target_train)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:35:07.594087Z","iopub.execute_input":"2022-12-19T09:35:07.594451Z","iopub.status.idle":"2022-12-19T09:35:08.867Z","shell.execute_reply.started":"2022-12-19T09:35:07.594419Z","shell.execute_reply":"2022-12-19T09:35:08.865639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The result is not so bad, yielding around 0.8 ROC_AUC on small unseen part of the initial train set.","metadata":{}},{"cell_type":"markdown","source":"## Ensembling","metadata":{}},{"cell_type":"markdown","source":"Intead of fine-tuning on the initial train set, we'll adjust the synthetic models' bias using logistic regression of their scores on the train set against the target:","metadata":{}},{"cell_type":"code","source":"target_train = [l[1] for l in main_set]\nprobabilities1 = net1.predict_proba(main_set)[:, 1]\nprobabilities2 = net2.predict_proba(main_set)[:, 1]\nprobabilities3 = net3.predict_proba(main_set)[:, 1]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fit_ensemble():\n    meta_X = list()\n    meta_X.append(probabilities1.reshape(-1, 1))\n    meta_X.append(probabilities2.reshape(-1, 1))\n    meta_X.append(probabilities3.reshape(-1, 1))\n    meta_X = np.hstack(meta_X)\n    blender = LogisticRegression(random_state=RANDOM_STATE, n_jobs=-1)\n    blender.fit(meta_X, target_train)\n    return blender\n\ndef predict_ensemble(blender, features):\n    meta_X = list()\n    meta_X.append(net1.predict_proba(features)[:, 1].reshape(-1, 1))\n    meta_X.append(net2.predict_proba(features)[:, 1].reshape(-1, 1))\n    meta_X.append(net3.predict_proba(features)[:, 1].reshape(-1, 1))\n    meta_X = np.hstack(meta_X)\n    return blender.predict_proba(meta_X)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:35:08.86878Z","iopub.execute_input":"2022-12-19T09:35:08.869267Z","iopub.status.idle":"2022-12-19T09:35:08.880452Z","shell.execute_reply.started":"2022-12-19T09:35:08.869221Z","shell.execute_reply":"2022-12-19T09:35:08.879117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"blender = fit_ensemble()","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:35:08.882508Z","iopub.execute_input":"2022-12-19T09:35:08.883196Z","iopub.status.idle":"2022-12-19T09:35:08.910594Z","shell.execute_reply.started":"2022-12-19T09:35:08.88315Z","shell.execute_reply":"2022-12-19T09:35:08.909453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = predict_ensemble(blender, test_set)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T09:36:08.462593Z","iopub.execute_input":"2022-12-19T09:36:08.463711Z","iopub.status.idle":"2022-12-19T10:43:03.262939Z","shell.execute_reply.started":"2022-12-19T09:36:08.463672Z","shell.execute_reply":"2022-12-19T10:43:03.259711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_results = pd.DataFrame(data = {'id': test_set.files, 'target': predictions[:,1]}).set_index('id')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:43:03.274373Z","iopub.execute_input":"2022-12-19T10:43:03.276207Z","iopub.status.idle":"2022-12-19T10:43:03.318609Z","shell.execute_reply.started":"2022-12-19T10:43:03.276114Z","shell.execute_reply":"2022-12-19T10:43:03.31604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_results.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:43:03.329008Z","iopub.execute_input":"2022-12-19T10:43:03.333468Z","iopub.status.idle":"2022-12-19T10:43:03.387828Z","shell.execute_reply.started":"2022-12-19T10:43:03.333389Z","shell.execute_reply":"2022-12-19T10:43:03.386957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This results in about 0.7 ROC_AUC on the public test set.","metadata":{}}]}