{"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":"# Credits\n\nThis notebook uses [Jun Koda](https://www.kaggle.com/junkoda)'s [spectrogram classification notebook](https://www.kaggle.com/code/junkoda/basic-spectrogram-image-classification/notebook) as a starter. There are minor edits from code by @werus23 in this notebook. The generated signals data is also the one generated by @werus23","metadata":{"papermill":{"duration":0.009018,"end_time":"2022-10-29T18:42:54.77188","exception":false,"start_time":"2022-10-29T18:42:54.762862","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Version\nV1: add custom training dataset to use for pretraining model. Do not train multiple folds on training data as loss/metric is unstable\n\nV2: train folds instead of training single model","metadata":{"papermill":{"duration":0.00604,"end_time":"2022-10-29T18:42:54.78452","exception":false,"start_time":"2022-10-29T18:42:54.77848","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Use timm pretrained image model\n! pip3 install timm","metadata":{"papermill":{"duration":14.907952,"end_time":"2022-10-29T18:43:09.698902","exception":false,"start_time":"2022-10-29T18:42:54.79095","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:13.754816Z","iopub.execute_input":"2022-11-08T15:58:13.755694Z","iopub.status.idle":"2022-11-08T15:58:29.705Z","shell.execute_reply.started":"2022-11-08T15:58:13.755552Z","shell.execute_reply":"2022-11-08T15:58:29.703478Z"},"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":{"papermill":{"duration":4.99002,"end_time":"2022-10-29T18:43:14.696533","exception":false,"start_time":"2022-10-29T18:43:09.706513","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:29.70755Z","iopub.execute_input":"2022-11-08T15:58:29.708011Z","iopub.status.idle":"2022-11-08T15:58:34.128488Z","shell.execute_reply.started":"2022-11-08T15:58:29.707966Z","shell.execute_reply":"2022-11-08T15:58:34.126712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# params\nPRETRAIN = True\nSINGLE_MODEL = False\nIMG_SIZE = None\nMODEL_NAME = 'tf_efficientnet_b5_ns'\n# MODEL_NAME = 'eca_nfnet_l3'\n# MODEL_NAME = 'levit_192'\n# IMG_SIZE = 224 # resize the images for a specific image models\nN_FOLD = 5\n\n# shapes to resize the images\nSHAPE_1 = (180,23)\nSHAPE_2 = (360,1)","metadata":{"papermill":{"duration":0.01824,"end_time":"2022-10-29T18:43:14.722338","exception":false,"start_time":"2022-10-29T18:43:14.704098","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:34.130386Z","iopub.execute_input":"2022-11-08T15:58:34.130842Z","iopub.status.idle":"2022-11-08T15:58:34.138026Z","shell.execute_reply.started":"2022-11-08T15:58:34.130797Z","shell.execute_reply":"2022-11-08T15:58:34.136631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"papermill":{"duration":0.022217,"end_time":"2022-10-29T18:43:14.752084","exception":false,"start_time":"2022-10-29T18:43:14.729867","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:34.141832Z","iopub.execute_input":"2022-11-08T15:58:34.14233Z","iopub.status.idle":"2022-11-08T15:58:34.181038Z","shell.execute_reply.started":"2022-11-08T15:58:34.142277Z","shell.execute_reply":"2022-11-08T15:58:34.179791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"papermill":{"duration":0.007263,"end_time":"2022-10-29T18:43:14.76778","exception":false,"start_time":"2022-10-29T18:43:14.760517","status":"completed"},"tags":[]}},{"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        \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\n    \nclass 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                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                img[ch] = p\n        img = ((img)/img.mean()).astype(np.float32)\n        return img, y","metadata":{"papermill":{"duration":0.040869,"end_time":"2022-10-29T18:43:14.816897","exception":false,"start_time":"2022-10-29T18:43:14.776028","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:34.183029Z","iopub.execute_input":"2022-11-08T15:58:34.183808Z","iopub.status.idle":"2022-11-08T15:58:34.213121Z","shell.execute_reply.started":"2022-11-08T15:58:34.183769Z","shell.execute_reply":"2022-11-08T15:58:34.211453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize data","metadata":{"papermill":{"duration":0.00778,"end_time":"2022-10-29T18:43:14.8322","exception":false,"start_time":"2022-10-29T18:43:14.82442","status":"completed"},"tags":[]}},{"cell_type":"code","source":"dataset = H5Dataset('train', df)\nimg_real, y = dataset[1]\ndataset = 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":{"papermill":{"duration":0.785067,"end_time":"2022-10-29T18:43:15.624932","exception":false,"start_time":"2022-10-29T18:43:14.839865","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:34.216561Z","iopub.execute_input":"2022-11-08T15:58:34.218092Z","iopub.status.idle":"2022-11-08T15:58:35.019626Z","shell.execute_reply.started":"2022-11-08T15:58:34.218034Z","shell.execute_reply":"2022-11-08T15:58:35.018363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_real.shape, img_label0_1.shape, img_label0_2.shape","metadata":{"papermill":{"duration":0.020255,"end_time":"2022-10-29T18:43:15.6531","exception":false,"start_time":"2022-10-29T18:43:15.632845","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:35.02345Z","iopub.execute_input":"2022-11-08T15:58:35.023951Z","iopub.status.idle":"2022-11-08T15:58:35.033922Z","shell.execute_reply.started":"2022-11-08T15:58:35.023907Z","shell.execute_reply":"2022-11-08T15:58:35.032543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1,2, figsize=(8, 8))\nplt.title('real image')\nax[0].set(title=f\"real\")\nax[0].set(title=f\"generated\")\nc0 = ax[0].imshow(img_real[0])\nc1 = ax[1].imshow(img_label1_1[0])\nfig.colorbar(c0, ax=ax[0])\nfig.colorbar(c0, ax=ax[1])\nplt.show()\n\nplt.title('value distribution')\nplt.hist(img_label0_1.flatten(),bins=100,alpha=0.1,color='b')\nplt.hist(img_label0_2.flatten(),bins=100,alpha=0.1,color='b')\nplt.hist(img_label1_1.flatten(),bins=100,alpha=0.1,color='r')\nplt.hist(img_label1_2.flatten(),bins=100,alpha=0.1,color='r')\nplt.hist(img_real.flatten(),bins=100,alpha=0.1,color='g')\nplt.show()","metadata":{"papermill":{"duration":1.671774,"end_time":"2022-10-29T18:43:17.334229","exception":false,"start_time":"2022-10-29T18:43:15.662455","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:35.035576Z","iopub.execute_input":"2022-11-08T15:58:35.035967Z","iopub.status.idle":"2022-11-08T15:58:36.654732Z","shell.execute_reply.started":"2022-11-08T15:58:35.035934Z","shell.execute_reply":"2022-11-08T15:58:36.653423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"papermill":{"duration":0.013102,"end_time":"2022-10-29T18:43:17.361438","exception":false,"start_time":"2022-10-29T18:43:17.348336","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, name, *, pretrained=False):\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,num_classes = 2).to(device)\n        self.use_head=True\n        \n        if name[:9] == 'eca_nfnet':\n            clsf = 'head'\n            n_features = model._modules['head'].fc.in_features\n            model._modules[clsf].fc = nn.Identity()\n        elif name[:15] == 'tf_efficientnet':\n            clsf = model.default_cfg['classifier']\n            n_features = model._modules[clsf].in_features\n            model._modules[clsf] = nn.Identity()\n        else:\n            self.use_head=False\n            #placeholder\n            n_features=1\n        self.fc = nn.Linear(n_features, 1)\n        self.model = model\n\n    def forward(self, x):\n        if IMG_SIZE:\n            x = F.interpolate(x,IMG_SIZE)\n        \n        x = self.model(x)\n        if self.use_head:\n            x = self.fc(x)\n        else:\n            x = x[:,0]\n        return x","metadata":{"papermill":{"duration":0.028956,"end_time":"2022-10-29T18:43:17.403239","exception":false,"start_time":"2022-10-29T18:43:17.374283","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:36.656704Z","iopub.execute_input":"2022-11-08T15:58:36.657145Z","iopub.status.idle":"2022-11-08T15:58:36.669373Z","shell.execute_reply.started":"2022-11-08T15:58:36.657108Z","shell.execute_reply":"2022-11-08T15:58:36.668124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and evaluate","metadata":{"papermill":{"duration":0.013288,"end_time":"2022-10-29T18:43:17.429963","exception":false,"start_time":"2022-10-29T18:43:17.416675","status":"completed"},"tags":[]}},{"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\n        with torch.no_grad():\n            y_pred = model(img)\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        gc.collect()\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    gc.collect()\n    return ret","metadata":{"papermill":{"duration":0.028938,"end_time":"2022-10-29T18:43:17.47208","exception":false,"start_time":"2022-10-29T18:43:17.443142","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:36.672855Z","iopub.execute_input":"2022-11-08T15:58:36.673321Z","iopub.status.idle":"2022-11-08T15:58:36.686943Z","shell.execute_reply.started":"2022-11-08T15:58:36.673281Z","shell.execute_reply":"2022-11-08T15:58:36.685498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pretrain","metadata":{"papermill":{"duration":0.012431,"end_time":"2022-10-29T18:43:17.497412","exception":false,"start_time":"2022-10-29T18:43:17.484981","status":"completed"},"tags":[]}},{"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_best_loss.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_best_loss.pytorch'\n        torch.save(model.state_dict(), ofilename)\n        print(ofilename, 'written')\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":{"papermill":{"duration":17323.878084,"end_time":"2022-10-29T23:32:01.388839","exception":false,"start_time":"2022-10-29T18:43:17.510755","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T15:58:36.688901Z","iopub.execute_input":"2022-11-08T15:58:36.689407Z","iopub.status.idle":"2022-11-08T21:43:14.534367Z","shell.execute_reply.started":"2022-11-08T15:58:36.689366Z","shell.execute_reply":"2022-11-08T21:43:14.529739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"papermill":{"duration":0.014338,"end_time":"2022-10-29T23:32:01.418654","exception":false,"start_time":"2022-10-29T23:32:01.404316","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"Folds","metadata":{"papermill":{"duration":0.013769,"end_time":"2022-10-29T23:32:01.446368","exception":false,"start_time":"2022-10-29T23:32:01.432599","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if not SINGLE_MODEL:\n    kfold = StratifiedKFold(n_splits=N_FOLD, random_state=42, shuffle=True)\n    \n    epochs = 2\n    batch_size = 32\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n\n    lr_max = 4e-4\n    epochs_warmup = 1.0\n\n    losses = []\n    scores = []\n\n    for ifold, (idx_train, idx_test) in enumerate(kfold.split(df,df.target)):\n        print('Fold %d/%d' % (ifold, N_FOLD))\n        torch.manual_seed(42 + ifold + 1)\n\n        # Train - val split\n        dataset_train = H5Dataset('train', df.iloc[idx_train])\n        dataset_val = H5Dataset('train', df.iloc[idx_test])\n\n        loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size,\n                         num_workers=0, pin_memory=False, shuffle=True)\n        loader_val = torch.utils.data.DataLoader(dataset_val, batch_size=batch_size,\n                         num_workers=0, pin_memory=False)\n\n        # Model and optimizer\n    #     model = Model(MODEL_NAME, pretrained=True)\n        model = Model(MODEL_NAME, pretrained=True)\n        if PRETRAIN:\n            model.load_state_dict(torch.load('model_pretrain.pytorch'))\n        model.to(device)\n        model.train()\n\n        optimizer = torch.optim.AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n\n        # Learning-rate schedule\n        nbatch = len(loader_train)\n        warmup = epochs_warmup * nbatch  # number of warmup steps\n        nsteps = epochs * nbatch        # number of total steps\n\n        scheduler = 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\n        time_val = 0.0\n        lrs = []\n\n        tb = time.time()\n        print('Epoch   loss          score   lr')\n\n        best_val_loss = 1e10\n        best_val_score=0\n        for 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                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            if val['loss']<best_val_loss:\n                best_val_loss = val['loss']\n                ofilename = 'model%d.pytorch' % ifold\n                torch.save(model.state_dict(), ofilename)\n                print(ofilename, 'written')\n            if val['score']>best_val_score:\n                best_val_score = val['score']\n        dt = time.time() - tb\n        print('Training done %.2f min total, %.2f min val' % (dt / 60, time_val / 60))\n        losses.append(best_val_loss)\n        scores.append(best_val_score)\n    print('AVG LOSS:', np.mean(np.array(best_val_loss)))\n    print('AVG SCORE:', np.mean(np.array(best_val_score)))","metadata":{"_kg_hide-input":false,"papermill":{"duration":0.049444,"end_time":"2022-10-29T23:32:01.509508","exception":false,"start_time":"2022-10-29T23:32:01.460064","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T21:43:14.541023Z","iopub.execute_input":"2022-11-08T21:43:14.542566Z","iopub.status.idle":"2022-11-08T23:09:48.194309Z","shell.execute_reply.started":"2022-11-08T21:43:14.542512Z","shell.execute_reply":"2022-11-08T23:09:48.189604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Single Model\n\ntrain 1 epoch with competition data","metadata":{"papermill":{"duration":0.013499,"end_time":"2022-10-29T23:32:01.537222","exception":false,"start_time":"2022-10-29T23:32:01.523723","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if SINGLE_MODEL:\n    epochs = 1\n    batch_size = 32\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n    lr_max = 4e-4\n    epochs_warmup = 1.0\n\n    dataset_train = H5Dataset('train',df)\n\n    loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size,\n                     num_workers=0, pin_memory=False, shuffle=True)\n\n    model = Model(MODEL_NAME, pretrained=True)\n    if PRETRAIN:\n        model.load_state_dict(torch.load('model_pretrain.pytorch'))\n    model.to(device)\n    model.train()\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n\n    # Learning-rate schedule\n    nbatch = len(loader_train)\n    warmup = epochs_warmup * nbatch  # number of warmup steps\n    nsteps = epochs * nbatch        # number of total steps\n\n    time_val = 0.0\n\n    tb = time.time()\n    print('Epoch   loss          score   lr')\n\n    for iepoch in range(epochs):\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            optimizer.step()\n    dt = time.time() - tb\n    print('Training done %.2f min total, %.2f min val' % (dt / 60, time_val / 60))","metadata":{"papermill":{"duration":507.654042,"end_time":"2022-10-29T23:40:29.205388","exception":false,"start_time":"2022-10-29T23:32:01.551346","status":"completed"},"tags":[],"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-11-08T23:09:48.202047Z","iopub.execute_input":"2022-11-08T23:09:48.202562Z","iopub.status.idle":"2022-11-08T23:09:48.225167Z","shell.execute_reply.started":"2022-11-08T23:09:48.20252Z","shell.execute_reply":"2022-11-08T23:09:48.223212Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and submit","metadata":{"papermill":{"duration":0.013547,"end_time":"2022-10-29T23:40:29.234204","exception":false,"start_time":"2022-10-29T23:40:29.220657","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Load model (if necessary)\nsubmit = pd.read_csv(di + '/sample_submission.csv')\nsubmit['target'] = 0\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=2, 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.to_csv('submission.csv', index=False)","metadata":{"papermill":{"duration":4247.88371,"end_time":"2022-10-30T00:51:17.131843","exception":false,"start_time":"2022-10-29T23:40:29.248133","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-08T23:09:48.229061Z","iopub.execute_input":"2022-11-08T23:09:48.229518Z","iopub.status.idle":"2022-11-09T01:20:26.801175Z","shell.execute_reply.started":"2022-11-08T23:09:48.229478Z","shell.execute_reply":"2022-11-09T01:20:26.796865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('target range [%.2f, %.2f]' % (submit['target'].min(), submit['target'].max()))","metadata":{"papermill":{"duration":0.036987,"end_time":"2022-10-30T00:51:17.189123","exception":false,"start_time":"2022-10-30T00:51:17.152136","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-09T01:20:26.811738Z","iopub.execute_input":"2022-11-09T01:20:26.813554Z","iopub.status.idle":"2022-11-09T01:20:26.833126Z","shell.execute_reply.started":"2022-11-09T01:20:26.813483Z","shell.execute_reply":"2022-11-09T01:20:26.831984Z"},"trusted":true},"execution_count":null,"outputs":[]}]}