{"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":"code","source":"!pip install riroriro\n!pip install git+https://github.com/PyFstat/PyFstat@python37","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-11-08T03:13:06.640858Z","iopub.execute_input":"2022-11-08T03:13:06.641377Z","iopub.status.idle":"2022-11-08T03:13:33.117758Z","shell.execute_reply.started":"2022-11-08T03:13:06.641333Z","shell.execute_reply":"2022-11-08T03:13:33.116553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport h5py\nimport gc\nimport glob\nimport math\nimport random\nimport warnings\nimport pyfstat\nimport librosa\nimport librosa.display\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport riroriro.inspiralfuns as ins\nimport riroriro.mergerfirstfuns as me1\nimport riroriro.matchingfuns as mat\nimport riroriro.mergersecondfuns as me2\n\nfrom pathlib import Path\nfrom scipy import stats\nfrom pprint import pprint\nfrom tqdm.notebook import tqdm\nfrom scipy import signal\nfrom scipy.fft import fftshift\nimport matplotlib.pyplot as plt\nfrom IPython.display import HTML, display\nfrom sklearn.model_selection import train_test_split\nimport torch","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:33.120608Z","iopub.execute_input":"2022-11-08T03:13:33.121025Z","iopub.status.idle":"2022-11-08T03:13:36.073082Z","shell.execute_reply.started":"2022-11-08T03:13:33.120983Z","shell.execute_reply":"2022-11-08T03:13:36.071861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ZipDataset(torch.utils.data.Dataset):\n    \"\"\"\n    img, y = dataset[i]\n      img (np.float32): 2 x 360 x 180\n      y (np.float32): label 0 or 1\n    \"\"\"\n    def __init__(self, pos_len=5000, neg_len=9900, path='../input/g2net-generated-data-1/',mod= 100, noise = 0.99):\n        self.path = path\n        self.mod = int(mod)\n        self.noise = noise\n        self.noise_type = 2\n        \n        self.pos_len = pos_len\n        self.neg_len = neg_len\n        self.len = pos_len + neg_len\n        \n        self.mixup = True\n        self.mixup_prob = 0.1\n        self.perm_pos = np.random.permutation(np.arange(self.pos_len))\n        self.perm_neg = np.random.permutation(np.arange(self.neg_len))\n    def __len__(self):\n        return self.len\n    def gen_noise(self, shape):\n        ns = 0.15\n        nr = 0.05\n\n        noise_shape = (360*4140)\n        noise_L_r = self.gen_noise(noise_shape)\n        noise_H_r = noise_L_r*(1-ns)+self.gen_noise(noise_shape)*ns\n        noise_L_i = noise_L_r*(1-nr)+self.gen_noise(noise_shape)*nr\n        noise_H_i = noise_H_r*(1-nr)+self.gen_noise(noise_shape)*nr\n\n        noise_r = np.stack([noise_L_r,noise_H_r]) *1e22\n        noise_i = np.stack([noise_L_i,noise_H_i]) *1e22\n        img_n = noise_r**2 + noise_i**2\n        return img_n\n    def get_negative(self,i):\n        file_name = f'../input/g2net-generated-signals/archive/0_data_{self.mod*(1+(i)//self.mod)}/signals_{i%self.mod}.npy'\n        img = np.load(file_name).astype(np.float64)\n        y=0.0\n        return img, y\n    def get_positive(self, i):\n        file_name = f'../input/g2net-generated-signals/archive/1_data_{self.mod*(1+(i)//self.mod)}/signals_{i%self.mod}.npy'\n        img = np.load(file_name).astype(np.float64)\n        noise_id = int(random.random()*self.neg_len)\n        noise_r = random.random()*0.05+0.95\n        img = (np.sqrt(img)*(1-noise_r)+np.sqrt(self.get_negative(noise_id)[0])*noise_r)**2\n        y=1.0\n        return img, y\n    def get_noise(self):\n        return self.noise\n    def get_mixup(self, i, t):\n        if t==1:\n            mix_img = (self.get_positive(i)[0] + self.get_positive(self.perm_pos[i])[0])/2\n            if random.random() < 1/self.pos_len:\n                self.pos_perm = np.random.permutation(np.arange(self.pos_len))\n        else:\n            mix_img = (self.get_negative(i)[0] + self.get_negative(self.perm_neg[i])[0])/2\n            if random.random() < 1/self.neg_len:\n                self.neg_perm = np.random.permutation(np.arange(self.neg_len))\n        return mix_img, t\n    def __getitem__(self, i):\n        if i<self.pos_len:\n            if self.mixup and random.random() < self.mixup_prob:\n                img, y = self.get_mixup(i,1)\n            else:\n                img, y = self.get_positive(i)\n        else:\n            i = i-self.pos_len\n            if self.mixup and random.random() < self.mixup_prob:\n                img, y = self.get_mixup(i,0)\n            else:\n                img, y = self.get_negative(i)\n        img = ((img)/img.mean() ).astype(np.float32)\n        return img, y","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:36.075298Z","iopub.execute_input":"2022-11-08T03:13:36.075662Z","iopub.status.idle":"2022-11-08T03:13:36.096487Z","shell.execute_reply.started":"2022-11-08T03:13:36.075629Z","shell.execute_reply":"2022-11-08T03:13:36.095222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = ZipDataset()\nimg_label0_1, y = dataset[10050]\nimg_label0_2, y = dataset[10051]\nimg_label1_1,y=dataset[2]\nimg_label1_2,y=dataset[4]","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:36.100035Z","iopub.execute_input":"2022-11-08T03:13:36.100764Z","iopub.status.idle":"2022-11-08T03:13:36.184006Z","shell.execute_reply.started":"2022-11-08T03:13:36.100725Z","shell.execute_reply":"2022-11-08T03:13:36.182849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_label1_2.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:36.185826Z","iopub.execute_input":"2022-11-08T03:13:36.186196Z","iopub.status.idle":"2022-11-08T03:13:36.195537Z","shell.execute_reply.started":"2022-11-08T03:13:36.186157Z","shell.execute_reply":"2022-11-08T03:13:36.194484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class H5Dataset(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):\n        self.data_type = data_type\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        \"\"\"\n        i (int): get ith data\n        \"\"\"\n        r = self.df.iloc[i]\n        y = np.float32(r.target)\n        file_id = r.id\n\n        img = np.empty((2, SHAPE_2[0], SHAPE_1[0]), dtype=np.float32)\n\n        filename = '%s/%s/%s.hdf5' % (di, self.data_type, file_id)\n        with h5py.File(filename, 'r') as f:\n            g = f[file_id]\n\n            for ch, s in enumerate(['H1', 'L1']):\n                a = g[s]['SFTs'][:, :SHAPE_1[0]*SHAPE_1[1]] * 1e22  # Fourier coefficient complex64\n\n                p = a.real**2 + a.imag**2  # power\n                p = np.mean(p.reshape(360, SHAPE_1[0],SHAPE_1[1]), axis=2)\n                p = np.mean(p.reshape(SHAPE_2[0],SHAPE_2[1],SHAPE_1[0]), axis=1)\n\n                img[ch] = p\n        img = ((img)/img.mean()).astype(np.float32)\n        return img, y","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:36.197536Z","iopub.execute_input":"2022-11-08T03:13:36.198283Z","iopub.status.idle":"2022-11-08T03:13:36.209847Z","shell.execute_reply.started":"2022-11-08T03:13:36.198246Z","shell.execute_reply":"2022-11-08T03:13:36.208855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip3 install timm","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:36.21305Z","iopub.execute_input":"2022-11-08T03:13:36.214099Z","iopub.status.idle":"2022-11-08T03:13:46.952045Z","shell.execute_reply.started":"2022-11-08T03:13:36.214053Z","shell.execute_reply":"2022-11-08T03:13:46.950796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport time\nimport h5py\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport random\nimport gc,os,sys,shutil\n\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import KFold,StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom timm.scheduler import CosineLRScheduler\n\nfrom transformers.file_utils import is_torch_tpu_available\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\ncriterion = nn.BCEWithLogitsLoss()\n\n# Train metadata\ndi = '/kaggle/input/g2net-detecting-continuous-gravitational-waves'\ndf = pd.read_csv(di + '/train_labels.csv')\ndf = df[df.target >= 0]  # Remove 3 unknowns (target = -1)","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:46.954317Z","iopub.execute_input":"2022-11-08T03:13:46.954685Z","iopub.status.idle":"2022-11-08T03:13:47.799614Z","shell.execute_reply.started":"2022-11-08T03:13:46.954651Z","shell.execute_reply":"2022-11-08T03:13:47.798438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PRETRAIN = True\nSINGLE_MODEL = True\nimg_size = None\n\nMODEL_NAME = 'efficientnet_b5'\n\n# shapes to resize the images\nSHAPE_1 = (180,23)\nSHAPE_2 = (360,1)","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:47.801285Z","iopub.execute_input":"2022-11-08T03:13:47.801718Z","iopub.status.idle":"2022-11-08T03:13:47.80859Z","shell.execute_reply.started":"2022-11-08T03:13:47.801667Z","shell.execute_reply":"2022-11-08T03:13:47.807495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:47.813155Z","iopub.execute_input":"2022-11-08T03:13:47.813629Z","iopub.status.idle":"2022-11-08T03:13:47.826593Z","shell.execute_reply.started":"2022-11-08T03:13:47.813534Z","shell.execute_reply":"2022-11-08T03:13:47.825555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, name, *, pretrained=True):\n        \"\"\"\n        name (str): timm model name, e.g. tf_efficientnet_b2_ns\n        \"\"\"\n        super().__init__()\n\n        # Use timm\n        model = timm.create_model(name, pretrained=pretrained, in_chans=2)\n\n        clsf = model.default_cfg['classifier']\n        n_features = model._modules[clsf].in_features\n        model._modules[clsf] = nn.Identity()\n\n        self.fc = nn.Linear(n_features, 1)\n        self.model = model\n\n    def forward(self, x):\n        x = self.model(x)\n        x = self.fc(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:47.828505Z","iopub.execute_input":"2022-11-08T03:13:47.829099Z","iopub.status.idle":"2022-11-08T03:13:47.840236Z","shell.execute_reply.started":"2022-11-08T03:13:47.829059Z","shell.execute_reply":"2022-11-08T03:13:47.839027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 img, y in loader_val:\n        n = y.size(0)\n        img = img.to(device)\n        y = y.to(device)\n\n        with torch.no_grad():\n            y_pred = model(img.to(device))\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(img))\n        \n        del loss, y_pred, img, 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":"2022-11-08T03:13:47.841484Z","iopub.execute_input":"2022-11-08T03:13:47.841873Z","iopub.status.idle":"2022-11-08T03:13:47.855805Z","shell.execute_reply.started":"2022-11-08T03:13:47.841833Z","shell.execute_reply":"2022-11-08T03:13:47.854613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"chkpt_batchs = 100\n\nepochs = 2\nbatch_size = 32\nweight_decay = 1e-6\nmax_grad_norm = 1000\n\nlr_max = 1e-4\nepochs_warmup = 1.0\n\ntorch.manual_seed(42)\n\n# Train - val split\ndataset_train = ZipDataset()\ndataset_val = H5Dataset('train', df)\n\nloader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size,\n                 num_workers=0, pin_memory=False, shuffle=True, drop_last=True)\nloader_val = torch.utils.data.DataLoader(dataset_val, batch_size=batch_size,\n                 num_workers=0, pin_memory=False)\n\n# Model and optimizer\nmodel = Model(MODEL_NAME, pretrained=True)\nmodel.to(device)\nmodel.train()\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n\n# Learning-rate schedule\nnbatch = len(loader_train)\nwarmup = epochs_warmup * nbatch  # number of warmup steps\nnsteps = epochs * nbatch        # number of total steps\n\nscheduler = CosineLRScheduler(optimizer,\n              warmup_t=warmup, warmup_lr_init=0.0, warmup_prefix=True, # 1 epoch of warmup\n              t_initial=(nsteps - warmup), lr_min=1e-6)                # 3 epochs of cosine\n\ntime_val = 0.0\n\ntb = time.time()\nbest_loss = 1e10\nprint('Epoch   loss          score   lr')\nfor iepoch in range(epochs):\n    loss_sum = 0.0\n    n_sum = 0\n\n    # Train\n    for ibatch, (img, y) in tqdm(enumerate(loader_train)):\n        n = y.size(0)\n        img = img.to(device)\n        y = y.to(device)\n\n        optimizer.zero_grad()\n\n        y_pred = model(img)\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\n        optimizer.step()\n\n        scheduler.step(iepoch * nbatch + ibatch + 1)\n        \n        if (ibatch%chkpt_batchs==chkpt_batchs-1):\n            val = evaluate(model, loader_val)\n            time_val += val['time']\n            loss_train = loss_sum / n_sum\n            dt = (time.time() - tb) / 60\n            print('Epoch %d %.4f %.4f %.4f  %.2f min' %\n                  (iepoch + 1, loss_train, val['loss'], val['score'],dt))\n            if val['loss'] < best_loss:\n                best_loss = val['loss']\n                # Save model\n                ofilename = 'model_pretrain.pytorch'\n                torch.save(model.state_dict(), ofilename)\n                print(ofilename, 'written')\n            del val\n        gc.collect()\n    val = evaluate(model, loader_val)\n    time_val += val['time']\n    loss_train = loss_sum / n_sum\n    dt = (time.time() - tb) / 60\n    print('Epoch %d %.4f %.4f %.4f  %.2f min' %\n          (iepoch + 1, loss_train, val['loss'], val['score'],dt))\n    if val['loss'] < best_loss:\n        best_loss = val['loss']\n        # Save model\n        ofilename = 'model_pretrain.pytorch'\n        torch.save(model.state_dict(), ofilename)\n        print(ofilename, 'written')\n    del val\n    gc.collect()\n\ndt = time.time() - tb\nprint('Training done %.2f min total, %.2f min val' % (dt / 60, time_val / 60))\n\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:47.857232Z","iopub.execute_input":"2022-11-08T03:13:47.857892Z","iopub.status.idle":"2022-11-08T04:08:08.9535Z","shell.execute_reply.started":"2022-11-08T03:13:47.857845Z","shell.execute_reply":"2022-11-08T04:08:08.950223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load model (if necessary)\nsubmit = pd.read_csv(di + '/sample_submission.csv')\nmodel = Model(MODEL_NAME, pretrained=False)\nmodel.to(device)\n\nif SINGLE_MODEL:\n    model.load_state_dict(torch.load('model_pretrain.pytorch'))\n    model.eval()\n    dataset_test = H5Dataset('test', submit)\n    loader_test = torch.utils.data.DataLoader(dataset_test, batch_size=64,\n                                              num_workers=0, pin_memory=False)\n    test = evaluate(model, loader_test, compute_score=False, pbar=len(submit))\n    submit['target'] += test['y_pred']\nelse:\n    for i in range(N_FOLD):\n        filename = f'model{i}.pytorch'\n        model.load_state_dict(torch.load(filename, map_location=device))\n        model.eval()\n\n        # Predict\n        dataset_test = H5Dataset('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        submit['target'] += test['y_pred']/ N_FOLD\n# Write prediction\nsubmit['target'] = test['y_pred']\nsubmit.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-08T03:13:01.474277Z","iopub.status.idle":"2022-11-08T03:13:01.474884Z","shell.execute_reply.started":"2022-11-08T03:13:01.474603Z","shell.execute_reply":"2022-11-08T03:13:01.474628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}