{"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":"# This Notebook\n\nThis notebook explores training a model on three features to solve the G2Net continuous gravitational waves task. The model is a custom multi-head CNN with three heads, and the three features are a noise reduced spectrogram, a Viterbi track and a Viterbi map. These features are described in greater detail in the 'Features' section of this notebook.","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install timm\n\nimport timm\nimport torch\nfrom torch import nn\nfrom scipy.stats import norm\nimport torchaudio\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nfrom sklearn.model_selection import KFold\nimport time\nimport h5py\nimport os\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import roc_auc_score\nfrom timm.scheduler import CosineLRScheduler\nimport gc","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-01-04T12:08:42.367376Z","iopub.execute_input":"2023-01-04T12:08:42.368187Z","iopub.status.idle":"2023-01-04T12:09:02.705856Z","shell.execute_reply.started":"2023-01-04T12:08:42.368061Z","shell.execute_reply":"2023-01-04T12:09:02.703911Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data\n\n## Generated Data\nA modified version of [George Chirita's notebook](https://www.kaggle.com/code/crischir/consolidatedg2netdataset/notebook) has been used to consolidate Chirita's multiple synthetic datasets and the original G2Net data into a single dataset, from which three features have been extracted. Class imbalance still exists in this consolidated dataset; a Weighted Random Sampler has been used to oversample the minority class to reduce the impact of this problem.\n\n## Features\nThis notebook's contribution is the exploration of a model with multiple inputs to solve the G2Net continuous gravitational waves task. Three features are used as inputs for the model:\n1. Noise reduced spectrograms\n    - In [their notebook](https://www.kaggle.com/code/laeyoung/g2net-large-kernel-inference), laeyoung demonstrated that using large kernels helps CNNs to extract signal shapes from the G2Net data.\n    - We have used the stem from their model to improve the signal-to-noise ratio of the G2Net data.\n    - This feature extraction process takes an extremely long time due to the high dimensionality of the input to laeyoung's stem. The process has thus been separated from this training notebook. It can be found [here](https://www.kaggle.com/code/hnsyprst/publish-consolidatedg2netdataset).\n2. Viterbi Tracks\n    - [Bayley et al. (2020)](https://arxiv.org/pdf/2007.08207.pdf) developed SOAP, an 'algorithm to search for continuous gravitational waves'. \n    - SOAP is based on the Viterbi algorithm.\n    - The Viterbi algorithm finds the most likely sequence of states through a series of states with given transition probabilities.\n    - In SOAP's context, this means finding the most likely path through a given spectrogram (the path that will give the highest sum of FFT power).\n    - If a signal is present in a given spectrogram, this path should correspond with that signal.\n    - These paths, which Bayley et al. call 'Viterbi tracks', have been extracted for each of the noise reduced spectrograms described above.\n    - Using the SOAP algorithm also takes a long time, so this feature extraction process has also been separated from this training notebook. It can be found [here](https://www.kaggle.com/code/hnsyprst/publish-vit-track-and-map-generation).\n3. Viterbi Maps\n    - The Viterbi map encodes the probability that the signal (the Viterbi track) is in each frequency bin at each time across a given spectrogram.\n    - SOAP also enables extraction of the Viterbi maps. The process of extracting these features can be found in the same notebook as the Viterbi track extraction ([here](https://www.kaggle.com/code/hnsyprst/publish-vit-track-and-map-generation)).\n    \nFor reference, each of the three features are displayed below for a single sample.\n\nBayley, J., Messenger, C. and Woan, G. (2020) ‘A robust machine learning algorithm to search for continuous gravitational waves’, *Physical Review D*, 102(8), p. 083024. Available at: https://doi.org/10.1103/PhysRevD.102.083024.","metadata":{}},{"cell_type":"code","source":"# Data locations\nTRAIN_SPECT_DIR = '/kaggle/input/g2net-consolidated-reduced-noise/consolidated'\nTRAIN_VIT_MAP_DIR = '/kaggle/input/g2net-consolidated-vits/vit_maps_archive/train'\nTRAIN_VIT_TRACK_DIR = '/kaggle/input/g2net-consolidated-vits/vit_tracks_archive/train'\n\nTEST_SPECT_DIR = '/kaggle/input/g2net-reduced-noise'\nTEST_VIT_MAP_DIR = '/kaggle/input/g2net-denoised-vit-maps'\nTEST_VIT_TRACK_DIR = '/kaggle/input/g2net-denoised-vit-tracks'","metadata":{"execution":{"iopub.status.busy":"2023-01-04T12:09:02.708701Z","iopub.execute_input":"2023-01-04T12:09:02.709716Z","iopub.status.idle":"2023-01-04T12:09:02.717034Z","shell.execute_reply.started":"2023-01-04T12:09:02.70967Z","shell.execute_reply":"2023-01-04T12:09:02.715533Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compress spectrograms to reduce input dimensionality\ndef load_spect(file):\n    a = file[:, :720]\n\n    p = a.real**2 + a.imag**2  # power\n    p /= np.mean(p)  # normalize\n    p = np.sum(p.reshape(360, 90, 8), axis=2)\n\n    file_new = p\n\n    file_rot = np.rot90(file_new, 3)\n    return file_rot\n\ndef load_track(file):\n    file = np.squeeze(file)\n    return file\n\ndef display_data(spect, vitmap, track):\n    fig_ng, ax_ng = plt.subplots(figsize=(16,12),nrows=3,constrained_layout=True)\n    fig_ng.suptitle('Three Types of Training Data', fontsize=16)\n    ax_ng[0].imshow(spect.T,aspect=\"auto\",origin=\"lower\",cmap=\"gray\")\n    ax_ng[0].set_title('Noise Reduced Spectrogram')\n    \n    ax_ng[1].imshow(vitmap.T,aspect=\"auto\",origin=\"lower\",cmap=\"YlGnBu\")\n    ax_ng[1].set_title('Viterbi Map')\n    \n    ax_ng[2].scatter(np.arange(spect.shape[0]), np.argmax(track, axis=1), color=\"red\", s=3)\n    ax_ng[2].set_ylim([0, spect.shape[1]])\n    ax_ng[2].set_xlim([0, spect.shape[0]])\n    ax_ng[2].set_title('Viterbi Track')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-01-04T12:09:02.718856Z","iopub.execute_input":"2023-01-04T12:09:02.719227Z","iopub.status.idle":"2023-01-04T12:09:02.739992Z","shell.execute_reply.started":"2023-01-04T12:09:02.719189Z","shell.execute_reply":"2023-01-04T12:09:02.738618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_file_id = '02887d232'\nexample_target_type = 'signals'\nexample_spect = load_spect(np.load('%s/%s/%s.npy' % (TRAIN_SPECT_DIR, example_target_type, example_file_id)))\nexample_vitmap = np.load('%s/%s/%s.npy' % (TRAIN_VIT_MAP_DIR, example_target_type, example_file_id))\nexample_track = load_track(np.load('%s/%s/%s.npy' % (TRAIN_VIT_TRACK_DIR, example_target_type, example_file_id)))\n\ndisplay_data(example_spect, example_vitmap, example_track)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-01-04T12:09:02.741673Z","iopub.execute_input":"2023-01-04T12:09:02.742862Z","iopub.status.idle":"2023-01-04T12:09:04.067006Z","shell.execute_reply.started":"2023-01-04T12:09:02.742817Z","shell.execute_reply":"2023-01-04T12:09:04.065382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To illustrate the difference between these features and the original spectrograms, compressed spectrograms from the L1 and H1 interferometers for the same sample examined above are presented without feature extraction below:","metadata":{}},{"cell_type":"code","source":"# Compress spectrograms to reduce input dimensionality\ndef load_original_spect(file):\n    img = np.empty((2, 90, 360), dtype=np.float32)\n    with h5py.File(file, 'r') as f:\n        g = f[os.path.splitext(os.path.basename(file))[0]]\n        for ch, s in enumerate(['H1', 'L1']):\n            a = g[s]['SFTs'][:, :4590] * 1e22  # Fourier coefficient complex64\n\n            p = a.real**2 + a.imag**2  # power\n            p /= np.mean(p)  # normalize\n            p = np.mean(p.reshape(360, 90, 51), axis=2)  # compress 5760 -> 90\n            p = np.rot90(p, 3)\n\n            img[ch] = p\n    return img\n\ndef display_original_data(H1, L1):\n    fig_ng, ax_ng = plt.subplots(figsize=(16,8),nrows=2,constrained_layout=True)\n    fig_ng.suptitle('Spectrogram before Feature Extraction', fontsize=16)\n    ax_ng[0].imshow(H1.T,aspect=\"auto\",origin=\"lower\",cmap=\"gray\")\n    ax_ng[0].set_title('H1 Spectrogram')\n    \n    ax_ng[1].imshow(L1.T,aspect=\"auto\",origin=\"lower\",cmap=\"gray\")\n    ax_ng[1].set_title('L1 Spectrogram')\n    \nspects = load_original_spect('/kaggle/input/g2net-detecting-continuous-gravitational-waves/train/%s.hdf5' % (example_file_id))\ndisplay_original_data(spects[0], spects[1])","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-01-04T12:13:13.069803Z","iopub.execute_input":"2023-01-04T12:13:13.070251Z","iopub.status.idle":"2023-01-04T12:13:14.357759Z","shell.execute_reply.started":"2023-01-04T12:13:13.070213Z","shell.execute_reply":"2023-01-04T12:13:14.356668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Loading and Augmentation\n\nIn [their notebook](https://www.kaggle.com/code/myso1987/g2net-basic-audio-data-augmentation/notebook), MYSO introduces some data augmentation techniques for the G2Net data. These techniques (time and frequency masking, horizontal and vertical flips and vertical translation) have been replicated here.","metadata":{}},{"cell_type":"code","source":"# Time and frequency masking data augmentation setup\ntransforms_time_mask = nn.Sequential(\n                torchaudio.transforms.TimeMasking(time_mask_param=10),\n            )\n\ntransforms_freq_mask = nn.Sequential(\n                torchaudio.transforms.FrequencyMasking(freq_mask_param=10),\n            )","metadata":{"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Dataset class for loading and augmenting spectrograms, Viterbi maps and Viterbi tracks\nclass Vit_Track_Dataset(torch.utils.data.Dataset):\n    \"\"\"\n    dataset = Dataset(data_type, df)\n\n    img, y = dataset[i]\n      img (np.float32): 2 x 360 x 128\n      y (np.float32): label 0 or 1\n    \"\"\"\n    def __init__(self, data_type, df, augment=False, augment_dict=None):\n        self.data_type = data_type\n        self.df = df\n        self.augment = augment\n        \n        if self.augment:\n            self.flip_rate = augment_dict['flip_rate']\n            self.fre_shift_rate = augment_dict['fre_shift_rate']\n            self.time_mask_num = augment_dict['time_mask_num']\n            self.freq_mask_num = augment_dict['freq_mask_num']\n\n    def __len__(self):\n        return len(self.df)\n    \n    def augmentations(self, input):\n        if np.random.rand() <= self.flip_rate: # horizontal flip\n            input = np.flip(input, axis=1).copy()\n        if np.random.rand() <= self.flip_rate: # vertical flip\n            input = np.flip(input, axis=2).copy()\n        if np.random.rand() <= self.fre_shift_rate: # vertical shift\n            input = np.roll(input, np.random.randint(low=0, high=input.shape[0]), axis=1)\n        \n        return input\n    \n    def spect_augmentations(self, spect):\n        for _ in range(self.time_mask_num): # tima masking\n            spect = transforms_time_mask(spect)\n        for _ in range(self.freq_mask_num): # frequency masking\n            spect = transforms_freq_mask(spect)\n        return spect\n        \n\n    def __getitem__(self, i):\n        \"\"\"\n        i (int): get ith data\n        \"\"\"\n        r = self.df.iloc[i]\n        file_id = r.id\n        \n        if self.data_type == 'train':\n            y = torch.from_numpy(np.asarray(r.label)).float()\n            target_type = 'signals' if y == 1 else 'noises'\n        \n            # Load spect\n            filename = '%s/%s/%s' % (TRAIN_SPECT_DIR, target_type, file_id)\n            spect = load_spect(np.load(filename))\n            spect = torch.from_numpy(spect.copy()).float()\n            spect = torch.unsqueeze(spect, 0)\n            spect = self.spect_augmentations(spect) if self.augment else spect\n            spect = spect.cpu().detach().numpy()\n\n            # Load vit map\n            filename = '%s/%s/%s' % (TRAIN_VIT_MAP_DIR, target_type, file_id)\n            vitmap = np.load(filename)\n            vitmap = np.expand_dims(vitmap, 0)\n\n            # Load vit track\n            filename = '%s/%s/%s' % (TRAIN_VIT_TRACK_DIR, target_type, file_id)\n            track = load_track(np.load(filename))\n            track = np.expand_dims(track, 0)\n\n            group = np.concatenate((spect, vitmap, track), axis=0)\n            group = self.augmentations(group) if self.augment else group\n            spect, vitmap, track = np.split(group, 3, axis=0)\n\n            spect = torch.from_numpy(spect.copy()).float()\n            track = torch.from_numpy(track.copy()).float()\n            vitmap = torch.from_numpy(vitmap.copy()).float()\n        else:\n            y = torch.from_numpy(np.asarray(r.target)).float()\n        \n            # Load spect\n            filename = '%s/%s/%s.npy' % (TEST_SPECT_DIR, self.data_type, file_id)\n            spect = load_spect(np.load(filename))\n            spect = torch.from_numpy(spect.copy()).float()\n            spect = torch.unsqueeze(spect, 0)\n            spect = spect.cpu().detach().numpy()\n\n            # Load vit map\n            filename = '%s/%s/%s.npy' % (TEST_VIT_MAP_DIR, self.data_type, file_id)\n            vitmap = np.load(filename)\n            vitmap = np.expand_dims(vitmap, 0)\n\n            # Load vit track\n            filename = '%s/%s/%s.npy' % (TEST_VIT_TRACK_DIR, self.data_type, file_id)\n            track = load_track(np.load(filename))\n            track = np.expand_dims(track, 0)\n\n            spect = torch.from_numpy(spect.copy()).float()\n            track = torch.from_numpy(track.copy()).float()\n            vitmap = torch.from_numpy(vitmap.copy()).float()\n            \n        return spect, vitmap, track, y","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:21:47.507803Z","iopub.execute_input":"2023-01-02T22:21:47.508188Z","iopub.status.idle":"2023-01-02T22:21:47.530322Z","shell.execute_reply.started":"2023-01-02T22:21:47.508156Z","shell.execute_reply":"2023-01-02T22:21:47.529215Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\n\nTwo main architectures were explored: one in which the three features are concatenated into three channels of a single input tensor per sample, which is then classified by a single model and one in which three models each recieve one feature type and combine their outputs. The second architecture produced better results.\n\nAs suggested by Jun Koda in [their notebook](https://www.kaggle.com/code/junkoda/basic-spectrogram-image-classification), EfficientNet was explored as the backbone for both of the architectures explored above. A custom model design achieved higher performance.","metadata":{}},{"cell_type":"code","source":"class ConvBlock(torch.nn.Module):\n    def __init__(self, is_big, in_c, out_c, out_size):\n        super().__init__()\n        self.is_big = is_big\n        \n        self.relu = nn.LeakyReLU()\n        \n        self.conv1 = nn.Conv2d(in_c, out_c, 3, stride=1, padding=1)\n        torch.nn.init.xavier_uniform_(self.conv1.weight)\n        # relu\n        self.conv2 = nn.Conv2d(out_c, out_c, 3, stride=1, padding=1)\n        torch.nn.init.xavier_uniform_(self.conv2.weight)\n        # relu\n        self.conv3 = nn.Conv2d(out_c, out_c, 3, stride=1, padding=1)\n        torch.nn.init.xavier_uniform_(self.conv3.weight)\n        # relu\n        self.pool1 = nn.AdaptiveAvgPool2d(out_size)\n    \n    def forward(self, input):\n        out = self.conv1(input)\n        out = self.relu(out)\n        \n        out = self.conv2(out)\n        out = self.relu(out)\n        \n        if self.is_big:\n            out = self.conv3(out)\n            out = self.relu(out)\n        \n        out = self.pool1(out)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:21:48.043009Z","iopub.execute_input":"2023-01-02T22:21:48.044025Z","iopub.status.idle":"2023-01-02T22:21:48.055326Z","shell.execute_reply.started":"2023-01-02T22:21:48.04399Z","shell.execute_reply":"2023-01-02T22:21:48.054277Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class InputModule(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.relu = nn.LeakyReLU()\n        \n        self.c1 = ConvBlock(False, 1, 32, (45, 180))\n        self.c2 = ConvBlock(False, 32, 64, (23, 90))\n        self.c3 = ConvBlock(True, 64, 128, (12, 45))\n        self.c4 = ConvBlock(True, 128, 256, (6, 23))\n        self.c5 = ConvBlock(True, 256, 256, (3, 12))\n        \n        # flatten\n        self.d1 = nn.Linear(4608, 1152)\n        self.d2 = nn.Linear(1152, 576)\n        \n    def forward(self, input):\n        out = self.c1(input)\n        \n        out = self.c2(out)\n        \n        out = self.c3(out)\n        \n        out = self.c4(out)\n        \n        out = self.c5(out)\n        \n        #out = torch.flatten(out, start_dim = 1)\n        #out = self.d1(out)\n        #out = self.relu(out)\n        \n        #out = self.d2(out)\n        #out = self.relu(out)\n        \n        return out","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:21:48.552184Z","iopub.execute_input":"2023-01-02T22:21:48.554228Z","iopub.status.idle":"2023-01-02T22:21:48.562605Z","shell.execute_reply.started":"2023-01-02T22:21:48.5542Z","shell.execute_reply":"2023-01-02T22:21:48.561733Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Multi_Messenger_Model(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.relu = nn.LeakyReLU()\n        \n        self.track1 = InputModule()\n        self.track2 = InputModule()\n        self.track3 = InputModule()\n        \n        self.c1 = nn.Conv2d(768, 256, 3, stride=1, padding=1)\n        torch.nn.init.xavier_uniform_(self.c1.weight)\n        self.c2 = nn.Conv2d(256, 128, 3, stride=1, padding=1)\n        torch.nn.init.xavier_uniform_(self.c2.weight)\n        \n        # Concatenate\n        \n        #self.d1 = nn.Linear(6912, 2304)\n        self.d1 = nn.Linear(4608, 1)\n    \n    def forward(self, input1, input2, input3):\n        out1 = self.track1(input1)\n        out2 = self.track2(input2)\n        out3 = self.track3(input3)\n        \n        out = torch.cat((out1, out2, out3), dim=1)\n        \n        out = self.c1(out)\n        out = self.relu(out)\n        \n        out = self.c2(out)\n        out = self.relu(out)\n        \n        out = torch.flatten(out, start_dim = 1)\n        out = self.d1(out)\n        #out = self.relu(out)\n        \n        #out = self.d2(out)\n        #out = self.relu(out)\n        \n        return out","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:21:49.077176Z","iopub.execute_input":"2023-01-02T22:21:49.077861Z","iopub.status.idle":"2023-01-02T22:21:49.08852Z","shell.execute_reply.started":"2023-01-02T22:21:49.077816Z","shell.execute_reply":"2023-01-02T22:21:49.087466Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training\n\nThe following code is modified from [Jun Koda's notebook](https://www.kaggle.com/code/junkoda/basic-spectrogram-image-classification). As mentioned above, a Weighted Random Sampler has been used to mitigate the problems caused by the class imbalance.","metadata":{}},{"cell_type":"code","source":"def evaluate(model, loader_val, *, compute_score=True, pbar=None):\n    \"\"\"\n    Predict and compute loss and score\n    \"\"\"\n    tb = time.time()\n    was_training = model.training\n    model.eval()\n\n    loss_sum = 0.0\n    n_sum = 0\n    y_all = []\n    y_pred_all = []\n\n    if pbar is not None:\n        pbar = tqdm(desc='Predict', nrows=78, total=pbar)\n\n    for spect, vitmap, track, y in loader_val:\n        n = y.size(0)\n        spect = spect.to(device)\n        vitmap = vitmap.to(device)\n        track = track.to(device)\n        y = y.to(device)\n        \n        #input = torch.cat((spect, vitmap, track), dim=1)\n\n        with torch.no_grad():\n            y_pred = model(spect, vitmap, track)\n        loss = criterion(y_pred.view(-1), y)\n\n        n_sum += n\n        loss_sum += n * loss.item()\n\n        y_all.append(y.cpu().detach().numpy())\n        y_pred_all.append(y_pred.sigmoid().squeeze().cpu().detach().numpy())\n\n        if pbar is not None:\n            pbar.update(len(spect))\n        \n        del loss, y_pred, spect, vitmap, track, y\n\n    loss_val = loss_sum / n_sum\n\n    y = np.concatenate(y_all)\n    y_pred = np.concatenate(y_pred_all)\n\n    score = roc_auc_score(y, y_pred) if compute_score else None\n\n    ret = {'loss': loss_val,\n           'score': score,\n           'y': y,\n           'y_pred': y_pred,\n           'time': time.time() - tb}\n    \n    model.train(was_training)  # back to train from eval if necessary\n\n    return ret","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:21:49.614172Z","iopub.execute_input":"2023-01-02T22:21:49.614525Z","iopub.status.idle":"2023-01-02T22:21:49.625586Z","shell.execute_reply.started":"2023-01-02T22:21:49.614494Z","shell.execute_reply":"2023-01-02T22:21:49.624365Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/g2net-consolidated-reduced-noise/train_labels.csv', index_col = 0)\ndf = df[df.label >= 0]  # Remove 3 unknowns (target = -1)\n\nnfold = 5\nkfold = KFold(n_splits=nfold, random_state=42, shuffle=True)\n\nepochs = 10\nbatch_size = 100\nnum_workers = 2\nweight_decay = 1e-6\nmax_grad_norm = 1000\n\nlr_max = 0.00005\ndevice = torch.device('cuda')\ncriterion = nn.BCEWithLogitsLoss()\n\nfor ifold, (idx_train, idx_test) in enumerate(kfold.split(df)):\n    print('Fold %d/%d' % (ifold, nfold))\n    torch.manual_seed(42 + ifold + 1)\n\n    # Train - val split\n    augment_dict = {'flip_rate': 0.6,\n                    'fre_shift_rate': 0.6,\n                    'time_mask_num': 2,\n                    'freq_mask_num': 2}\n    dataset_train = Vit_Track_Dataset('train', df.iloc[idx_train], augment=True, augment_dict=augment_dict)\n    dataset_val = Vit_Track_Dataset('train', df.iloc[idx_test])\n    \n    unique_labels, counts = np.unique(df.iloc[idx_train]['label'], return_counts = True)\n    print('Unique labels: {}'.format(unique_labels))\n\n    class_weights = [sum(counts) / c for c in counts]\n    sample_weights = [class_weights[s] for s in df.iloc[idx_train]['label']]\n    sampler = torch.utils.data.WeightedRandomSampler(sample_weights, len(df.iloc[idx_train]['label']), replacement=True)\n\n    loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size,\n                     num_workers=num_workers, pin_memory=True, sampler=sampler, drop_last=True)\n    loader_val = torch.utils.data.DataLoader(dataset_val, batch_size=batch_size,\n                     num_workers=num_workers, pin_memory=True)\n\n    # Model and optimizer\n    model = Multi_Messenger_Model()\n    model.to(device)\n    model.train()\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n    \n    time_val = 0.0\n    lrs = []\n\n    tb = time.time()\n    print('Epoch   loss          score   lr')\n    for iepoch in range(epochs):\n        loss_sum = 0.0\n        n_sum = 0\n\n        # Train\n        for ibatch, (spect, vitmap, track, y) in enumerate(loader_train):\n            n = y.size(0)\n            spect = spect.to(device)\n            vitmap = vitmap.to(device)\n            track = track.to(device)\n            y = y.to(device)\n            \n            #input = torch.cat((spect, vitmap, track), dim=1)\n\n            optimizer.zero_grad()\n\n            #y_pred = model(input)\n            y_pred = model(spect, vitmap, track)\n            loss = criterion(y_pred.view(-1), y)\n\n            loss_train = loss.item()\n            loss_sum += n * loss_train\n            n_sum += n\n\n            loss.backward()\n\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(),\n                                                       max_grad_norm)\n            optimizer.step()\n            lrs.append(optimizer.param_groups[0]['lr'])            \n\n        # Evaluate\n        val = evaluate(model, loader_val)\n        time_val += val['time']\n        loss_train = loss_sum / n_sum\n        lr_now = optimizer.param_groups[0]['lr']\n        dt = (time.time() - tb) / 60\n        print('Epoch %d %.4f %.4f %.4f  %.2e  %.2f min' %\n              (iepoch + 1, loss_train, val['loss'], val['score'], lr_now, dt))\n\n    dt = time.time() - tb\n    print('Training done %.2f min total, %.2f min val' % (dt / 60, time_val / 60))\n\n    # Save model\n    ofilename = 'model_1_%d.pytorch' % ifold\n    torch.save(model.state_dict(), ofilename)\n    print(ofilename, 'written')","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:23:14.521466Z","iopub.execute_input":"2023-01-02T22:23:14.521865Z","iopub.status.idle":"2023-01-02T22:55:38.217227Z","shell.execute_reply.started":"2023-01-02T22:23:14.521831Z","shell.execute_reply":"2023-01-02T22:55:38.215497Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"di = '../input/g2net-detecting-continuous-gravitational-waves'\nsubmit = pd.read_csv(di + '/sample_submission.csv')\nCOLAB = False\nif COLAB == False:\n    # Load model (if necessary)\n    \n    submit['target'] = 0\n    for i in range(nfolds):\n        model = Multi_Messenger_Model()\n        filename = f'model_1_{i}.pytorch'\n        model.to(device)\n        model.load_state_dict(torch.load(filename, map_location=device))\n        model.eval()\n\n        # Predict\n        df = pd.read_csv(TEST_VIT_TRACK_DIR + '/testindex.csv', index_col = 0)\n        dataset_test = Vit_Track_Dataset('test', submit)\n        loader_test = torch.utils.data.DataLoader(dataset_test, batch_size=64,\n                                                num_workers=num_workers, pin_memory=True)\n\n        test = evaluate(model, loader_test, compute_score=False, pbar=len(submit))\n\n        # Write prediction\n        submit['target'] += test['y_pred']/nfolds\nsubmit.to_csv('submission.csv', index=False)\nprint('target range [%.2f, %.2f]' % (submit['target'].min(), submit['target'].max()))","metadata":{"execution":{"iopub.status.busy":"2023-01-02T22:55:38.220469Z","iopub.execute_input":"2023-01-02T22:55:38.2208Z","iopub.status.idle":"2023-01-02T22:55:38.894956Z","shell.execute_reply.started":"2023-01-02T22:55:38.220769Z","shell.execute_reply":"2023-01-02T22:55:38.893186Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]}]}