{"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":"# Basic spectrogram image classification with Basic Audio Data Augmentation for G2Net.","metadata":{"papermill":{"duration":0.006279,"end_time":"2022-11-10T15:38:36.489264","exception":false,"start_time":"2022-11-10T15:38:36.482985","status":"completed"},"tags":[]}},{"cell_type":"code","source":"COLAB = False\n\nif COLAB == True:\n    from google.colab import drive\n    drive.mount('/content/drive')\n    %cd '/content/drive/MyDrive/Colab Notebooks/kaggle/G2Net2022/code'","metadata":{"id":"w5UprQs8ga3c","outputId":"5374b854-2ac0-42d5-9c0c-dede89cdaf8a","papermill":{"duration":0.021613,"end_time":"2022-11-10T15:38:36.526325","exception":false,"start_time":"2022-11-10T15:38:36.504712","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:17:38.064065Z","iopub.execute_input":"2022-12-27T04:17:38.064702Z","iopub.status.idle":"2022-12-27T04:17:38.097012Z","shell.execute_reply.started":"2022-12-27T04:17:38.06457Z","shell.execute_reply":"2022-12-27T04:17:38.095918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os \nos.system('python -m pip install --no-index --find-links=/kaggle/input/wheels-download/ timm')\nos.system('python -m pip install --no-index --find-links=/kaggle/input/wheels-download/ torchinfo')\n\nimport gc\nimport cv2\nimport glob\nimport h5py\nimport time\nimport timm\nimport torch\nimport torchaudio\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt \nimport torchvision.transforms as TF\n\n\n\nfrom torch import nn\nfrom tqdm.auto import tqdm\nfrom scipy.stats import norm\nfrom torchinfo import summary\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import roc_auc_score\nfrom timm.scheduler import CosineLRScheduler\n\n\n\ncriterion = nn.BCEWithLogitsLoss()\n\n\nbar = '=='\ndevice = torch.device('cuda') if torch.cuda.is_available() else 'cpu'\nprint(bar*20)\nprint(f'PyTorch Version :{torch.__version__}')\nprint(f'Device :{device}')","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:17:38.098653Z","iopub.execute_input":"2022-12-27T04:17:38.100245Z","iopub.status.idle":"2022-12-27T04:18:05.064055Z","shell.execute_reply.started":"2022-12-27T04:17:38.100205Z","shell.execute_reply":"2022-12-27T04:18:05.062713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.system('nvidia-smi')","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:05.06577Z","iopub.execute_input":"2022-12-27T04:18:05.066669Z","iopub.status.idle":"2022-12-27T04:18:05.07876Z","shell.execute_reply.started":"2022-12-27T04:18:05.06663Z","shell.execute_reply":"2022-12-27T04:18:05.077659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG :\n    model_name = 'tf_efficientnet_b5_ns'#'tf_efficientnetv2_b0' # tf_efficientnet_b7_ns \n    pretrained = True\n    Not_USE_LargeKernel = False # if your use the largekernel it will be 'False';because torchscript errors for it\n    if device == 'cuda' :\n        apex = True\n    else :\n        apex = False\n    \n    ONE_Fold = False\n    nfold = 5\n    epochs = 30#25\n    batch_size = 36\n    num_workers = 1\n    weight_decay = 1e-6\n    max_grad_norm = 1000#1000\n\n    lr_max = 4e-4\n    epochs_warmup = 1.0\n\n\n    ## setting of audio data augmentation \n    beta = 0.4\n    mixup_prob = None#0.01\n    flip_rate = 0.5 # probability of applying the horizontal flip and vertical flip \n    fre_shift_rate = 1.0 # probability of applying the vertical shift\n    time_mask_num = 1 # number of time masking\n    freq_mask_num = 2 # number of frequency masking\n    \n    # csv \n    H1 = pd.read_csv('/kaggle/input/g2net-images/H1.csv')\n    L1 = pd.read_csv('/kaggle/input/g2net-images/L1.csv')\n    ","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:05.081799Z","iopub.execute_input":"2022-12-27T04:18:05.082844Z","iopub.status.idle":"2022-12-27T04:18:05.212729Z","shell.execute_reply.started":"2022-12-27T04:18:05.082796Z","shell.execute_reply":"2022-12-27T04:18:05.21163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(CFG.H1)\ndisplay(CFG.L1)","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:05.214234Z","iopub.execute_input":"2022-12-27T04:18:05.215205Z","iopub.status.idle":"2022-12-27T04:18:05.24579Z","shell.execute_reply.started":"2022-12-27T04:18:05.215171Z","shell.execute_reply":"2022-12-27T04:18:05.244979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"id":"tbekTVh-gYWI","papermill":{"duration":0.005176,"end_time":"2022-11-10T15:38:52.940041","exception":false,"start_time":"2022-11-10T15:38:52.934865","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Train metadata\ndi = '../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-12-27T04:18:05.24709Z","iopub.execute_input":"2022-12-27T04:18:05.247722Z","iopub.status.idle":"2022-12-27T04:18:05.265086Z","shell.execute_reply.started":"2022-12-27T04:18:05.247686Z","shell.execute_reply":"2022-12-27T04:18:05.263869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_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":{"id":"7Z3rNynhq1gr","papermill":{"duration":0.015606,"end_time":"2022-11-10T15:38:52.960722","exception":false,"start_time":"2022-11-10T15:38:52.945116","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:18:05.266494Z","iopub.execute_input":"2022-12-27T04:18:05.266824Z","iopub.status.idle":"2022-12-27T04:18:05.271841Z","shell.execute_reply.started":"2022-12-27T04:18:05.266796Z","shell.execute_reply":"2022-12-27T04:18:05.270928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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, tfms=False, CFG=CFG):\n        self.data_type = data_type\n        self.df = df\n        self.tfms = tfms\n        self.cfg = CFG\n\n    def __len__(self):\n        return len(self.df)\n    \n\n    def __getitem__(self, i):\n        \n        # choice imgaes\n        if np.random.randint(0,2) == 1 :\n            \n            L1_signal = self.cfg.L1['signal_path'].iloc[i]\n            H1_signal = self.cfg.H1['signal_path'].iloc[i]\n        \n            signal_img_h1 = cv2.imread(H1_signal, -1)\n            signal_img_l1 = cv2.imread(L1_signal, -1)\n            img = np.mean([signal_img_h1, signal_img_l1], axis=0)\n            img = np.stack([img, img])\n            \n            y = torch.tensor(1.0, dtype=torch.float32)\n            \n        else :\n            \n            L1_noise = self.cfg.L1['noise_path'].iloc[i]\n            H1_noise = self.cfg.H1['noise_path'].iloc[i]\n        \n            img_h1 = cv2.imread(H1_noise, -1)\n            img_l1 = cv2.imread(L1_noise, -1)\n            img = np.mean([img_h1, img_l1], axis=0)\n            img = np.stack([img, img])\n            \n            y = torch.tensor(0.0, dtype=torch.float32)\n        \n        \n        if self.tfms:\n            if np.random.rand() <= self.cfg.flip_rate: # horizontal flip\n                img = np.flip(img, axis=1).copy()\n            if np.random.rand() <= self.cfg.flip_rate: # vertical flip\n                img = np.flip(img, axis=2).copy()\n            if np.random.rand() <= self.cfg.fre_shift_rate: # vertical shift\n                img = np.roll(img, np.random.randint(low=0, high=img.shape[1]), axis=1)\n            \n            img = torch.from_numpy(img)\n\n            for _ in range(self.cfg.time_mask_num): # tima masking\n                img = transforms_time_mask(img)\n            for _ in range(self.cfg.freq_mask_num): # frequency masking\n                img = transforms_freq_mask(img)\n                img = img.type(torch.float32)\n        \n        else:\n            img = torch.from_numpy(img)\n            img = img.type(torch.float32)\n                \n        return img, y","metadata":{"id":"3kOMMsyagYWN","papermill":{"duration":0.022247,"end_time":"2022-11-10T15:38:52.988165","exception":false,"start_time":"2022-11-10T15:38:52.965918","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:18:55.112133Z","iopub.execute_input":"2022-12-27T04:18:55.112798Z","iopub.status.idle":"2022-12-27T04:18:55.138572Z","shell.execute_reply.started":"2022-12-27T04:18:55.112746Z","shell.execute_reply":"2022-12-27T04:18:55.137223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = Dataset('train', df, True)\nprint(dataset[0][0].shape)\nprint(dataset[0][1])\nplt.figure(figsize=(100, 9))\nplt.imshow(dataset[0][0].numpy()[0])\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:58.015254Z","iopub.execute_input":"2022-12-27T04:18:58.016111Z","iopub.status.idle":"2022-12-27T04:18:58.717727Z","shell.execute_reply.started":"2022-12-27T04:18:58.016061Z","shell.execute_reply":"2022-12-27T04:18:58.716794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"ejqEYZTxgYWP","papermill":{"duration":0.008419,"end_time":"2022-11-10T15:38:57.654463","exception":false,"start_time":"2022-11-10T15:38:57.646044","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class LargeKernel_debias(nn.Conv2d):\n    def forward(self, input: torch.Tensor):\n        \n        #print(f'输入进大内核的数据形状:{input.shape}')\n        finput = input.flatten(0, 1)[:, None]\n        target = abs(self.weight)\n        target = target / target.sum((-1, -2), True)\n        joined_kernel = torch.cat([self.weight, target], 0)\n        reals = target.new_zeros(\n            [1, 1] + [s + p * 2 for p, s in zip(self.padding, input.shape[-2:])]\n        )\n        reals[\n            [slice(None)] * 2 + [slice(p, -p) if p != 0 else slice(None) for p in self.padding]\n        ].fill_(1)\n        \n        output, power = torch.nn.functional.conv2d(\n            finput, joined_kernel, padding=self.padding\n        ).chunk(2, 1)\n        ratio = torch.div(*torch.nn.functional.conv2d(reals, joined_kernel).chunk(2, 1))\n        \n        power_ = torch.mul(power, ratio)\n        output = torch.sub(output,power_)\n        out = output.unflatten(0, input.shape[:2]).flatten(1, 2)\n        #print(f'大内核的输出形状:{out.shape}')\n        #print(f'大内核的内容:{out}')\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:06.059495Z","iopub.execute_input":"2022-12-27T04:18:06.060028Z","iopub.status.idle":"2022-12-27T04:18:06.070763Z","shell.execute_reply.started":"2022-12-27T04:18:06.059995Z","shell.execute_reply":"2022-12-27T04:18:06.069402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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(Model, self).__init__()\n\n        # Use timm\n        model = timm.create_model(name, pretrained=pretrained, in_chans=32)\n        \n        # torch.Size([16, 1, 31, 255])\n        model.conv_stem = nn.Sequential(\n            nn.Identity(),\n            nn.AvgPool2d((1, 9), (1, 8), (0, 4), count_include_pad=False),\n            LargeKernel_debias(1, 16, [31, 255], 1, [31//2, 255//2], 1, 1, False),\n            model.conv_stem,\n        )\n        \n        clsf = model.default_cfg['classifier']\n        n_features = model._modules[clsf].in_features\n        model._modules[clsf] = nn.Identity()\n        self.model = model\n        \n        self.fc = nn.Linear(n_features, 1)\n        \n        \n\n    def forward(self, x):\n        x = self.model(x)\n        x = self.fc(x)\n        return x","metadata":{"id":"W3JQE_UbgYWP","papermill":{"duration":0.01924,"end_time":"2022-11-10T15:38:57.682025","exception":false,"start_time":"2022-11-10T15:38:57.662785","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:18:06.072333Z","iopub.execute_input":"2022-12-27T04:18:06.072768Z","iopub.status.idle":"2022-12-27T04:18:06.088093Z","shell.execute_reply.started":"2022-12-27T04:18:06.072734Z","shell.execute_reply":"2022-12-27T04:18:06.086842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Model(CFG.model_name)\nsummary(model)","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:06.089193Z","iopub.execute_input":"2022-12-27T04:18:06.089582Z","iopub.status.idle":"2022-12-27T04:18:06.313499Z","shell.execute_reply.started":"2022-12-27T04:18:06.08955Z","shell.execute_reply":"2022-12-27T04:18:06.31213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and evaluate","metadata":{"id":"eO6YqnT0gYWQ","papermill":{"duration":0.008209,"end_time":"2022-11-10T15:38:57.698494","exception":false,"start_time":"2022-11-10T15:38:57.690285","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        img = img.to(device)\n        y = y.to(device)\n\n        with torch.no_grad():\n                y_pred = model(img.to(device))\n\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":{"id":"B-xWfjWYgYWR","papermill":{"duration":0.021904,"end_time":"2022-11-10T15:38:57.728574","exception":false,"start_time":"2022-11-10T15:38:57.70667","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:18:06.314838Z","iopub.execute_input":"2022-12-27T04:18:06.315195Z","iopub.status.idle":"2022-12-27T04:18:06.328537Z","shell.execute_reply.started":"2022-12-27T04:18:06.315164Z","shell.execute_reply":"2022-12-27T04:18:06.327386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"id":"3arVeLkRgYWS","papermill":{"duration":0.008011,"end_time":"2022-11-10T15:38:57.74485","exception":false,"start_time":"2022-11-10T15:38:57.736839","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\nkfold = KFold(n_splits=CFG.nfold, random_state=42, shuffle=True)\n\n\nfor ifold, (idx_train, idx_test) in enumerate(kfold.split(df)):\n    print('Fold %d/%d' % (ifold, CFG.nfold))\n    torch.manual_seed(42 + ifold + 1)\n\n    # Train - val split\n    dataset_train = Dataset('train', df.iloc[idx_train], tfms=True, CFG=CFG)\n    dataset_val = Dataset('train', df.iloc[idx_test],CFG=CFG)\n\n    loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=CFG.batch_size,\n                     num_workers=CFG.num_workers, pin_memory=True, shuffle=True, drop_last=True)\n    loader_val = torch.utils.data.DataLoader(dataset_val, batch_size=CFG.batch_size,\n                     num_workers=CFG.num_workers, pin_memory=True)\n\n    # Model and optimizer\n    model = Model(CFG.model_name, pretrained=True)\n    model.to(device)\n    model.train()\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.apex) # 混合精度计算\n    \n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr_max, weight_decay=CFG.weight_decay)\n\n    # Learning-rate schedule\n    nbatch = len(loader_train)\n    warmup = CFG.epochs_warmup * nbatch  # number of warmup steps\n    nsteps = CFG.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    for iepoch in range(CFG.epochs):\n        loss_sum = 0.0\n        n_sum = 0\n\n        # Train\n        for ibatch, (img, y) in enumerate(loader_train):\n            n = y.size(0)\n            img = img.to(device)\n            y = y.to(device)\n\n            #optimizer.zero_grad()\n            with torch.cuda.amp.autocast(enabled=CFG.apex): \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            # 用scaler，scale loss(FP16)，backward得到scaled的梯度(FP16)\n            scaler.scale(loss).backward()\n            #loss.backward()\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(),\n                                                       CFG.max_grad_norm)\n            \n            # scaler 更新参数，会先自动unscale梯度\n            # 如果有nan或inf，自动跳过\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            \n            #optimizer.step()\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    \n    \n    finnal_score = val['score']\n    print(f'=======CV {finnal_score}========')\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_dict = 'model%d.pytorch' % ifold\n    ofilename = 'model%d.bin' % ifold\n    \n    # JIT\n    if CFG.Not_USE_LargeKernel :\n        saved_model = torch.jit.script(model)\n        saved_model.save(ofilename)\n        print('==>> PyTorch JIT Mdoel Saved!!')\n    \n    # Dict\n    torch.save(model.state_dict(), ofilename_dict)\n    print('==>> PyTorch Dict Mdoel Saved!!')\n    \n    if CFG.ONE_Fold :\n        break  # 1 fold only\n    ","metadata":{"id":"tkWJ1eXpgYWS","outputId":"ed0dd2e0-114f-4dd5-c5a6-2d8511e78359","papermill":{"duration":3523.610355,"end_time":"2022-11-10T16:37:41.363719","exception":false,"start_time":"2022-11-10T15:38:57.753364","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:19:12.324797Z","iopub.execute_input":"2022-12-27T04:19:12.325295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.title('LR Schedule: Cosine with linear warmup')\nplt.xlabel('steps')\nplt.ylabel('learning rate')\nplt.plot(lrs)\nplt.show()","metadata":{"id":"gXnX0qHogYWT","papermill":{"duration":0.257289,"end_time":"2022-11-10T16:37:41.63154","exception":false,"start_time":"2022-11-10T16:37:41.374251","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:18:15.101134Z","iopub.status.idle":"2022-12-27T04:18:15.101693Z","shell.execute_reply.started":"2022-12-27T04:18:15.101427Z","shell.execute_reply":"2022-12-27T04:18:15.101447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and submit","metadata":{"id":"xRXxBCwRgYWU","papermill":{"duration":0.010356,"end_time":"2022-11-10T16:37:41.652597","exception":false,"start_time":"2022-11-10T16:37:41.642241","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#if COLAB == False:\n    \n    # Load model (if necessary)\n    #model = Model(CFG.model_name, pretrained=False)\n    #filename = 'model0.pytorch'\n    #model.to(device)\n    #model.load_state_dict(torch.load(filename, map_location=device))\n    #model.eval()\n\n    # Predict\n    #submit = pd.read_csv(di + '/sample_submission.csv')\n    #dataset_test = Dataset('test', submit)\n    #loader_test = torch.utils.data.DataLoader(dataset_test, batch_size=64,\n                                              #num_workers=CFG.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']\n    #submit.to_csv('submission.csv', index=False)\n    #print('target range [%.2f, %.2f]' % (submit['target'].min(), submit['target'].max()))","metadata":{"id":"FMAAA0_pgYWV","papermill":{"duration":1477.026695,"end_time":"2022-11-10T17:02:18.689691","exception":false,"start_time":"2022-11-10T16:37:41.662996","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-27T04:18:15.104974Z","iopub.status.idle":"2022-12-27T04:18:15.1062Z","shell.execute_reply.started":"2022-12-27T04:18:15.105862Z","shell.execute_reply":"2022-12-27T04:18:15.105897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(path) :\n    \n    model = Model(CFG.model_name, pretrained=False)\n    model.to(device)\n    model.load_state_dict(torch.load(path, map_location=device))\n    model.eval()\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:15.107733Z","iopub.status.idle":"2022-12-27T04:18:15.108668Z","shell.execute_reply.started":"2022-12-27T04:18:15.108322Z","shell.execute_reply":"2022-12-27T04:18:15.108373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Predict\nsubmit = pd.read_csv(di + '/sample_submission.csv')\ndataset_test = Dataset('test', submit)\nloader_test = torch.utils.data.DataLoader(dataset_test,\n                                          batch_size=64,\n                                          num_workers=CFG.num_workers,\n                                          pin_memory=True)\n\n\nFOLDS = glob.glob('/kaggle/working/*.pytorch')\nwith torch.no_grad():\n    fold_preds = []\n    for fold in tqdm(FOLDS, desc='Fold prediction'):\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        model = load_model(fold)\n        preds = torch.concat([torch.sigmoid(model(X[0].to(device))).cpu() for X in tqdm(loader_test)], dim=0).numpy()\n        #print(preds)\n        fold_preds.append(preds)\n        \nif CFG.ONE_Fold :\n    preds = np.stack(fold_preds).squeeze()\nelse :\n    preds = np.stack(fold_preds).squeeze().mean(axis=0)","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:15.111322Z","iopub.status.idle":"2022-12-27T04:18:15.112673Z","shell.execute_reply.started":"2022-12-27T04:18:15.112158Z","shell.execute_reply":"2022-12-27T04:18:15.112202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Write prediction\nsubmit['target'] = preds\nsubmit.to_csv('submission.csv', index=False)\nprint('target range [%.2f, %.2f]' % (submit['target'].min(), submit['target'].max()))","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:15.115127Z","iopub.status.idle":"2022-12-27T04:18:15.115745Z","shell.execute_reply.started":"2022-12-27T04:18:15.115466Z","shell.execute_reply":"2022-12-27T04:18:15.115495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit.target.plot.hist()","metadata":{"execution":{"iopub.status.busy":"2022-12-27T04:18:15.118171Z","iopub.status.idle":"2022-12-27T04:18:15.119299Z","shell.execute_reply.started":"2022-12-27T04:18:15.119061Z","shell.execute_reply":"2022-12-27T04:18:15.119095Z"},"trusted":true},"execution_count":null,"outputs":[]}]}