{"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","metadata":{}},{"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","metadata":{}},{"cell_type":"code","source":"# Use timm pretrained image model\n! pip3 install timm","metadata":{"execution":{"iopub.status.busy":"2022-11-13T03:52:58.204099Z","iopub.execute_input":"2022-11-13T03:52:58.205118Z","iopub.status.idle":"2022-11-13T03:53:09.546686Z","shell.execute_reply.started":"2022-11-13T03:52:58.205033Z","shell.execute_reply":"2022-11-13T03:53:09.545409Z"},"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-13T03:53:09.548784Z","iopub.execute_input":"2022-11-13T03:53:09.549282Z","iopub.status.idle":"2022-11-13T03:53:15.422503Z","shell.execute_reply.started":"2022-11-13T03:53:09.549238Z","shell.execute_reply":"2022-11-13T03:53:15.420779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# params\nPRETRAIN = True\n#SINGLE_MODEL = True\n\nSINGLE_MODEL = False\nIMG_SIZE = None\n#MODEL_NAME = 'tf_efficientnet_b4_ns'\n\n#MODEL_NAME = 'tf_efficientnet_b5_ns'\n\nMODEL_NAME = 'tf_efficientnet_b5_ns'\n\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":{"execution":{"iopub.status.busy":"2022-11-13T03:53:15.425356Z","iopub.execute_input":"2022-11-13T03:53:15.426276Z","iopub.status.idle":"2022-11-13T03:53:15.436817Z","shell.execute_reply.started":"2022-11-13T03:53:15.42624Z","shell.execute_reply":"2022-11-13T03:53:15.435092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"execution":{"iopub.status.busy":"2022-11-13T03:53:15.442479Z","iopub.execute_input":"2022-11-13T03:53:15.4438Z","iopub.status.idle":"2022-11-13T03:53:15.473933Z","shell.execute_reply.started":"2022-11-13T03:53:15.443628Z","shell.execute_reply":"2022-11-13T03:53:15.471828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"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\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-13T03:53:15.476077Z","iopub.execute_input":"2022-11-13T03:53:15.476783Z","iopub.status.idle":"2022-11-13T03:53:15.520652Z","shell.execute_reply.started":"2022-11-13T03:53:15.47674Z","shell.execute_reply":"2022-11-13T03:53:15.517075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize data","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2022-11-13T03:53:15.525728Z","iopub.execute_input":"2022-11-13T03:53:15.527461Z","iopub.status.idle":"2022-11-13T03:53:16.241537Z","shell.execute_reply.started":"2022-11-13T03:53:15.527292Z","shell.execute_reply":"2022-11-13T03:53:16.239597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_real.shape, img_label0_1.shape, img_label0_2.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-13T03:53:16.246967Z","iopub.execute_input":"2022-11-13T03:53:16.248505Z","iopub.status.idle":"2022-11-13T03:53:16.274161Z","shell.execute_reply.started":"2022-11-13T03:53:16.248135Z","shell.execute_reply":"2022-11-13T03:53:16.270425Z"},"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":{"execution":{"iopub.status.busy":"2022-11-13T03:53:16.277602Z","iopub.execute_input":"2022-11-13T03:53:16.281269Z","iopub.status.idle":"2022-11-13T03:53:20.352687Z","shell.execute_reply.started":"2022-11-13T03:53:16.281127Z","shell.execute_reply":"2022-11-13T03:53:20.351628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2022-11-13T03:53:20.354047Z","iopub.execute_input":"2022-11-13T03:53:20.354654Z","iopub.status.idle":"2022-11-13T03:53:20.365963Z","shell.execute_reply.started":"2022-11-13T03:53:20.354618Z","shell.execute_reply":"2022-11-13T03:53:20.364722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and evaluate","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 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":{"execution":{"iopub.status.busy":"2022-11-13T03:53:20.367497Z","iopub.execute_input":"2022-11-13T03:53:20.367862Z","iopub.status.idle":"2022-11-13T03:53:20.393694Z","shell.execute_reply.started":"2022-11-13T03:53:20.367821Z","shell.execute_reply":"2022-11-13T03:53:20.3921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pretrain","metadata":{}},{"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')\n\n'''\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()\n\n'''","metadata":{"execution":{"iopub.status.busy":"2022-11-13T03:53:20.39548Z","iopub.execute_input":"2022-11-13T03:53:20.395953Z","iopub.status.idle":"2022-11-13T03:53:21.995375Z","shell.execute_reply.started":"2022-11-13T03:53:20.395903Z","shell.execute_reply":"2022-11-13T03:53:21.994035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"markdown","source":"Folds","metadata":{}},{"cell_type":"code","source":"'''\nif not SINGLE_MODEL:\n    kfold = StratifiedKFold(n_splits=N_FOLD, random_state=42, shuffle=True)\n    \n    epochs = 6\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, drop_last=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            if val['score']>best_val_score:\n                best_val_score = val['score']\n                ofilename = 'model%d.pytorch' % ifold\n                torch.save(model.state_dict(), ofilename)\n                print(ofilename, 'written')\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)))\n    \n'''","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-11-13T03:53:21.997109Z","iopub.execute_input":"2022-11-13T03:53:21.997477Z","iopub.status.idle":"2022-11-13T03:53:22.008012Z","shell.execute_reply.started":"2022-11-13T03:53:21.997446Z","shell.execute_reply":"2022-11-13T03:53:22.00668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Single Model\n\ntrain 1 epoch with competition data","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2022-11-13T03:53:22.012938Z","iopub.execute_input":"2022-11-13T03:53:22.013879Z","iopub.status.idle":"2022-11-13T03:53:22.027167Z","shell.execute_reply.started":"2022-11-13T03:53:22.013823Z","shell.execute_reply":"2022-11-13T03:53:22.026123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and submit","metadata":{}},{"cell_type":"code","source":"# Load model (if necessary)\nsubmit = pd.read_csv(di + '/sample_submission.csv')\nsubmit['target']=0\n\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'/kaggle/input/generated-5fold/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# Write prediction\nsubmit['target'] =submit['target'] / N_FOLD\nsubmit.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-13T04:09:02.626176Z","iopub.execute_input":"2022-11-13T04:09:02.627129Z","iopub.status.idle":"2022-11-13T09:20:33.684518Z","shell.execute_reply.started":"2022-11-13T04:09:02.62707Z","shell.execute_reply":"2022-11-13T09:20:33.680872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('target range [%.2f, %.2f]' % (submit['target'].min(), submit['target'].max()))","metadata":{"execution":{"iopub.status.busy":"2022-11-13T09:20:33.695184Z","iopub.execute_input":"2022-11-13T09:20:33.697658Z","iopub.status.idle":"2022-11-13T09:20:33.714393Z","shell.execute_reply.started":"2022-11-13T09:20:33.697582Z","shell.execute_reply":"2022-11-13T09:20:33.712702Z"},"trusted":true},"execution_count":null,"outputs":[]}]}