{"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 torch==1.12.1+cu113 torchvision==0.13.1+cu113 torchaudio==0.12.1 --extra-index-url https://download.pytorch.org/whl/cu113\n\n!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-19T04:34:43.923139Z","iopub.execute_input":"2023-06-19T04:34:43.923400Z","iopub.status.idle":"2023-06-19T04:35:09.078791Z","shell.execute_reply.started":"2023-06-19T04:34:43.923376Z","shell.execute_reply":"2023-06-19T04:35:09.077605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntorch.cuda.is_available()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T04:35:09.081636Z","iopub.execute_input":"2023-06-19T04:35:09.082735Z","iopub.status.idle":"2023-06-19T04:35:12.167189Z","shell.execute_reply.started":"2023-06-19T04:35:09.082692Z","shell.execute_reply":"2023-06-19T04:35:12.166313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.version.cuda,torch.__version__","metadata":{"execution":{"iopub.status.busy":"2023-06-19T04:35:12.168471Z","iopub.execute_input":"2023-06-19T04:35:12.169434Z","iopub.status.idle":"2023-06-19T04:35:12.176855Z","shell.execute_reply.started":"2023-06-19T04:35:12.169397Z","shell.execute_reply":"2023-06-19T04:35:12.175980Z"},"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-19T04:35:12.179654Z","iopub.execute_input":"2023-06-19T04:35:12.180410Z","iopub.status.idle":"2023-06-19T04:35:12.492013Z","shell.execute_reply.started":"2023-06-19T04:35:12.180337Z","shell.execute_reply":"2023-06-19T04:35:12.491120Z"},"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=500\nTIMES_REAL=4\nTIMES_TRAIN=8\nis_mixed_precision = True\nINPUT_PATH_NP = '/kaggle/input/data-creation-v1'\nINPUT_PATH = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction'\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    \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        subject = row.Subject\n        t = row.type\n        \n        if t == 0:\n            wav = np.load(f'{INPUT_PATH_NP}/tdcsfog_np/{row.Id}_sig.npy')\n            tgt = np.load(f'{INPUT_PATH_NP}/tdcsfog_np/{row.Id}_tgt.npy')\n        else:\n            wav = np.load(f'{INPUT_PATH_NP}/defog_np//{row.Id}_sig.npy')\n            tgt = np.load(f'{INPUT_PATH_NP}/defog_np/{row.Id}_tgt.npy')\n            \n        \n        wav = wav/40.\n        \n        label = tgt\n        wav_df = pd.DataFrame(wav)\n        tgt_df = pd.DataFrame(label)\n        \n        wavs = []\n        tgts = []\n        if self.is_train:\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            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        sample = {\"wav\": wav, \"label\":label, \"actual_len\":actual_len}\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-19T04:35:12.493213Z","iopub.execute_input":"2023-06-19T04:35:12.493539Z","iopub.status.idle":"2023-06-19T04:35:12.535564Z","shell.execute_reply.started":"2023-06-19T04:35:12.493507Z","shell.execute_reply":"2023-06-19T04:35:12.534734Z"},"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-19T04:35:12.536685Z","iopub.execute_input":"2023-06-19T04:35:12.537425Z","iopub.status.idle":"2023-06-19T04:35:12.547709Z","shell.execute_reply.started":"2023-06-19T04:35:12.537392Z","shell.execute_reply":"2023-06-19T04:35:12.546779Z"},"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.LSTM = nn.GRU(input_size=128, hidden_size=128, num_layers=4, \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.LSTM(x)\n        x = self.fc1(x)\n    \n        \n        return x,x,x","metadata":{"execution":{"iopub.status.busy":"2023-06-19T04:35:12.549310Z","iopub.execute_input":"2023-06-19T04:35:12.549639Z","iopub.status.idle":"2023-06-19T04:35:12.564212Z","shell.execute_reply.started":"2023-06-19T04:35:12.549610Z","shell.execute_reply":"2023-06-19T04:35:12.563151Z"},"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-19T04:35:12.565710Z","iopub.execute_input":"2023-06-19T04:35:12.566097Z","iopub.status.idle":"2023-06-19T04:35:12.575947Z","shell.execute_reply.started":"2023-06-19T04:35:12.566067Z","shell.execute_reply":"2023-06-19T04:35:12.575118Z"},"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.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=True)\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-19T04:35:12.577041Z","iopub.execute_input":"2023-06-19T04:35:12.577490Z","iopub.status.idle":"2023-06-19T04:35:12.587490Z","shell.execute_reply.started":"2023-06-19T04:35:12.577459Z","shell.execute_reply":"2023-06-19T04:35:12.586690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_model(epoch,model,ckpt_path='./',name='',val_rmse=0):\n    path = os.path.join(ckpt_path, '{}_{}.pth'.format(name, epoch))\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-19T04:35:12.591524Z","iopub.execute_input":"2023-06-19T04:35:12.591806Z","iopub.status.idle":"2023-06-19T04:35:12.598231Z","shell.execute_reply.started":"2023-06-19T04:35:12.591784Z","shell.execute_reply":"2023-06-19T04:35:12.597152Z"},"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-19T04:35:12.599963Z","iopub.execute_input":"2023-06-19T04:35:12.600369Z","iopub.status.idle":"2023-06-19T04:35:12.608084Z","shell.execute_reply.started":"2023-06-19T04:35:12.600270Z","shell.execute_reply":"2023-06-19T04:35:12.607203Z"},"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    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\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    \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)/ iters_to_accumulate\n                \n    preds = torch.sigmoid(preds)\n    \n    loss = loss.item()\n    return loss,preds.detach().cpu().numpy(),y.long().detach().cpu().numpy()\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-19T04:35:12.609577Z","iopub.execute_input":"2023-06-19T04:35:12.609894Z","iopub.status.idle":"2023-06-19T04:35:12.629404Z","shell.execute_reply.started":"2023-06-19T04:35:12.609863Z","shell.execute_reply":"2023-06-19T04:35:12.628721Z"},"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 pbar :\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-19T04:35:12.630582Z","iopub.execute_input":"2023-06-19T04:35:12.631343Z","iopub.status.idle":"2023-06-19T04:35:13.124390Z","shell.execute_reply.started":"2023-06-19T04:35:12.631312Z","shell.execute_reply":"2023-06-19T04:35:13.123454Z"},"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    \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 > best_ap :\n            print(f'Saving for AP {AP}')\n            save_model(e,model,ckpt_path=savedir,name=mdl_name,val_rmse=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-19T04:35:13.126005Z","iopub.execute_input":"2023-06-19T04:35:13.126335Z","iopub.status.idle":"2023-06-19T04:35:13.136626Z","shell.execute_reply.started":"2023-06-19T04:35:13.126305Z","shell.execute_reply":"2023-06-19T04:35:13.135700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hparams = {\n    # Optional hparams\n    \"backbone\": 'wavenet_4096', #'', #'tf_efficientnetv2_b2',\n    \"learning_rate\": [5e-4],\n    \"max_epochs\": 71,\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':5,\n    'final_div_factor':10,\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-19T04:35:13.137946Z","iopub.execute_input":"2023-06-19T04:35:13.138916Z","iopub.status.idle":"2023-06-19T04:35:13.148926Z","shell.execute_reply.started":"2023-06-19T04:35:13.138881Z","shell.execute_reply":"2023-06-19T04:35:13.148043Z"},"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-19T04:35:13.152010Z","iopub.execute_input":"2023-06-19T04:35:13.152269Z","iopub.status.idle":"2023-06-19T04:35:13.160060Z","shell.execute_reply.started":"2023-06-19T04:35:13.152246Z","shell.execute_reply":"2023-06-19T04:35:13.159167Z"},"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-19T04:35:13.161573Z","iopub.execute_input":"2023-06-19T04:35:13.161937Z","iopub.status.idle":"2023-06-19T04:35:13.172965Z","shell.execute_reply.started":"2023-06-19T04:35:13.161907Z","shell.execute_reply":"2023-06-19T04:35:13.172149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nevents = pd.read_csv(f'{INPUT_PATH}/events.csv')\nevents = events[~events.Type.isnull()]\n\ndefog = pd.read_csv(f'{INPUT_PATH}/defog_metadata.csv')\ndefog = defog[defog.Id.isin(events.Id)].reset_index(drop=True)\n\ntdcsfog = pd.read_csv(f'{INPUT_PATH}/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-19T04:35:13.175945Z","iopub.execute_input":"2023-06-19T04:35:13.176241Z","iopub.status.idle":"2023-06-19T04:35:13.224256Z","shell.execute_reply.started":"2023-06-19T04:35:13.176219Z","shell.execute_reply":"2023-06-19T04:35:13.223460Z"},"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-19T04:35:13.225557Z","iopub.execute_input":"2023-06-19T04:35:13.225893Z","iopub.status.idle":"2023-06-19T04:35:13.235258Z","shell.execute_reply.started":"2023-06-19T04:35:13.225864Z","shell.execute_reply":"2023-06-19T04:35:13.234184Z"},"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-19T04:35:13.236881Z","iopub.execute_input":"2023-06-19T04:35:13.237286Z","iopub.status.idle":"2023-06-19T04:35:13.263350Z","shell.execute_reply.started":"2023-06-19T04:35:13.237257Z","shell.execute_reply":"2023-06-19T04:35:13.262489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nversion='6'\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('FOLD',fn,'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-19T04:35:13.264704Z","iopub.execute_input":"2023-06-19T04:35:13.265052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}