{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"pip install h5py","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:24.28455Z","iopub.execute_input":"2022-12-06T17:53:24.285088Z","iopub.status.idle":"2022-12-06T17:53:37.227409Z","shell.execute_reply.started":"2022-12-06T17:53:24.284956Z","shell.execute_reply":"2022-12-06T17:53:37.226111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:37.230855Z","iopub.execute_input":"2022-12-06T17:53:37.23119Z","iopub.status.idle":"2022-12-06T17:53:47.748174Z","shell.execute_reply.started":"2022-12-06T17:53:37.231161Z","shell.execute_reply":"2022-12-06T17:53:47.746922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom scipy.stats import norm\nimport h5py\nimport timm\nimport matplotlib.pyplot as plt\nimport seaborn\nimport time\n\nimport torch\nimport torchaudio\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom torch.optim.lr_scheduler import StepLR\n\nfrom tqdm import tqdm\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom timm.scheduler import CosineLRScheduler\n\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndi = '/kaggle/input/g2net-detecting-continuous-gravitational-waves'\n\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:47.750172Z","iopub.execute_input":"2022-12-06T17:53:47.750578Z","iopub.status.idle":"2022-12-06T17:53:51.83884Z","shell.execute_reply.started":"2022-12-06T17:53:47.750525Z","shell.execute_reply":"2022-12-06T17:53:51.837949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = pd.read_csv('../input/g2net-detecting-continuous-gravitational-waves/train_labels.csv')\nsubmission = pd.read_csv('../input/g2net-detecting-continuous-gravitational-waves/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.841823Z","iopub.execute_input":"2022-12-06T17:53:51.84275Z","iopub.status.idle":"2022-12-06T17:53:51.86987Z","shell.execute_reply.started":"2022-12-06T17:53:51.842711Z","shell.execute_reply":"2022-12-06T17:53:51.868934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.871828Z","iopub.execute_input":"2022-12-06T17:53:51.87221Z","iopub.status.idle":"2022-12-06T17:53:51.890806Z","shell.execute_reply.started":"2022-12-06T17:53:51.872174Z","shell.execute_reply":"2022-12-06T17:53:51.889973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels['target'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.892095Z","iopub.execute_input":"2022-12-06T17:53:51.892514Z","iopub.status.idle":"2022-12-06T17:53:51.908051Z","shell.execute_reply.started":"2022-12-06T17:53:51.892477Z","shell.execute_reply":"2022-12-06T17:53:51.907072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Removing the negative labels\ntrain_labels = train_labels[train_labels.target>=0]\ntrain_labels.target.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.909346Z","iopub.execute_input":"2022-12-06T17:53:51.913172Z","iopub.status.idle":"2022-12-06T17:53:51.923648Z","shell.execute_reply.started":"2022-12-06T17:53:51.913144Z","shell.execute_reply":"2022-12-06T17:53:51.922729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.925366Z","iopub.execute_input":"2022-12-06T17:53:51.925888Z","iopub.status.idle":"2022-12-06T17:53:51.935585Z","shell.execute_reply.started":"2022-12-06T17:53:51.925848Z","shell.execute_reply":"2022-12-06T17:53:51.934561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_train = '../input/g2net-detecting-continuous-gravitational-waves/train/001121a05.hdf5'","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.937336Z","iopub.execute_input":"2022-12-06T17:53:51.938223Z","iopub.status.idle":"2022-12-06T17:53:51.944915Z","shell.execute_reply.started":"2022-12-06T17:53:51.938183Z","shell.execute_reply":"2022-12-06T17:53:51.943681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f = h5py.File(example_train, 'r')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.949623Z","iopub.execute_input":"2022-12-06T17:53:51.949905Z","iopub.status.idle":"2022-12-06T17:53:51.961546Z","shell.execute_reply.started":"2022-12-06T17:53:51.949881Z","shell.execute_reply":"2022-12-06T17:53:51.960533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f.keys()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.963412Z","iopub.execute_input":"2022-12-06T17:53:51.964242Z","iopub.status.idle":"2022-12-06T17:53:51.976125Z","shell.execute_reply.started":"2022-12-06T17:53:51.964197Z","shell.execute_reply":"2022-12-06T17:53:51.975084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"group1 = f['001121a05']","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.977649Z","iopub.execute_input":"2022-12-06T17:53:51.978312Z","iopub.status.idle":"2022-12-06T17:53:51.986189Z","shell.execute_reply.started":"2022-12-06T17:53:51.978279Z","shell.execute_reply":"2022-12-06T17:53:51.985262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"group1","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:51.98775Z","iopub.execute_input":"2022-12-06T17:53:51.988453Z","iopub.status.idle":"2022-12-06T17:53:52.001092Z","shell.execute_reply.started":"2022-12-06T17:53:51.988419Z","shell.execute_reply":"2022-12-06T17:53:52.000081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"group1.keys()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.002596Z","iopub.execute_input":"2022-12-06T17:53:52.00333Z","iopub.status.idle":"2022-12-06T17:53:52.013746Z","shell.execute_reply.started":"2022-12-06T17:53:52.003295Z","shell.execute_reply":"2022-12-06T17:53:52.012674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"h1 = group1['H1']","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.015403Z","iopub.execute_input":"2022-12-06T17:53:52.016164Z","iopub.status.idle":"2022-12-06T17:53:52.022442Z","shell.execute_reply.started":"2022-12-06T17:53:52.016129Z","shell.execute_reply":"2022-12-06T17:53:52.021569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"h1.keys()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.023938Z","iopub.execute_input":"2022-12-06T17:53:52.024843Z","iopub.status.idle":"2022-12-06T17:53:52.036929Z","shell.execute_reply.started":"2022-12-06T17:53:52.024805Z","shell.execute_reply":"2022-12-06T17:53:52.03607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SFT_1 = h1['SFTs']","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.038133Z","iopub.execute_input":"2022-12-06T17:53:52.038418Z","iopub.status.idle":"2022-12-06T17:53:52.048756Z","shell.execute_reply.started":"2022-12-06T17:53:52.038393Z","shell.execute_reply":"2022-12-06T17:53:52.047851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SFT_1.shape","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.05047Z","iopub.execute_input":"2022-12-06T17:53:52.050869Z","iopub.status.idle":"2022-12-06T17:53:52.074398Z","shell.execute_reply.started":"2022-12-06T17:53:52.05083Z","shell.execute_reply":"2022-12-06T17:53:52.073411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SFT_1[0:2]","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.076912Z","iopub.execute_input":"2022-12-06T17:53:52.07757Z","iopub.status.idle":"2022-12-06T17:53:52.090057Z","shell.execute_reply.started":"2022-12-06T17:53:52.077534Z","shell.execute_reply":"2022-12-06T17:53:52.088914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with h5py.File(example_train, 'r') as f:\n    \n    group_1_list = list(f.keys())\n    print(\"First layer groups:\", group_1_list)\n    \n    group_1 = f[group_1_list[0]]\n    group_2_list = list(group_1.keys())\n    print(\"Second layer groups:\", group_2_list)\n    \n    group_2 = group_1[group_2_list[0]]\n    group_3_list = list(group_2.keys())\n    print(\"Third layer groups:\", group_3_list)\n    \n    print(\"key:\", group_2_list[0], \", shape:\", group_1[group_2_list[0]]['SFTs'].shape)\n    print(\"key:\", group_2_list[1], \", shape:\", group_1[group_2_list[1]]['SFTs'].shape)\n    print(\"key:\", group_2_list[2], \", shape:\", group_1[group_2_list[2]].shape)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.091628Z","iopub.execute_input":"2022-12-06T17:53:52.092566Z","iopub.status.idle":"2022-12-06T17:53:52.104683Z","shell.execute_reply.started":"2022-12-06T17:53:52.092527Z","shell.execute_reply":"2022-12-06T17:53:52.1033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count = 0\ntrain_files = []\ntest_files = []\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        count += 1\n        path = os.path.join(dirname, filename)\n        \n        if 'test' in dirname:\n            test_files.append(path)\n            \n        if 'train' in dirname:\n            train_files.append(path)\n            \n        if count%1000 == 0:\n            print(count, 'data files loaded')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:52.106573Z","iopub.execute_input":"2022-12-06T17:53:52.107325Z","iopub.status.idle":"2022-12-06T17:53:54.177499Z","shell.execute_reply.started":"2022-12-06T17:53:52.107287Z","shell.execute_reply":"2022-12-06T17:53:54.176395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files[0]","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.178905Z","iopub.execute_input":"2022-12-06T17:53:54.179614Z","iopub.status.idle":"2022-12-06T17:53:54.188949Z","shell.execute_reply.started":"2022-12-06T17:53:54.179575Z","shell.execute_reply":"2022-12-06T17:53:54.18607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#defining a configuration\nclass CFG:\n    model_name = 'tf_efficientnet_b5_ns'\n    target_size = 1\n    transform = True\n    flip_rate = 0.5\n    fre_shift_rate = 1.0\n    time_mask_num = 1\n    freq_mask_num = 2\n    nfold = 5\n    is_cross_validate = True\n    batch_size = 32\n    epochs = 25\n    num_workers = 2\n    lr = 1e-3\n    weight_decay = 1e-6\n    train = True\n    seed = 42\n    score_method = 'roc_auc_score'\n    scheduler_type = 'CosineLRScheduler'\n    optimizer_type = 'AdamW'\n    loss_type = 'BCEWithLogitsLoss'\n    max_grad_norm = 1000\n    lr_max = 4e-4\n    epochs_warmup = 1.0","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.190728Z","iopub.execute_input":"2022-12-06T17:53:54.191647Z","iopub.status.idle":"2022-12-06T17:53:54.199072Z","shell.execute_reply.started":"2022-12-06T17:53:54.191612Z","shell.execute_reply":"2022-12-06T17:53:54.198134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_criterion():\n    if CFG.loss_type == 'CrossEntropyLoss':\n        return nn.CrossEntropyLoss()\n    if CFG.loss_type == 'BCEWithLogitsLoss':\n        return nn.BCEWithLogitsLoss()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.200783Z","iopub.execute_input":"2022-12-06T17:53:54.201508Z","iopub.status.idle":"2022-12-06T17:53:54.21029Z","shell.execute_reply.started":"2022-12-06T17:53:54.201472Z","shell.execute_reply":"2022-12-06T17:53:54.209164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_optimizer(model):\n    if CFG.optimizer_type == 'Adam':\n        optimizer = torch.optim.Adam(model.parameters(), lr=CFG.lr_max, weight_decay=CFG.weight_decay, amsgrad=False)\n    if CFG.optimizer_type == 'AdamW':\n        optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr_max, weight_decay=CFG.weight_decay)\n    return optimizer","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.211787Z","iopub.execute_input":"2022-12-06T17:53:54.212436Z","iopub.status.idle":"2022-12-06T17:53:54.222306Z","shell.execute_reply.started":"2022-12-06T17:53:54.212402Z","shell.execute_reply":"2022-12-06T17:53:54.221294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_scheduler(optimizer, warmup, nsteps):\n    if CFG.scheduler_type == 'StepLR':\n        scheduler = StepLR(optimizer, step_size=2, gamma=0.1, verbose=True)\n    if CFG.scheduler_type == 'CosineLRScheduler':\n        scheduler = CosineLRScheduler(optimizer,\n                                      warmup_t=warmup, warmup_lr_init=0.0, warmup_prefix=True,\n                                      t_initial=(nsteps - warmup), lr_min=1e-6) \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.223793Z","iopub.execute_input":"2022-12-06T17:53:54.224474Z","iopub.status.idle":"2022-12-06T17:53:54.238103Z","shell.execute_reply.started":"2022-12-06T17:53:54.224415Z","shell.execute_reply":"2022-12-06T17:53:54.237061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_score(y_true, y_pred):\n    if CFG.score_method == \"roc_auc_score\":\n        score = roc_auc_score(y_true, y_pred)\n    if CFG.score_method == \"accuracy_score\":\n        score = accuracy_score(y_true, y_pred)\n    return score","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.239619Z","iopub.execute_input":"2022-12-06T17:53:54.240231Z","iopub.status.idle":"2022-12-06T17:53:54.249315Z","shell.execute_reply.started":"2022-12-06T17:53:54.240196Z","shell.execute_reply":"2022-12-06T17:53:54.248374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def normalize(X):\n    X = (X[..., None].view(X.real.dtype) ** 2).sum(-1)\n    POS = int(X.size * 0.99903)\n    EXP = norm.ppf((POS + 0.4) / (X.size + 0.215))\n    scale = np.partition(X.flatten(), POS, -1)[POS]\n    X /= scale / EXP.astype(scale.dtype) ** 2\n    return X","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.255063Z","iopub.execute_input":"2022-12-06T17:53:54.25573Z","iopub.status.idle":"2022-12-06T17:53:54.262329Z","shell.execute_reply.started":"2022-12-06T17:53:54.255696Z","shell.execute_reply":"2022-12-06T17:53:54.261325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transform(img):\n    transforms_time_mask = nn.Sequential(\n                torchaudio.transforms.TimeMasking(time_mask_param=10),\n            )\n    transforms_freq_mask = nn.Sequential(\n                torchaudio.transforms.FrequencyMasking(freq_mask_param=10),\n            )\n    if np.random.rand() <= CFG.flip_rate: # horizontal flip\n        img = np.flip(img, axis=1).copy()\n    if np.random.rand() <= CFG.flip_rate: # vertical flip\n        img = np.flip(img, axis=2).copy()\n    if np.random.rand() <= 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(CFG.time_mask_num): # tima masking\n        img = transforms_time_mask(img)\n    for _ in range(CFG.freq_mask_num): # frequency masking\n        img = transforms_freq_mask(img)\n        \n    return img","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.26361Z","iopub.execute_input":"2022-12-06T17:53:54.264506Z","iopub.status.idle":"2022-12-06T17:53:54.274819Z","shell.execute_reply.started":"2022-12-06T17:53:54.264468Z","shell.execute_reply":"2022-12-06T17:53:54.273758Z"},"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):\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, 360, 128), 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'][:, :4096] * 1e22  # Fourier coefficient complex64\n\n                #p = a.real**2 + a.imag**2  # power\n                #p /= np.mean(p)  # normalize\n                p = normalize(a)\n                p = np.mean(p.reshape(360, 128, 32), axis=2)  # compress 4096 -> 128\n\n                img[ch] = p\n                \n        if CFG.transform:\n            img = transform(img)\n\n        return img, y","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.27623Z","iopub.execute_input":"2022-12-06T17:53:54.27692Z","iopub.status.idle":"2022-12-06T17:53:54.293502Z","shell.execute_reply.started":"2022-12-06T17:53:54.276885Z","shell.execute_reply":"2022-12-06T17:53:54.292496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = Dataset('train', train_labels)\nimg, y = dataset[10]\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 300:360]) # zooming in for dataset[10]\nplt.colorbar()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:54.295071Z","iopub.execute_input":"2022-12-06T17:53:54.295702Z","iopub.status.idle":"2022-12-06T17:53:55.282484Z","shell.execute_reply.started":"2022-12-06T17:53:54.295669Z","shell.execute_reply":"2022-12-06T17:53:55.281481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataset)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:55.283521Z","iopub.execute_input":"2022-12-06T17:53:55.283837Z","iopub.status.idle":"2022-12-06T17:53:55.291191Z","shell.execute_reply.started":"2022-12-06T17:53:55.283804Z","shell.execute_reply":"2022-12-06T17:53:55.290074Z"},"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().__init__()\n\n        # Use timm\n        model = timm.create_model(name, pretrained=pretrained, in_chans=2)\n\n        clsf = model.default_cfg['classifier']\n        n_features = model._modules[clsf].in_features\n        model._modules[clsf] = nn.Identity()\n\n        self.fc = nn.Linear(n_features, 1)\n        self.model = model\n\n    def forward(self, x):\n        x = self.model(x)\n        x = self.fc(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:55.292773Z","iopub.execute_input":"2022-12-06T17:53:55.293359Z","iopub.status.idle":"2022-12-06T17:53:55.303348Z","shell.execute_reply.started":"2022-12-06T17:53:55.293324Z","shell.execute_reply":"2022-12-06T17:53:55.30237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"        \ndef train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    model.train() # switch to training mode\n    nbatch = len(train_loader)\n    running_loss = 0\n    count = 0\n    tb = time.time()\n    \n    pbar = tqdm(train_loader, total=len(train_loader))\n    pbar.set_description(f\"[{epoch+1}/{CFG.epochs}] Train\")\n        \n    for ibatch, (images, labels) in enumerate(pbar):\n        images = images.to(device)\n        labels = labels.to(device)\n        y_preds = model(images)\n        loss = criterion(y_preds.view(-1), labels)\n        running_loss += loss.item()*labels.shape[0]\n        count += labels.shape[0]\n        \n        loss.backward()\n        grad_norm = nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        optimizer.step()\n        scheduler.step(epoch * nbatch + ibatch + 1)\n        optimizer.zero_grad()\n            \n    lr_now = optimizer.param_groups[0]['lr']\n    dt = (time.time()-tb)/60\n    train_dict = {'loss': running_loss/count,\n                     'lr_now': lr_now,\n                     'time': dt}\n        \n    return train_dict\n\n\ndef valid_fn(valid_loader, model, criterion, device, compute_score=True):\n    \n    tb = time.time() \n    model.eval() # switch to evaluation mode\n    preds = []\n    y_all = []\n    running_loss = 0\n    count = 0\n    \n    pbar = tqdm(valid_loader, total=len(valid_loader))\n    pbar.set_description(\"Validation\")\n    \n    for images, labels in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images)\n        loss = criterion(y_preds.view(-1), labels)\n        running_loss += loss.item()*labels.shape[0]\n        count += labels.shape[0]\n        # record accuracy\n        y_all.append(labels.cpu().detach().numpy())\n        preds.append(y_preds.sigmoid().to('cpu').numpy())\n    \n    del loss, images, labels, y_preds\n    \n    y_ground = np.concatenate(y_all)\n    y_pred = np.concatenate(preds)\n    score = get_score(y_ground, y_pred) if compute_score else None \n    val_loss = running_loss/count\n    \n    val_dict = {'loss': val_loss,\n               'score': score,\n               'y': y,\n               'y_pred': y_pred,\n                'time': (time.time() - tb)/60\n               }\n    \n    return val_dict","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:55.304978Z","iopub.execute_input":"2022-12-06T17:53:55.305358Z","iopub.status.idle":"2022-12-06T17:53:55.322309Z","shell.execute_reply.started":"2022-12-06T17:53:55.305322Z","shell.execute_reply":"2022-12-06T17:53:55.321471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train loop\ndef train_loop(data):\n    \n    if CFG.is_cross_validate:\n        \n        kfold = StratifiedKFold(n_splits=CFG.nfold, random_state=42, shuffle=True)\n        for ifold, (idx_train, idx_test) in enumerate(kfold.split(data['id'], data['target'])):\n\n            print('Fold %d/%d' %(ifold, CFG.nfold))\n            torch.manual_seed(CFG.seed + ifold + 1)\n            # create dataset\n            train_dataset = Dataset('train', data.iloc[idx_train])\n            valid_dataset = Dataset('train', data.iloc[idx_test])\n\n            # create dataloader\n            train_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True, \n                                      num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n            valid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size, shuffle=False, \n                                      num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n            # create model and transfer to device\n\n            model = Model(CFG.model_name, pretrained=True)\n            model.to(device)\n\n            # select optimizer, scheduler and criterion\n            optimizer = get_optimizer(model)\n            nbatch = len(train_loader)\n            warmup = CFG.epochs_warmup*nbatch\n            nsteps = CFG.epochs*nbatch\n            scheduler = get_scheduler(optimizer, warmup, nsteps)\n            criterion = get_criterion()\n\n            time_val = 0.0\n            best_score = 0.0 \n            tb = time.time()\n            # start training\n            for epoch in range(CFG.epochs):\n                # train\n                train_dict = train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device)\n                # validation\n                val_dict = valid_fn(valid_loader, model, criterion, device)\n\n                time_val += val_dict['time']\n\n                print('Epoch = %d train_loss = %.4f val_loss = %.4f val_score = %.4f lr = %.2e time = %.2f min' % (epoch+1, train_dict['loss'], val_dict['loss'], \n                                                                 val_dict['score'], train_dict['lr_now'], train_dict['time']))\n                \n                val_score = val_dict['score']\n                if val_score > best_score:\n                    best_score = val_score\n                    output_file = 'model_%d_best.pytorch'%ifold\n                    torch.save(model.state_dict(), output_file)\n                    print(output_file, 'written')\n                    \n            torch.save(model.state_dict(), 'model_%d.pytorch'%ifold)\n            print('model_%d.pytorch'%ifold, 'written')\n            dt = (time.time() - tb)/60\n            print('Training done %.2f min total, %.2f min val'% (dt, time_val))\n","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:55.325815Z","iopub.execute_input":"2022-12-06T17:53:55.326289Z","iopub.status.idle":"2022-12-06T17:53:55.339932Z","shell.execute_reply.started":"2022-12-06T17:53:55.326262Z","shell.execute_reply":"2022-12-06T17:53:55.339313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# main\ndef main():\n    if CFG.train: \n        # train\n        train_loop(train_labels)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:55.34135Z","iopub.execute_input":"2022-12-06T17:53:55.341926Z","iopub.status.idle":"2022-12-06T17:53:55.354673Z","shell.execute_reply.started":"2022-12-06T17:53:55.34189Z","shell.execute_reply":"2022-12-06T17:53:55.353604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    main()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T17:53:55.358093Z","iopub.execute_input":"2022-12-06T17:53:55.358372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train the model on whole data and save it\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}