{"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":"TPU = False\n\nimport os\n# os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2022-11-01T14:57:24.743065Z","iopub.execute_input":"2022-11-01T14:57:24.743761Z","iopub.status.idle":"2022-11-01T14:57:24.793915Z","shell.execute_reply.started":"2022-11-01T14:57:24.743621Z","shell.execute_reply":"2022-11-01T14:57:24.792677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use timm pretrained image model\nif TPU:\n    ! pip install torch==1.7.1 torchvision==0.8.2 torchaudio==0.7.2\n    ! pip install cloud-tpu-client==0.10 https://storage.googleapis.com/tpu-pytorch/wheels/torch_xla-1.7-cp37-cp37m-linux_x86_64.whl\n! pip3 install timm\n! pip install git+https://github.com/huggingface/accelerate","metadata":{"execution":{"iopub.status.busy":"2022-11-01T14:57:24.798611Z","iopub.execute_input":"2022-11-01T14:57:24.799032Z","iopub.status.idle":"2022-11-01T14:58:01.810413Z","shell.execute_reply.started":"2022-11-01T14:57:24.798992Z","shell.execute_reply":"2022-11-01T14:58:01.80907Z"},"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 gc\nimport h5py\nimport timm\nimport torch\nif TPU:\n    import torch_xla\n    import torch_xla.core.xla_model as xm\n    import torch_xla.utils.serialization as xser\n\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport random\nimport gc,os,sys,shutil\n\nfrom accelerate import Accelerator, notebook_launcher # main interface, distributed launcher\nfrom accelerate.utils import set_seed # reproducability across devices\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\nif TPU:\n    from transformers.file_utils import is_torch_tpu_available\n\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-01T14:58:01.812271Z","iopub.execute_input":"2022-11-01T14:58:01.81301Z","iopub.status.idle":"2022-11-01T14:58:07.442713Z","shell.execute_reply.started":"2022-11-01T14:58:01.812967Z","shell.execute_reply":"2022-11-01T14:58:07.441693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# device = xm.xla_device() if TPU else torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# torch.set_default_tensor_type('torch.FloatTensor')\n\n# device","metadata":{"execution":{"iopub.status.busy":"2022-11-01T14:58:07.445393Z","iopub.execute_input":"2022-11-01T14:58:07.44576Z","iopub.status.idle":"2022-11-01T14:58:07.450539Z","shell.execute_reply.started":"2022-11-01T14:58:07.445721Z","shell.execute_reply":"2022-11-01T14:58:07.449561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# params\nPRETRAIN = True\nSINGLE_MODEL = True\nIMG_SIZE = None\nMODEL_NAME = 'tf_efficientnet_b7_ns'\nMODEL_FILE = MODEL_NAME+'.pth'\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-01T14:58:07.451832Z","iopub.execute_input":"2022-11-01T14:58:07.452466Z","iopub.status.idle":"2022-11-01T14:58:07.487804Z","shell.execute_reply.started":"2022-11-01T14:58:07.452432Z","shell.execute_reply":"2022-11-01T14:58:07.486771Z"},"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-01T14:58:07.48948Z","iopub.execute_input":"2022-11-01T14:58:07.48989Z","iopub.status.idle":"2022-11-01T14:58:07.515609Z","shell.execute_reply.started":"2022-11-01T14:58:07.489832Z","shell.execute_reply":"2022-11-01T14:58:07.514588Z"},"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-01T14:58:07.517357Z","iopub.execute_input":"2022-11-01T14:58:07.517744Z","iopub.status.idle":"2022-11-01T14:58:08.145302Z","shell.execute_reply.started":"2022-11-01T14:58:07.517712Z","shell.execute_reply":"2022-11-01T14:58:08.14423Z"},"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-01T14:58:08.146561Z","iopub.execute_input":"2022-11-01T14:58:08.146916Z","iopub.status.idle":"2022-11-01T14:58:08.155716Z","shell.execute_reply.started":"2022-11-01T14:58:08.146879Z","shell.execute_reply":"2022-11-01T14:58:08.154614Z"},"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-01T14:58:08.15724Z","iopub.execute_input":"2022-11-01T14:58:08.15784Z","iopub.status.idle":"2022-11-01T14:58:09.80166Z","shell.execute_reply.started":"2022-11-01T14:58:08.157804Z","shell.execute_reply":"2022-11-01T14:58:09.799574Z"},"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_b_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-01T14:58:09.806131Z","iopub.execute_input":"2022-11-01T14:58:09.806472Z","iopub.status.idle":"2022-11-01T14:58:09.816872Z","shell.execute_reply.started":"2022-11-01T14:58:09.806443Z","shell.execute_reply":"2022-11-01T14:58:09.815363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and evaluate","metadata":{}},{"cell_type":"code","source":"def evaluate(model, loader_val, accelerator=None,  *, compute_score=True, pbar=None):\n    \"\"\"\n    Predict and compute loss and score\n    \"\"\"\n    tb = time.time()\n    was_training = model.training\n    \n    model.eval()\n\n    loss_sum = 0.0\n    n_sum = 0\n    y_all = []\n    y_pred_all = []\n    # metrics = []\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        gc.collect()\n        # torch.cuda.empty_cache()\n        \n        # img = img.to(device)\n        # y = y.to(device)\n        \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        if accelerator is not None:\n            _, _ = accelerator.gather_for_metrics((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-01T16:38:03.099263Z","iopub.execute_input":"2022-11-01T16:38:03.099685Z","iopub.status.idle":"2022-11-01T16:38:03.112714Z","shell.execute_reply.started":"2022-11-01T16:38:03.099655Z","shell.execute_reply":"2022-11-01T16:38:03.111411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pretrain","metadata":{}},{"cell_type":"code","source":"def pretraining_loop(mixed_precision:str=\"fp16\", seed:int=42, batch_size:int=64):\n    # Set our seed\n    set_seed(42)\n    \n    # Initialize the Accelerator, our main interface\n    # it will handle mixed precision automatically for us\n    accelerator = Accelerator(mixed_precision=mixed_precision)\n    \n    chkpt_batchs = 100\n\n    epochs = 3\n    batch_size = 32\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n    num_workers = 2\n\n    lr_max = 1e-4\n    epochs_warmup = 1\n\n    torch.manual_seed(42)\n\n    # Train - val split\n    dataset_train = ZipDataset()\n    dataset_val = H5Dataset('train', df)\n\n    loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size,\n                     num_workers=num_workers, pin_memory=True, shuffle=True, drop_last=True)\n    loader_val = torch.utils.data.DataLoader(dataset_val, batch_size=batch_size,\n                     num_workers=num_workers, pin_memory=True)\n    \n    with accelerator.main_process_first():\n        # Model and optimizer\n        model = Model(MODEL_NAME, pretrained=True)\n        # if os.path.exists(MODEL_FILE):\n        #    model.load_state_dict(torch.load(\"model_pretrain.pytorch\"))\n\n        # model.to(device)\n        \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    \n    model, optimizer, loader_train, loader_val, scheduler = accelerator.prepare(model, optimizer, loader_train, loader_val, scheduler)\n\n    time_val = 0.0\n\n    tb = time.time()\n    best_loss = 1e10\n    print('Epoch   loss          score   lr')\n    for iepoch in range(epochs):\n        model.train()\n        \n        loss_sum = 0.0\n        n_sum = 0\n\n        # Train\n        for ibatch, (img, y) in tqdm(enumerate(loader_train)):\n            gc.collect()\n            torch.cuda.empty_cache()\n\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            accelerator.backward(loss)\n\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(),\n                                                       max_grad_norm)\n            if TPU:\n                xm.optimizer_step(optimizer, barrier=True)\n            else:\n                optimizer.step()\n\n            scheduler.step(iepoch * nbatch + ibatch + 1)\n\n            if (ibatch%chkpt_batchs==chkpt_batchs-1):\n                val = evaluate(accelerator, model, loader_val)\n                time_val += val['time']\n                loss_train = loss_sum / n_sum\n                dt = (time.time() - tb) / 60\n\n                print('Epoch %d %.4f %.4f %.4f  %.2f min' %\n                      (iepoch + 1, loss_train, val['loss'], val['score'],dt))\n\n                if val['loss'] < best_loss:\n                    best_loss = val['loss']\n                    # Save model\n                    ofilename = \"Pretrain_HF\"+MODEL_FILE\n                    accelerator.save(accelerator.unwrap_model(model).state_dict(), MODEL_FILE)\n                    print(ofilename, 'written')\n                del val\n\n            gc.collect()\n\n        val = evaluate(accelerator, model, loader_val)\n        \n        time_val += val['time']\n        \n        loss_train = loss_sum / n_sum\n        \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        \n        accelerator.wait_for_everyone()\n        \n        if val['loss'] < best_loss:\n            best_loss = val['loss']\n            # Save model\n            \n            model = accelerator.unwrap_model(model)\n            \n            ofilename = MODEL_FILE\n            # torch.save(model.state_dict(), MODEL_FILE)\n            accelerator.save(accelerator.unwrap_model(model).state_dict(), \"Pretrain_HF\"+MODEL_FILE)\n            \n            print(ofilename, 'written')\n            \n        del val\n        gc.collect()\n\n    dt = time.time() - tb\n    print('Training done %.2f min total, %.2f min val' % (dt / 60, time_val / 60))\n    accelerator.save(accelerator.unwrap_model(model).state_dict(), \"Pretrain_HF\"+MODEL_FILE)\n    torch.save(model.state_dict(), \"backup.pth\")\n    \n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-11-01T15:18:58.367623Z","iopub.execute_input":"2022-11-01T15:18:58.368042Z","iopub.status.idle":"2022-11-01T15:18:58.39115Z","shell.execute_reply.started":"2022-11-01T15:18:58.368002Z","shell.execute_reply":"2022-11-01T15:18:58.390191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"notebook_launcher(pretraining_loop, (\"fp16\", 42, 64), num_processes=2)","metadata":{"execution":{"iopub.status.busy":"2022-11-01T15:18:59.012062Z","iopub.execute_input":"2022-11-01T15:18:59.012731Z","iopub.status.idle":"2022-11-01T15:56:12.020903Z","shell.execute_reply.started":"2022-11-01T15:18:59.012699Z","shell.execute_reply":"2022-11-01T15:56:12.018873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"markdown","source":"# Single Model\n\ntrain 1 epoch with competition data","metadata":{}},{"cell_type":"code","source":"def finetuning_loop(mixed_precision=\"fp16\", seed=42, batch_size=64):\n    \n    set_seed(42)\n    accelerator = Accelerator(mixed_precision=mixed_precision)\n    \n    if SINGLE_MODEL:\n        epochs = 2\n        batch_size = 16\n        weight_decay = 1e-6\n        max_grad_norm = 1000\n        lr_max = 4e-4\n        num_workers=2\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=num_workers, pin_memory=True, shuffle=True)\n        with accelerator.main_process_first():\n            model = Model(MODEL_NAME, pretrained=True)\n\n            if PRETRAIN:\n                model.load_state_dict(torch.load(\"Pretrain_HF\"+MODEL_FILE))\n\n        # model.to(device)\n        model.train()\n\n        optimizer = torch.optim.AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n        \n        model, optimizer, loader_train = accelerator.prepare(model, optimizer, loader_train)\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            loss_sum = 0.0\n            n_sum = 0\n            # Train\n            for ibatch, (img, y) in tqdm(enumerate(loader_train)):\n                gc.collect()\n                torch.cuda.empty_cache()\n\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                accelerator.backward(loss)\n\n                grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(),\n                                                           max_grad_norm)\n                optimizer.step()\n        accelerator.save(accelerator.unwrap_model(model).state_dict(), \"FineTune_HF\"+MODEL_FILE)\n        dt = time.time() - tb\n        print('Training done %.2f min total, %.2f min val' % (dt / 60, time_val / 60))\n\n    else:\n        kfold = StratifiedKFold(n_splits=N_FOLD, random_state=42, shuffle=True)\n\n        epochs = 16\n        batch_size = 32\n        weight_decay = 1e-6\n        max_grad_norm = 1000\n        num_workers = 2\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=num_workers, pin_memory=True, shuffle=True, drop_last=True)\n            loader_val = torch.utils.data.DataLoader(dataset_val, batch_size=batch_size,\n                             num_workers=num_workers, pin_memory=True)\n\n            # Model and optimizer\n            \n            model = Model(MODEL_NAME, pretrained=True)\n            if PRETRAIN:\n                model.load_state_dict(torch.load(\"Pretrain_HF\"+MODEL_FILE))\n\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            model, optimizer, loader_train, loader_val, scheduler = accelerator.prepare(model, optimizer, loader_train, loader_val, scheduler)\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            loss_sum = 0.0\n            n_sum = 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                    gc.collect()\n                    torch.cuda.empty_cache()\n\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                    accelerator.backward(loss)\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(accelerator, 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_FILE% ifold\n                    # torch.save(model.state_dict(), ofilename)\n                    accelerator.save(accelerator.unwrap_model(model).state_dict(), \"Finetune_HF\"+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            \n            accelerator.save(model, \"Finetune_HF\"+MODEL_FILE)\n            \n        print('AVG LOSS:', np.mean(np.array(best_val_loss)))\n        print('AVG SCORE:', np.mean(np.array(best_val_score)))\n        ","metadata":{"execution":{"iopub.status.busy":"2022-11-01T16:22:57.32195Z","iopub.execute_input":"2022-11-01T16:22:57.322372Z","iopub.status.idle":"2022-11-01T16:22:57.351062Z","shell.execute_reply.started":"2022-11-01T16:22:57.322339Z","shell.execute_reply":"2022-11-01T16:22:57.349911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"notebook_launcher(finetuning_loop, (\"fp16\", 42, 64), num_processes=2)","metadata":{"execution":{"iopub.status.busy":"2022-11-01T16:22:57.803232Z","iopub.execute_input":"2022-11-01T16:22:57.803599Z","iopub.status.idle":"2022-11-01T16:25:59.018934Z","shell.execute_reply.started":"2022-11-01T16:22:57.803568Z","shell.execute_reply":"2022-11-01T16:25:59.015796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and submit","metadata":{}},{"cell_type":"code","source":"#def prediction_loop(mixed_precision=\"fp16\", seed=42, batch_size=64):\n    # Load model (if necessary)\nsubmit = pd.read_csv(di + '/sample_submission.csv')\n\n#set_seed(42)\n\n#accelerator = Accelerator(mixed_precision=mixed_precision)\n\n#with accelerator.main_process_first():\nmodel = Model(MODEL_NAME, pretrained=False)\n\nnum_workers=2\n\n\nif SINGLE_MODEL:\n#    with accelerator.main_process_first():\n    model.load_state_dict(torch.load('./FineTune_HFtf_efficientnet_b7_ns.pth'))\n\n    model.eval()\n    dataset_test = H5Dataset('test', submit)\n\n    loader_test = torch.utils.data.DataLoader(dataset_test, batch_size=16,\n                                              num_workers=num_workers, pin_memory=True)\n\n    #model, loader_test = accelerator.prepare(model, loader_test)\n\n    test = evaluate(model, loader_test, accelerator=None, compute_score=False, pbar=len(submit))\n    submit['target'] += test['y_pred']\n\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 = nn.DataParallel(model, device_ids=[0])\n\n        model.eval()\n\n        # Predict\n        dataset_test = H5Dataset('test', submit)\n        loader_test = torch.utils.data.DataLoader(dataset_test, batch_size=16,\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-01T16:39:17.855968Z","iopub.execute_input":"2022-11-01T16:39:17.856516Z","iopub.status.idle":"2022-11-01T18:27:54.917638Z","shell.execute_reply.started":"2022-11-01T16:39:17.856474Z","shell.execute_reply":"2022-11-01T18:27:54.915404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"backup_finetune.pth\")","metadata":{"execution":{"iopub.status.busy":"2022-11-01T18:27:55.198415Z","iopub.execute_input":"2022-11-01T18:27:55.200841Z","iopub.status.idle":"2022-11-01T18:27:55.652931Z","shell.execute_reply.started":"2022-11-01T18:27:55.20081Z","shell.execute_reply":"2022-11-01T18:27:55.651885Z"},"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-01T18:27:55.654525Z","iopub.execute_input":"2022-11-01T18:27:55.655126Z","iopub.status.idle":"2022-11-01T18:27:55.665766Z","shell.execute_reply.started":"2022-11-01T18:27:55.655087Z","shell.execute_reply":"2022-11-01T18:27:55.664584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}