{"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 /kaggle/input/audiomentations-v0290/resampy-0.4.2-py3-none-any.whl\n!pip install  /kaggle/input/audiomentations-v0290/librosa-0.9.2-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:30.720243Z","iopub.execute_input":"2023-06-19T15:55:30.720844Z","iopub.status.idle":"2023-06-19T15:55:52.871554Z","shell.execute_reply.started":"2023-06-19T15:55:30.720814Z","shell.execute_reply":"2023-06-19T15:55:52.870390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntorch.cuda.is_available()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:52.874661Z","iopub.execute_input":"2023-06-19T15:55:52.875430Z","iopub.status.idle":"2023-06-19T15:55:52.883176Z","shell.execute_reply.started":"2023-06-19T15:55:52.875360Z","shell.execute_reply":"2023-06-19T15:55:52.882145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\nimport random\n\nimport numpy as np\nimport pandas as pd\nimport time\nimport os\nimport matplotlib.pyplot as plt\n# These transformations will be passed to our model class\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport yaml\nfrom tqdm.auto import tqdm\nimport glob\nfrom torch.distributions import Beta\nfrom torchvision.ops import sigmoid_focal_loss","metadata":{"papermill":{"duration":6.939158,"end_time":"2021-12-13T10:09:03.241098","exception":false,"start_time":"2021-12-13T10:08:56.301940","status":"completed"},"scrolled":true,"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:52.884657Z","iopub.execute_input":"2023-06-19T15:55:52.885278Z","iopub.status.idle":"2023-06-19T15:55:52.895401Z","shell.execute_reply.started":"2023-06-19T15:55:52.885247Z","shell.execute_reply":"2023-06-19T15:55:52.894379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_DIR = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction'\nINPUT_DIR_NP = '/kaggle/input/data-creation-v1'\nPRETRAIN_CKPT='/kaggle/input/pre-train-weight-dir/wavenet_4096-fold0_14.pth'","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:52.896788Z","iopub.execute_input":"2023-06-19T15:55:52.897308Z","iopub.status.idle":"2023-06-19T15:55:52.907357Z","shell.execute_reply.started":"2023-06-19T15:55:52.897277Z","shell.execute_reply":"2023-06-19T15:55:52.906502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from joblib import Parallel, delayed\nimport os\nfrom os.path import exists\n\nWAV_SIZE=2000\nSTEP_SIZE=300\nTIMES_REAL=4\nTIMES_TRAIN=8\nis_mixed_precision = True\n\n\n\nclass GaitDataset(torch.utils.data.Dataset):\n\n    def __init__(self, df, is_train=False,transforms=None):\n        self.is_train = is_train\n        self.data = df\n\n    def __len__(self):\n        if self.is_train:\n            return len(self.data)*TIMES_TRAIN\n        else:\n            return len(self.data)\n    \n    def get_defog_data(self,row):\n        wid = row.Id\n        files = glob.glob(f'{INPUT_DIR_NP}/defog_np/{wid}*.npy')\n        \n        if self.is_train:\n            rix = np.random.randint(0,len(files))\n        \n            file = files[rix]\n            data = np.load(file,allow_pickle=True).astype(np.float32)\n            \n            val_data = data\n            mask = data[:,-1]\n            \n            vix_s = np.where(mask==1)[0][0]\n            vix_e = np.where(mask[vix_s:]==0)[0][0]\n            \n            if vix_e < WAV_SIZE:\n                six = np.random.randint(0,len(data)-WAV_SIZE)\n            else:\n                six = np.random.randint(0,vix_s+vix_e-WAV_SIZE-1)\n            \n            wav = data[six:six+WAV_SIZE,0:3]\n            label = data[six:six+WAV_SIZE,3:6]\n            mask = data[six:six+WAV_SIZE,-1]\n\n            wav = wav.reshape(1,wav.shape[0],-1)\n            label= label.reshape(1,label.shape[0],-1)\n            mask= mask.reshape(1,-1)\n            #print('wav',wav.shape,mask.shape)\n        \n        return wav,label,mask\n    \n    def __getitem__(self, idx):\n        if self.is_train:\n            idx = np.random.randint(0,len(self.data))\n            \n        row = self.data.iloc[idx]\n        wid = row.Id\n        \n        subject = row.Subject\n        t = row.type\n        \n        if t == 0:\n            wav = np.load(f'{INPUT_DIR_NP}/tdcsfog_np/{row.Id}_sig.npy')\n            tgt = np.load(f'{INPUT_DIR_NP}/tdcsfog_np/{row.Id}_tgt.npy')\n            mask = np.ones((wav.shape[0]))\n        \n\n        if self.is_train:\n            if row.type == 1 and np.random.uniform() < 0:\n                wav,label,mask = self.get_defog_data(row)\n                actual_len = -1\n                \n                #print('wav',wav.shape,label.shape)\n            else:\n                if row.type == 1:\n                    wav = np.load(f'{INPUT_DIR_NP}/defog_np//{wid}_sig.npy')\n                    tgt = np.load(f'{INPUT_DIR_NP}/defog_np/{wid}_tgt.npy')\n                    mask = pd.read_csv(f'{INPUT_DIR}/train/defog/{wid}.csv').Valid.values\n                    \n                label = tgt\n                \n                wav_df = pd.DataFrame(wav)\n                tgt_df = pd.DataFrame(label)\n\n                wavs = []\n                tgts = []\n                for w in wav_df.rolling(WAV_SIZE,step=STEP_SIZE):\n                    if w.shape[0] == WAV_SIZE:\n                        wavs.append(w.values)\n\n                if len(wavs) ==0:\n                    wavs = [wav]\n\n                for w in tgt_df.rolling(WAV_SIZE,step=STEP_SIZE):\n                    if w.shape[0] == WAV_SIZE:\n                        tgts.append(w.values)\n\n                if len(tgts) ==0:\n                    tgts = [label]\n\n                wav = np.stack(wavs,axis=0)\n                label = np.stack(tgts,axis=0)\n                actual_len=-1\n        \n        else:\n            if row.type == 1:\n                wav = np.load(f'{INPUT_DIR_NP}/defog_np//{wid}_sig.npy')\n                tgt = np.load(f'{INPUT_DIR_NP}/defog_np/{wid}_tgt.npy')\n                mask = pd.read_csv(f'{INPUT_DIR}/train/defog/{wid}.csv').Valid.values\n                    \n            label = tgt\n            actual_len = len(wav)\n            nchunk = (len(wav)//WAV_SIZE)+1\n            wav = wav.reshape(-1,len(wav),3)\n            label = label.reshape(-1,len(label),3)\n            \n        \n        if self.is_train and len(wav)>1:\n            if row.type == 0:\n                rix = np.random.randint(0,len(wav))\n                wav = wav[rix:rix+1]\n                label = label[rix:rix+1]\n            else:\n                rix = np.random.randint(0,len(wav),TIMES_REAL)\n                wav = wav[rix]\n                label = label[rix]\n        \n        #print('wav',wav.shape, label.shape)\n        \n        wav = wav/40.\n        \n        sample = {\"wav\": wav, \"label\":label, \"actual_len\":actual_len, \"mask\":mask}\n        \n        #print('label',label.shape,tgt.shape)\n        #print('wav',wav.shape)\n\n        return sample\n        \ndef collate_wrapper(batch):\n    out = {}\n    wavs = []\n    labels = []\n    s_ix1s = []\n    e_ix1s = []\n    for item in batch:\n        wavs.append(item['wav'])\n        labels.append(item['label'])\n        \n    out['wav'] = torch.from_numpy(np.concatenate(wavs,axis=0))\n    out['label'] = torch.from_numpy(np.concatenate(labels,axis=0))\n    \n    return out\n\ndef getDataLoader(params,train_x,val_x,train_transforms=None,val_transforms=None):\n    \n    train_dataset = GaitDataset(\n            df=train_x, is_train=True, transforms=train_transforms\n        )\n    val_dataset = GaitDataset(df=val_x, transforms=val_transforms)\n    \n    trainDataLoader = torch.utils.data.DataLoader(\n                            train_dataset,\n                            batch_size=params['batch_size'],\n                            num_workers=params['num_workers'],\n                            shuffle=True,collate_fn = collate_wrapper,\n                            pin_memory=False,\n                            worker_init_fn=lambda id: np.random.seed(torch.initial_seed() // 2 ** 32 + id)\n                        )\n    valDataLoader = torch.utils.data.DataLoader(\n                        val_dataset,\n                        batch_size=1,\n                        num_workers=params['num_workers'],\n                        shuffle=False,\n                        pin_memory=False,\n                    )\n    \n    return trainDataLoader,valDataLoader","metadata":{"papermill":{"duration":0.04712,"end_time":"2021-12-13T10:09:03.372538","exception":false,"start_time":"2021-12-13T10:09:03.325418","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:52.910908Z","iopub.execute_input":"2023-06-19T15:55:52.911299Z","iopub.status.idle":"2023-06-19T15:55:52.943784Z","shell.execute_reply.started":"2023-06-19T15:55:52.911276Z","shell.execute_reply":"2023-06-19T15:55:52.942945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Mixup(nn.Module):\n    def __init__(self, mix_beta=1):\n\n        super(Mixup, self).__init__()\n        self.beta_distribution = Beta(mix_beta, mix_beta)\n\n    def forward(self, X, Y, weight=None):\n\n        bs = X.shape[0]\n        n_dims = len(X.shape)\n        perm = torch.randperm(bs)\n        coeffs = self.beta_distribution.rsample(torch.Size((bs,))).to(X.device)\n\n        if n_dims == 2:\n            X = coeffs.view(-1, 1) * X + (1 - coeffs.view(-1, 1)) * X[perm]\n        elif n_dims == 3:\n            X = coeffs.view(-1, 1, 1) * X + (1 - coeffs.view(-1, 1, 1)) * X[perm]\n        else:\n            X = coeffs.view(-1, 1, 1, 1) * X + (1 - coeffs.view(-1, 1, 1, 1)) * X[perm]\n\n        Y = coeffs.view(-1, 1,1) * Y + (1 - coeffs.view(-1, 1, 1)) * Y[perm]\n\n        if weight is None:\n            return X, Y\n        else:\n            weight = coeffs.view(-1) * weight + (1 - coeffs.view(-1)) * weight[perm]\n            return X, Y, weight","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:52.945076Z","iopub.execute_input":"2023-06-19T15:55:52.945430Z","iopub.status.idle":"2023-06-19T15:55:52.958187Z","shell.execute_reply.started":"2023-06-19T15:55:52.945398Z","shell.execute_reply":"2023-06-19T15:55:52.957254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Wave_Block(nn.Module):\n\n    def __init__(self, in_channels, out_channels, dilation_rates, kernel_size):\n        super(Wave_Block, self).__init__()\n        self.num_rates = dilation_rates\n        self.convs = nn.ModuleList()\n        self.filter_convs = nn.ModuleList()\n        self.gate_convs = nn.ModuleList()\n\n        self.convs.append(nn.Conv1d(in_channels, out_channels, kernel_size=1))\n        dilation_rates = [2 ** i for i in range(dilation_rates)]\n        for dilation_rate in dilation_rates:\n            self.filter_convs.append(\n                nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size, padding=int((dilation_rate*(kernel_size-1))/2), dilation=dilation_rate))\n            self.gate_convs.append(\n                nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size, padding=int((dilation_rate*(kernel_size-1))/2), dilation=dilation_rate))\n            self.convs.append(nn.Conv1d(out_channels, out_channels, kernel_size=1))\n\n    def forward(self, x):\n        x = self.convs[0](x)\n        res = x\n        for i in range(self.num_rates):\n            x = torch.tanh(self.filter_convs[i](x)) * torch.sigmoid(self.gate_convs[i](x))\n            x = self.convs[i + 1](x)\n            res = res + x\n        return res\n# detail \nclass Classifier(nn.Module):\n    def __init__(self, inch=3, kernel_size=3):\n        super().__init__()\n        self.LSTM1 = nn.GRU(input_size=128, hidden_size=128, num_layers=4,dropout=0.2, \n                           batch_first=True, bidirectional=True)\n        \n        #self.wave_block1 = Wave_Block(inch, 16, 12, kernel_size)\n        self.wave_block2 = Wave_Block(inch, 32, 8, kernel_size)\n        self.wave_block3 = Wave_Block(32, 64, 4, kernel_size)\n        self.wave_block4 = Wave_Block(64, 128, 1, kernel_size)\n        self.fc1 = nn.Linear(256, 3)\n\n    def forward(self, x):\n        x = x.permute(0, 2, 1)\n        #x = self.wave_block1(x)\n        x = self.wave_block2(x)\n        x = self.wave_block3(x)\n\n        x = self.wave_block4(x)\n        x = x.permute(0, 2, 1)\n        x, h = self.LSTM1(x)\n        x = self.fc1(x)\n    \n        \n        return x,x,x","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:52.959682Z","iopub.execute_input":"2023-06-19T15:55:52.960021Z","iopub.status.idle":"2023-06-19T15:55:52.975970Z","shell.execute_reply.started":"2023-06-19T15:55:52.959991Z","shell.execute_reply":"2023-06-19T15:55:52.975009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.cuda.amp as amp\nclass AmpNet(Classifier):\n    \n    def __init__(self,params):\n        super(AmpNet, self).__init__()\n    @torch.cuda.amp.autocast()\n    def forward(self,*args):\n        return super(AmpNet, self).forward(*args)\n\nis_mixed_precision = True  #True #False","metadata":{"papermill":{"duration":0.080072,"end_time":"2021-12-13T10:09:03.957995","exception":false,"start_time":"2021-12-13T10:09:03.877923","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:52.977394Z","iopub.execute_input":"2023-06-19T15:55:52.977726Z","iopub.status.idle":"2023-06-19T15:55:52.989497Z","shell.execute_reply.started":"2023-06-19T15:55:52.977696Z","shell.execute_reply":"2023-06-19T15:55:52.988424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getOptimzersScheduler(model,params,steps_in_epoch=25,pct_start=0.1):\n    \n    \n    mdl_parameters = [\n                {'params': model.wave_block2.parameters(), 'lr': 1e-4},\n                {'params': model.wave_block3.parameters(), 'lr': 1e-4},\n                {'params': model.wave_block4.parameters(), 'lr': 1e-4},\n                {'params': model.LSTM1.parameters(), 'lr': 1e-4},\n                {'params': model.fc1.parameters(), 'lr': 1e-4},\n                #{'params': model.fc.parameters(), 'lr': 1e-4},\n                #{'params': model.attention.parameters(), 'lr': 1e-5},\n            ]\n    \n    optimizer = torch.optim.Adam(mdl_parameters, lr=params['learning_rate'][0])\n    \n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer,steps_per_epoch=1,\n                                                    pct_start=pct_start,\n                                                    max_lr=params['learning_rate'],\n                                                    epochs  = params['max_epochs'], \n                                                    div_factor = params['div_factor'], \n                                                    final_div_factor=params['final_div_factor'],\n                                                    verbose=False)\n    \n    return optimizer,scheduler,False","metadata":{"papermill":{"duration":0.032959,"end_time":"2021-12-13T10:09:04.015870","exception":false,"start_time":"2021-12-13T10:09:03.982911","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:52.990758Z","iopub.execute_input":"2023-06-19T15:55:52.991171Z","iopub.status.idle":"2023-06-19T15:55:53.002063Z","shell.execute_reply.started":"2023-06-19T15:55:52.991142Z","shell.execute_reply":"2023-06-19T15:55:53.001010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_model(epoch,model,ckpt_path='./',name='',score=0):\n    path = os.path.join(ckpt_path, '{}_{}_{:.2f}.pth'.format(name, epoch,score))\n    torch.save(model.state_dict(), path, _use_new_zipfile_serialization=False)\n    \ndef load_model(model,ckpt_path):\n    state = torch.load(ckpt_path)\n    print(model.load_state_dict(state,strict=False))\n    return model","metadata":{"papermill":{"duration":0.032012,"end_time":"2021-12-13T10:09:04.072285","exception":false,"start_time":"2021-12-13T10:09:04.040273","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.003606Z","iopub.execute_input":"2023-06-19T15:55:53.003975Z","iopub.status.idle":"2023-06-19T15:55:53.016333Z","shell.execute_reply.started":"2023-06-19T15:55:53.003942Z","shell.execute_reply":"2023-06-19T15:55:53.015395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def focal_loss(pred,target):\n    return 32*sigmoid_focal_loss(pred,target,reduction='mean')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:53.017821Z","iopub.execute_input":"2023-06-19T15:55:53.018246Z","iopub.status.idle":"2023-06-19T15:55:53.026526Z","shell.execute_reply.started":"2023-06-19T15:55:53.018203Z","shell.execute_reply":"2023-06-19T15:55:53.025428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def training_step(model, batch, batch_idx,optimizer,scheduler,isStepScheduler=False):\n    # Load images and labels\n    x = batch[\"wav\"].float()\n    y = batch[\"label\"].float()\n    \n   \n    mixup = Mixup()\n    \n    ##Mixup Aug\n    if np.random.uniform(0,1) < 0.:\n        x,y = mixup(x,y)\n    \n    #print('x',x.shape,ys1.shape,ye1.shape)\n    \n    if GPU:\n        x, y  = x.cuda(non_blocking=True), y.cuda(non_blocking=True)\n\n    criterion = focal_loss #torch.nn.BCEWithLogitsLoss(reduction=\"mean\") \n    reg_criterion = torch.nn.MSELoss(reduction='mean')\n    #criterion = FocalLoss()\n    \n    #optimizer.zero_grad()\n    iters_to_accumulate=2\n    # Forward \n\n    if is_mixed_precision:\n        with amp.autocast():\n            preds, sp, ep = model(x)\n            b,s,c = y.shape\n            y = y.reshape(b*s,c)\n            preds = preds.reshape(b*s,-1)\n            loss = criterion(preds,y)/ iters_to_accumulate\n            \n            #print('pred',sp.shape,ys1.shape, ep.shape)\n            #rloss1 = reg_criterion(sp,ys1)\n            #rloss2 = reg_criterion(ep,ye1)\n            #loss = loss + rloss1 + rloss2\n            \n            scaler.scale(loss).backward()\n            \n            if (batch_idx + 1) % iters_to_accumulate == 0:\n                #print('accumulating')\n            # may unscale_ here if desired (e.g., to allow clipping unscaled gradients)\n                scaler.unscale_(optimizer)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n            \n            loss = loss.item()\n    else:\n        preds = model(x,att_mask)\n        loss = criterion(preds.flatten(), y.flatten())\n        loss.backward()\n        optimizer.step()\n        loss = loss.item()\n        \n    if isStepScheduler:\n        scheduler.step()\n\n    # Calculate validation IOU (global)\n    #preds = preds.detach()\n    #y = y.detach().cpu()\n    return loss\n\ndef validation_step(model, batch, batch_idx):\n    # Load images and labels\n    x = batch[\"wav\"].float()\n    y = batch[\"label\"].float()\n    mask = batch['mask'].long().unsqueeze(-1)\n    \n    actual_len = batch['actual_len'].long()\n    iters_to_accumulate= 2\n    \n    if GPU:\n        x, y = x.cuda(non_blocking=True), y.cuda(non_blocking=True)\n        mask = mask.cuda()\n\n    criterion = focal_loss #torch.nn.BCEWithLogitsLoss(reduction=\"mean\") \n    #criterion = FocalLoss()\n\n    x = x[0]\n    y =y[0]\n    actual_len=actual_len[0]\n    mask = mask[0]\n    \n    BS=20\n    \n    preds_list = []\n    tgt_list = []\n\n    # Forward pass & softmax\n    with torch.no_grad():\n        if is_mixed_precision:\n            with amp.autocast():\n                num_iter =  x.shape[0]//BS\n                if num_iter== 0:\n                    num_iter=1\n                for b in range(num_iter):\n                    preds,_,_ = model(x[b*BS:(b+1)*BS])\n                    yb = y[b*BS:(b+1)*BS]\n                    b,s,c = yb.shape\n                    yb = yb.reshape(b*s,c)\n                    preds = preds.reshape(b*s,-1)\n                    preds_list.append(preds)\n                    tgt_list.append(yb)\n                   \n                preds = torch.cat(preds_list,dim=0)\n                y = torch.cat(tgt_list,dim=0)\n                \n                #print('preds',preds.shape,y.shape)\n                \n                y = y[0:actual_len]\n                preds = preds[0:actual_len]\n                loss = criterion(preds, y)\n                \n                loss = (loss*mask).sum()/mask.sum()\n                \n    preds = torch.sigmoid(preds)\n    \n    preds = preds.cpu().numpy()\n    y= y.long().cpu().numpy()\n    mask = mask.long().cpu().numpy()[:,0]\n    \n    preds = preds[mask==1]\n    y = y[mask==1]\n    \n    \n    loss = loss.item()\n    return loss,preds,y\n","metadata":{"papermill":{"duration":0.040406,"end_time":"2021-12-13T10:09:04.193173","exception":false,"start_time":"2021-12-13T10:09:04.152767","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.029522Z","iopub.execute_input":"2023-06-19T15:55:53.029800Z","iopub.status.idle":"2023-06-19T15:55:53.050238Z","shell.execute_reply.started":"2023-06-19T15:55:53.029777Z","shell.execute_reply":"2023-06-19T15:55:53.049217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score,f1_score,precision_score,average_precision_score\ndef train_epoch(model,trainDataLoader,optimizer,scheduler,isStepScheduler=True):\n    total_intersection=0\n    total_union=0\n    total_loss=0\n    model.train()\n    torch.set_grad_enabled(True)\n    total_step=0\n    ious = []\n    \n    \n    pbar = tqdm(enumerate(trainDataLoader),total=len(trainDataLoader))\n    for bi,data in pbar:\n        loss= training_step(model,data,bi,optimizer,scheduler)\n        total_loss+=loss\n        total_step+=1\n        pbar.set_postfix({'loss':total_loss/total_step})\n        \n    if not isStepScheduler: #in case epoch based scheduler\n        scheduler.step()\n            \n    total_loss /= total_step\n    return total_loss\n        \n\ndef val_epoch(model,valDataLoader):\n    total_intersection=0\n    total_union=0\n    total_loss=0\n    \n    total_step=0\n    model.eval()\n    preds = []\n    targets = []\n    #pbar=tqdm(enumerate(valDataLoader),total=len(valDataLoader))\n    for bi,data in enumerate(valDataLoader) :\n        loss, pred ,tgt = validation_step(model,data,bi)\n        total_loss+=loss\n        total_step+=1\n        preds.extend(pred)\n        targets.extend(tgt)\n        \n        #pbar.set_postfix({'loss':total_loss/total_step})\n        \n    preds = np.stack(preds)\n    preds = np.clip(preds,0,1)\n    targets = np.stack(targets)\n    \n    #preds = preds[targets!=0]\n    #targets = targets[targets!=0]\n    \n    print('targets',targets.shape, preds.shape)\n    aps = []\n    for i in range(3):\n        score = average_precision_score(targets[:,i],preds[:,i])\n        aps.append(score) \n    \n    APx = average_precision_score(targets,preds,average='macro')\n    AP = np.mean(aps)\n    \n    del targets,preds\n    gc.collect()\n    \n    print('AP', AP, APx)\n    total_loss /= total_step\n    return total_loss,AP","metadata":{"papermill":{"duration":0.036917,"end_time":"2021-12-13T10:09:04.254423","exception":false,"start_time":"2021-12-13T10:09:04.217506","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.051917Z","iopub.execute_input":"2023-06-19T15:55:53.052261Z","iopub.status.idle":"2023-06-19T15:55:53.066880Z","shell.execute_reply.started":"2023-06-19T15:55:53.052212Z","shell.execute_reply":"2023-06-19T15:55:53.065854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GPU=True\ndef training_loop(params,train_x,val_x,savedir='./',mdl_name='resnet34'):\n    \n    #create model\n    model = AmpNet(params).cuda()\n    #load model\n    model = load_model(model,PRETRAIN_CKPT)\n    #get loaders\n    train_transforms=None\n    val_transforms = None\n    trainDataLoader,valDataLoader = getDataLoader(params,train_x,val_x,train_transforms,val_transforms)\n    \n    optimizer,scheduler,isStepScheduler = getOptimzersScheduler(model,params,\n                                                                steps_in_epoch=len(trainDataLoader),\n                                                                pct_start=0.1)\n    best_ap= 0\n    #control loop\n    for e in range(params['max_epochs']):\n        train_loss = train_epoch(model,trainDataLoader,optimizer,scheduler,isStepScheduler)\n        loss, AP = val_epoch(model,valDataLoader)\n        #logging here\n        #print(e,'Train Result',f'loss={train_loss}')\n        print(e,'Val Result',f'AP={AP} ')\n        if AP > 0.25 :\n            print(f'Saving for AP {AP}')\n            save_model(e,model,ckpt_path=savedir,name=mdl_name,score=best_ap)\n            best_ap=AP\n        else:\n            print(f'Not Saving for AP {AP}')\n        ","metadata":{"papermill":{"duration":0.033772,"end_time":"2021-12-13T10:09:04.368672","exception":false,"start_time":"2021-12-13T10:09:04.334900","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.071114Z","iopub.execute_input":"2023-06-19T15:55:53.071386Z","iopub.status.idle":"2023-06-19T15:55:53.081716Z","shell.execute_reply.started":"2023-06-19T15:55:53.071347Z","shell.execute_reply":"2023-06-19T15:55:53.080829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hparams = {\n    # Optional hparams\n    \"backbone\": 'wavenet_2000_pretrain', #'', #'tf_efficientnetv2_b2',\n    #\"learning_rate\": [1e-4,3e-4,5e-4,5e-4,5e-4],\n    \"learning_rate\": [1e-4,1e-4,1e-4,5e-4,1e-3],\n    \"max_epochs\": 91,\n    \"batch_size\": 16,\n    \"num_workers\": 0,\n    \"val_sanity_checks\": 0,\n    \"fast_dev_run\": False,\n    \"output_path\": f\"\",\n    \"gpu\": torch.cuda.is_available(),\n    'div_factor':10,\n    'final_div_factor':50,\n}","metadata":{"papermill":{"duration":0.031307,"end_time":"2021-12-13T10:09:04.424259","exception":false,"start_time":"2021-12-13T10:09:04.392952","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.083924Z","iopub.execute_input":"2023-06-19T15:55:53.084454Z","iopub.status.idle":"2023-06-19T15:55:53.095580Z","shell.execute_reply.started":"2023-06-19T15:55:53.084418Z","shell.execute_reply":"2023-06-19T15:55:53.094681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scaler = amp.GradScaler()","metadata":{"papermill":{"duration":0.034696,"end_time":"2021-12-13T10:09:04.580936","exception":false,"start_time":"2021-12-13T10:09:04.546240","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.097040Z","iopub.execute_input":"2023-06-19T15:55:53.097355Z","iopub.status.idle":"2023-06-19T15:55:53.105782Z","shell.execute_reply.started":"2023-06-19T15:55:53.097326Z","shell.execute_reply":"2023-06-19T15:55:53.104910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nimport random\nseed=42\ndef set_seed(seed=42):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.use_deterministic_algorithms = True\n    random.seed(0)\n    np.random.seed(0)\nset_seed(seed)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:53.107125Z","iopub.execute_input":"2023-06-19T15:55:53.107570Z","iopub.status.idle":"2023-06-19T15:55:53.115863Z","shell.execute_reply.started":"2023-06-19T15:55:53.107539Z","shell.execute_reply":"2023-06-19T15:55:53.114949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nevents = pd.read_csv(f'{INPUT_DIR}/events.csv')\nevents = events[~events.Type.isnull()]\n\ndefog = pd.read_csv(f'{INPUT_DIR}/defog_metadata.csv')\ndefog = defog[defog.Id.isin(events.Id)].reset_index(drop=True)\n\ntdcsfog = pd.read_csv(f'{INPUT_DIR}/tdcsfog_metadata.csv')\ntdcsfog = tdcsfog[tdcsfog.Id.isin(events.Id)].reset_index(drop=True)\n\ndefog['type'] = 1\ntdcsfog['type'] = 0\n\ntrain = pd.concat([tdcsfog]).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:53.117397Z","iopub.execute_input":"2023-06-19T15:55:53.117737Z","iopub.status.idle":"2023-06-19T15:55:53.146277Z","shell.execute_reply.started":"2023-06-19T15:55:53.117708Z","shell.execute_reply":"2023-06-19T15:55:53.145501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set(tdcsfog.Subject.unique()).intersection(set(defog.Subject.unique()))","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:53.147547Z","iopub.execute_input":"2023-06-19T15:55:53.148229Z","iopub.status.idle":"2023-06-19T15:55:53.155299Z","shell.execute_reply.started":"2023-06-19T15:55:53.148199Z","shell.execute_reply":"2023-06-19T15:55:53.154440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold, StratifiedKFold, StratifiedGroupKFold,GroupKFold\n#kf = GroupKFold(n_splits=5,random_state=42,shuffle=True)\nkf = GroupKFold(n_splits=5)\nfor i, (train_index, test_index) in enumerate(kf.split(tdcsfog.Id,tdcsfog.Medication,groups=tdcsfog.Subject)):\n    tdcsfog.loc[test_index,'fold'] =i\n    \n#kf = GroupKFold(n_splits=5,random_state=42,shuffle=True)\nkf = GroupKFold(n_splits=5)\nfor i, (train_index, test_index) in enumerate(kf.split(defog.Id,defog.Medication,groups=defog.Subject)):\n    defog.loc[test_index,'fold'] =i","metadata":{"execution":{"iopub.status.busy":"2023-06-19T15:55:53.156546Z","iopub.execute_input":"2023-06-19T15:55:53.157403Z","iopub.status.idle":"2023-06-19T15:55:53.175357Z","shell.execute_reply.started":"2023-06-19T15:55:53.157341Z","shell.execute_reply":"2023-06-19T15:55:53.174479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nversion='1'\nfn=0\nfor fn in [1,4,0,2,3]:  \n    set_seed()\n\n    mdl_name=hparams['backbone']\n    savedir = f'trained-models-{mdl_name}-v{version}'\n    Path(savedir).mkdir(exist_ok=True, parents=True)\n    \n    val = pd.concat([defog[defog.fold==fn],tdcsfog[tdcsfog.fold==fn]])\n    tr = pd.concat([defog[defog.fold!=fn],tdcsfog[tdcsfog.fold!=fn]])\n    \n    print('Train',tr.shape,'Val',val.shape)\n    \n   \n    training_loop(hparams,tr,val,savedir=savedir,mdl_name=f'{mdl_name}-fold{fn}')\n    gc.collect()\n\n    #break","metadata":{"papermill":{"duration":0.04818,"end_time":"2021-12-13T13:17:52.100451","exception":false,"start_time":"2021-12-13T13:17:52.052271","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-19T15:55:53.177064Z","iopub.execute_input":"2023-06-19T15:55:53.177519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}