{"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":"cfg = {\n    \"cv_splits\":3,\n    \"param_search\":{\"use_sampled\":[\"True\",\"False\"],\n                    \"num_clusters\":[4,6,8],\n                   \"conv_layers\":[1,2,3],\n                    \"conv_dropout\":[0,0.1,0.2,0.3],\n                    \"batch_size\":[8,16,32],\n                    \"embed_samples\":[64,512], #min max\n                    \"lr\":[1e-4,1e-3], #min max\n                    \"wdecay\":[5e-4,1e-7] # min max\n                    \n                   },\n    #\"embed_samples\":64,\n    #\"batch_size\":8,\n    \"num_epochs\":25,\n    \"es_epochs\":8,\n    \"scheduler_patience\":3\n    #\"start_lr\":1e-4,\n    #\"max_lr\":2.5E-02,\n    #\"wdecay\":1e-6\n}","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:53:41.452345Z","iopub.execute_input":"2022-09-29T15:53:41.452719Z","iopub.status.idle":"2022-09-29T15:53:41.459438Z","shell.execute_reply.started":"2022-09-29T15:53:41.452686Z","shell.execute_reply":"2022-09-29T15:53:41.458475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install pykeops --no-index --find-links=../input/strip-ai-library-fetcher/pykeops/","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:35:29.13777Z","iopub.execute_input":"2022-09-29T15:35:29.138489Z","iopub.status.idle":"2022-09-29T15:35:38.54246Z","shell.execute_reply.started":"2022-09-29T15:35:29.138452Z","shell.execute_reply":"2022-09-29T15:35:38.541199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\nimport optuna\nimport math\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport random\nimport os\nfrom pathlib import Path\nimport time\nfrom pykeops.torch import LazyTensor\nfrom torch.utils.data import Dataset,DataLoader\nimport cv2\nfrom sklearn.model_selection import StratifiedKFold\nfrom tqdm.notebook import tqdm_notebook\nfrom sklearn.metrics import accuracy_score,brier_score_loss,precision_recall_fscore_support,roc_auc_score\nimport csv\nfrom torchvision import transforms\nimport gc\nimport itertools\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-29T15:35:38.545101Z","iopub.execute_input":"2022-09-29T15:35:38.545505Z","iopub.status.idle":"2022-09-29T15:35:52.522292Z","shell.execute_reply.started":"2022-09-29T15:35:38.54546Z","shell.execute_reply":"2022-09-29T15:35:52.521059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed_value):\n    random.seed(seed_value)\n    np.random.seed(seed_value)\n    torch.manual_seed(seed_value)\n    os.environ['PYTHONHASHSEED'] = str(seed_value)    \n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = True\n\nseed = 0\nseed_everything(seed)","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:40:20.306111Z","iopub.execute_input":"2022-09-29T15:40:20.306557Z","iopub.status.idle":"2022-09-29T15:40:20.314844Z","shell.execute_reply.started":"2022-09-29T15:40:20.306522Z","shell.execute_reply":"2022-09-29T15:40:20.313852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Helper functions","metadata":{}},{"cell_type":"code","source":"\ndef KMeans(x, K=10, Niter=10, verbose=False):\n    \"\"\"Implements Lloyd's algorithm for the Euclidean metric.\"\"\"\n\n    start = time.time()\n    \n    N, D = x.shape  # Number of samples, dimension of the ambient space\n    ## Avoid getting an eror when an image doesnt have enough valid tiles.\n    if N < K:\n        print(f\"Image does not have enough tiles! Reverting to max number of clusters({N}) possible..\")\n        K = N\n    c = x[:K, :].clone()  # Simplistic initialization for the centroids\n    #print(c.shape)\n    x_i = LazyTensor(x.view(N, 1, D))  # (N, 1, D) samples\n    c_j = LazyTensor(c.view(1, K, D))  # (1, K, D) centroids\n\n    # K-means loop:\n    # - x  is the (N, D) point cloud,\n    # - cl is the (N,) vector of class labels\n    # - c  is the (K, D) cloud of cluster centroids\n    for i in range(Niter):\n\n        # E step: assign points to the closest cluster -------------------------\n        D_ij = ((x_i - c_j) ** 2).sum(-1)  # (N, K) symbolic squared distances\n        cl = D_ij.argmin(dim=1).long().view(-1)  # Points -> Nearest cluster\n\n        # M step: update the centroids to the normalized cluster average: ------\n        # Compute the sum of points per cluster:\n        c.zero_()\n        c.scatter_add_(0, cl[:, None].repeat(1, D), x)\n\n        # Divide by the number of points per cluster:\n        Ncl = torch.bincount(cl, minlength=K).type_as(c).view(K, 1)\n        c /= Ncl  # in-place division to compute the average\n\n    if verbose:  # Fancy display -----------------------------------------------\n        if torch.cuda.is_available():\n            torch.cuda.synchronize()\n        end = time.time()\n        print(\n            f\"K-means for the Euclidean metric with {N:,} points in dimension {D:,}, K = {K:,}:\"\n        )\n        print(\n            \"Timing for {} iterations: {:.5f}s = {} x {:.5f}s\\n\".format(\n                Niter, end - start, Niter, (end - start) / Niter\n            )\n        )\n\n    return cl, c\n\ndef select_pooling(pooling):\n    if pooling==\"max\":\n        return nn.AdaptiveMaxPool2d((1,1))\n    elif pooling==\"avg\":\n        return nn.AdaptiveAvgPool2d((1,1))\n\ndef select_conv_layers(conv_layers,conv_dropout):\n        if conv_layers == 1:\n            return nn.Sequential(nn.Conv2d(1280,32,1),\n                                         nn.ReLU(),\n                                 nn.Dropout2d(conv_dropout))\n        elif conv_layers == 2:\n            return nn.Sequential(nn.Conv2d(1280,640,1),\n                                nn.ReLU(),\n                                 nn.Dropout2d(conv_dropout),\n                                 nn.Conv2d(640,32,1),\n                                nn.ReLU(),\n                                nn.Dropout2d(conv_dropout))\n        elif conv_layers == 3:\n            return nn.Sequential(nn.Conv2d(1280,640,1),\n                                nn.ReLU(),\n                                 nn.Dropout2d(conv_dropout),\n                                 nn.Conv2d(640,320,1),\n                                nn.ReLU(),\n                                nn.Dropout2d(conv_dropout),\n                                nn.Conv2d(320,32,1),\n                                nn.ReLU(),\n                                nn.Dropout2d(conv_dropout))","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:40:20.643097Z","iopub.execute_input":"2022-09-29T15:40:20.643525Z","iopub.status.idle":"2022-09-29T15:40:20.663861Z","shell.execute_reply.started":"2022-09-29T15:40:20.643492Z","shell.execute_reply":"2022-09-29T15:40:20.662355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Datasets","metadata":{}},{"cell_type":"code","source":"# To read already tiled and embedded dataset   \nclass TrainDataset(Dataset):\n    def __init__(self, df,embed_dir,num_clusters,embed_samples):\n        #self.idxs = df[\"idx\"].values\n        self.patient_ids = df['patient_id'].values\n        self.labels = df['label'].values\n        self.embed_dir = embed_dir\n        self.num_clusters = num_clusters\n        self.embed_samples = embed_samples\n\n    def __len__(self):\n        return len(self.patient_ids)\n\n    def __getitem__(self, idx):\n        #i = self.idxs[idx]\n        patient_id = self.patient_ids[idx]\n        #embeds = torch.load(f\"{self.embed_dir}/{i}_('{patient_id}',).pt\").cuda() # Fix formatting error\n        embeds = torch.load(f\"{self.embed_dir}/('{patient_id}',).pt\").cuda() # Fix formatting error\n        embeds = torch.permute(embeds,(1,2,0)).squeeze() # (1xmxd -> mxd)\n        if len(embeds) > self.embed_samples:\n            idxs = torch.randperm(len(embeds))[:self.embed_samples]\n        else:\n            idxs = range(0,len(embeds))\n        #print(idxs)\n        embeds = embeds[idxs]\n        #print(embeds.shape)\n        clusters,_ = KMeans(embeds,self.num_clusters)\n        cluster_embeds = []\n        for c in range(self.num_clusters):\n            idxs = torch.where(clusters == c,1,0).nonzero()\n            embeds_cluster = embeds[idxs]\n            #print(embeds_cluster.shape)\n            cluster_embeds.append(torch.permute(embeds_cluster,(2,0,1))) # mxdx1\n        label = torch.tensor(float(self.labels[idx][8:10]))\n        #label = torch.tensor(float(self.labels[idx]))\n        label = torch.FloatTensor([1-label,label]).cuda().float().unsqueeze(0)\n        return cluster_embeds, label\n\nclass TestDataset(Dataset):\n    def __init__(self, df,embed_dir,num_clusters,embed_samples):\n        #self.idxs = df[\"idx\"].values\n        self.patient_ids = df['patient_id'].values\n        self.labels = df['label'].values\n        self.embed_dir = embed_dir\n        self.num_clusters = num_clusters\n        self.embed_samples = embed_samples\n\n    def __len__(self):\n        return len(self.patient_ids)\n\n    def __getitem__(self, idx):\n        #i = self.idxs[idx]\n        patient_id = self.patient_ids[idx]\n        embeds = torch.load(f\"{self.embed_dir}/('{patient_id}',).pt\").cuda()\n        if len(embeds) > self.embed_samples:\n            idxs = torch.randperm(len(embeds))[:self.embed_samples]\n        else:\n            idxs = range(0,len(embeds))\n        patient_id  = patient_id[0]\n        embeds = torch.permute(embeds,(1,2,0)).squeeze() # (1xmxd -> mxd)\n        clusters,_ = KMeans(embeds.squeeze(),self.num_clusters)\n        cluster_embeds = []\n        for c in range(self.num_clusters):\n            idxs = torch.where(clusters == c,1,0).nonzero()\n            embeds_cluster = embeds[idxs]\n            #print(embeds_cluster.shape)\n            cluster_embeds.append(torch.permute(embeds_cluster,(2,0,1))) # mxdx1\n        label = torch.tensor(float(self.labels[idx][8:10]))\n        label = torch.FloatTensor([1-label,label]).cuda().float().unsqueeze(0)\n        \n        return cluster_embeds, label","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:40:20.967819Z","iopub.execute_input":"2022-09-29T15:40:20.968175Z","iopub.status.idle":"2022-09-29T15:40:20.985313Z","shell.execute_reply.started":"2022-09-29T15:40:20.968145Z","shell.execute_reply":"2022-09-29T15:40:20.983994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Models","metadata":{}},{"cell_type":"code","source":"class Flatten(nn.Module):\n    def __init__(self, dim=1):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, x): \n        input_shape = x.shape\n        output_shape = [input_shape[i] for i in range(self.dim)] + [-1]\n        return x.view(*output_shape)\n\nclass DeepAttnMIL(nn.Module):\n    def __init__(self,num_clusters,pooling,conv_layers,conv_dropout):\n        super().__init__()\n        self.num_clusters = num_clusters\n        self.mi_conv = nn.Sequential(select_conv_layers(conv_layers,conv_dropout),\n                                     select_pooling(pooling)) # Should result in mx64x1x1 dim\n        self.attention = nn.Sequential(\n                                        nn.Linear(32, 16), # V\n                                        nn.Tanh(),\n                                        nn.Linear(16, 1)  # W\n                                        )\n        self.fc = nn.Sequential(\n            #nn.Dropout(0.5),\n            nn.Linear(32,2)\n        )\n\n    def forward(self,x_batch):\n        batched = []\n        for x in x_batch:\n            # x is list of tensors \n            cluster_embeddings = []\n            for i in range(self.num_clusters):\n                cluster = x[i] # dxmx1\n                #print(\"cluster \",cluster.device)\n                #print(cluster.shape)\n                if cluster.shape[1] == 0:\n                    continue\n                output = self.mi_conv(cluster) # L x 1 x 1\n                #print(\"output \",output.device)\n                #print(\"mi_conv output shape: \", output.shape)\n                output = output.view(output.size()[1], -1) #1 x L\n                #print(\"mi_conv view shape: \", output.shape)\n                cluster_embeddings.append(output)\n            cluster_reps = torch.stack(cluster_embeddings,dim=0) # num_clusters x L\n            cluster_reps = cluster_reps.view(cluster_reps.shape[0],-1)\n            #print(\"cluster_reps \",cluster_reps.device)\n            #print(\"Concatted clusters shape: \", cluster_reps.shape)\n            A = self.attention(cluster_reps) # num_clusters x 1\n            #print(\"A:\",A.shape)\n            A = torch.transpose(A, 1, 0)  # 1 x num_clusters\n            #print(\"A transposed:\", A.shape)\n            A = torch.softmax(A,dim=1)\n            #print(A.shape)\n            M = torch.mm(A, cluster_reps)  # 1 x L\n            #print(\"M \", M.device)\n            batched.append(M)\n        batched = torch.stack(batched,dim=0)\n        #print(\"batch dims\",batched.shape)\n        #print(\"batched\", batched.device)\n        #print(\"M:\",M.shape)\n        #print(M.shape)\n        Y_pred = self.fc(batched).view(-1,2)\n        return Y_pred\n            ","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:40:21.349737Z","iopub.execute_input":"2022-09-29T15:40:21.350448Z","iopub.status.idle":"2022-09-29T15:40:21.364352Z","shell.execute_reply.started":"2022-09-29T15:40:21.350412Z","shell.execute_reply":"2022-09-29T15:40:21.363387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model_name,model, dataloaders_dict, criterion, optimizer,scheduler, num_epochs,save_model,use_es):\n    results = {}\n    best_val_loss = np.inf\n    best_model_brier = np.inf\n    best_model_acc = 0\n    best_model_auc = 0\n    val_losses = []\n    train_losses = []\n    early_stopper = 0\n    for epoch in range(num_epochs):\n        print(\"*\"*10 +f\" Starting Epoch {epoch + 1}/{num_epochs} \" +\"*\"*10)\n        train_loss = np.inf\n        val_loss = np.inf\n        for phase in ['train', 'val']:\n            predictions = []\n            probs = []\n            labels = []\n            if phase == 'train': \n                model.train() \n            else: model.eval()\n            epoch_loss = 0.0\n            dataloader = dataloaders_dict[phase]\n            print(f\"Phase: {phase}\")\n            for embeds,label in tqdm_notebook(dataloader,leave=False):\n                classes = torch.stack(label,dim=0).view(-1,2)\n                #print(classes)\n                optimizer.zero_grad()                \n                with torch.set_grad_enabled(phase == 'train'):\n                    \n                    output = model(embeds)\n                    #print(output.shape)\n                    soft_output = torch.softmax(output,dim=1)#.view(-1,2)\n                    #print(soft_output.shape)\n                    prob_out = soft_output.detach().cpu().numpy()\n                    #print(prob_out)\n                    probs.append(prob_out[:,1])\n                    _,preds = torch.max(torch.FloatTensor(prob_out),1)\n                    #print(preds)\n                    predictions.append(preds.detach().cpu().numpy())\n                    _,lbl = torch.max(classes,1)\n                    #print(lbl)\n                    labels.append(lbl.detach().cpu().numpy())\n                    #print(labels)\n                    #print(\"output\",output)\n                    #print(\"classes\",classes)\n                    loss = criterion(output, classes)\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n                        \n                    epoch_loss += loss.item() * len(output)                 \n            data_size = len(dataloader.dataset)\n            labels = np.concatenate(labels)\n            preds = np.concatenate(predictions)\n            probs = np.concatenate(probs)\n            # Epoch scores\n            epoch_loss = epoch_loss / data_size\n            acc = accuracy_score(labels,preds)\n            brier_loss = brier_score_loss(labels,probs)\n            auc = roc_auc_score(labels, probs,average=\"weighted\")\n            if phase == 'train':\n                train_loss = epoch_loss\n                train_losses.append(epoch_loss)\n            else:\n                val_loss = epoch_loss\n                scheduler.step(val_loss)\n                val_losses.append(epoch_loss)\n            print(f'{model_name}:\\\n                   Loss: {epoch_loss:.4f} |  Brier: {brier_loss:.4f} | \\n \\\n                   Acc: {acc:.4f} | AUC:{auc:.4f}')   \n            \n        if best_val_loss > val_loss:\n            print(\"Updating the best model\")\n            if save_model == True:\n                torch.save(model.state_dict(), f\"{model_name}.pth\")\n            best_val_loss = val_loss\n            best_model_acc = acc\n            best_model_brier = brier_loss\n            best_model_auc = auc\n            early_stopper = 0\n        else:\n            early_stopper +=1\n        results = {\"model_name\":model_name,\n                   \"val_loss\":best_val_loss,\"brier\":best_model_brier,\n                   \"acc\":best_model_acc,\"AUC\":best_model_auc,\n                  \"train_losses\":train_losses,\"val_losses\":val_losses}\n        if use_es and early_stopper == cfg[\"es_epochs\"]:\n            print(\"Stopped training since val loss didn't improve\")\n            return results\n    return results\n","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:40:21.560945Z","iopub.execute_input":"2022-09-29T15:40:21.561234Z","iopub.status.idle":"2022-09-29T15:40:21.576875Z","shell.execute_reply.started":"2022-09-29T15:40:21.561206Z","shell.execute_reply":"2022-09-29T15:40:21.575782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def objective_cv(trial):\n    \n    # Sample params\n    #use_sampled = trial.suggest_categorical(\"use_sampled\",cfg[\"param_search\"][\"use_sampled\"])\n    num_clusters = trial.suggest_categorical('num_clusters', cfg[\"param_search\"][\"num_clusters\"])\n    conv_layers = trial.suggest_categorical(\"conv_layers\",cfg[\"param_search\"][\"conv_layers\"])\n    batch_size = trial.suggest_categorical(\"batch_size\",cfg[\"param_search\"][\"batch_size\"])\n                    \n    conv_dropout = trial.suggest_categorical(\"conv_dropout\",cfg[\"param_search\"][\"conv_dropout\"])\n    embed_samples = trial.suggest_int('embed_samples', cfg[\"param_search\"][\"embed_samples\"][0],\n                                      cfg[\"param_search\"][\"embed_samples\"][1],step=32)\n    lr = trial.suggest_loguniform(\"lr\",cfg[\"param_search\"][\"lr\"][0],cfg[\"param_search\"][\"lr\"][1])\n#     wdecay = trial.suggest_loguniform(\"wdecay\",cfg[\"param_search\"][\"wdecay\"][0],\n#                                       cfg[\"param_search\"][\"wdecay\"][1])\n    pooling = \"max\" # Better compared to \"avg\" (from last cv attempt)\n    \n    \n    cv = StratifiedKFold(n_splits=cfg[\"cv_splits\"])\n    losses = []\n    \n    # Load csv\n    ## Embedded, transformed and over-sampled dataset\n    sampled = pd.read_csv(\"../input/strip-ai-embeddings-training-data/embedded_train.csv\")\n    ## Embedded dataset\n    unsampled = pd.read_csv(\"../input/strip-ai-effnet-embeddings-validation/embedded_val.csv\")\n    X = unsampled\n    X[\"patient_id\"] = X[\"patient_id\"].apply(lambda x: x.strip(\"('',)\"))\n    y = unsampled.label\n    for train_idxs, val_idxs in cv.split(X, y):\n        val_data = unsampled.iloc[val_idxs]\n        train_data = unsampled.iloc[train_idxs]\n        #Create datasets and dataloaders\n        trainset = TrainDataset(train_data,\n                               embed_dir=\"../input/strip-ai-effnet-embeddings-validation/output/val\",\n                               num_clusters = num_clusters,\n                               embed_samples=embed_samples)\n        train_loader = DataLoader(trainset,batch_size=batch_size,collate_fn=collate_fn)\n        valset = TestDataset(val_data,\n                            embed_dir=\"../input/strip-ai-effnet-embeddings-validation/output/val\",\n                            num_clusters = num_clusters,\n                            embed_samples=embed_samples)\n        val_loader = DataLoader(valset,batch_size=batch_size,collate_fn=collate_fn)\n        dataloader_dict = {\"train\":train_loader,\"val\":val_loader}\n        # Create model using the sampled hyperparams\n        model = DeepAttnMIL(num_clusters,pooling,conv_layers,conv_dropout).to('cuda:0')\n        model_name = f\"{num_clusters}_{conv_layers}_{batch_size}_{conv_dropout}_{embed_samples}_{lr:.6f}\"\n        criterion = torch.nn.CrossEntropyLoss().cuda()\n        optimizer = torch.optim.SGD(model.parameters(), lr=lr)\n        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,patience=cfg[\"scheduler_patience\"],\n                                                               min_lr=1e-7,verbose=True)\n        results = train_model(model_name,model,dataloader_dict,criterion,\n                              optimizer,scheduler,cfg[\"num_epochs\"],save_model=False,use_es=True)\n        losses.append(results['val_loss'])    \n    return np.mean(losses)\n\n","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:52:11.155613Z","iopub.execute_input":"2022-09-29T15:52:11.155991Z","iopub.status.idle":"2022-09-29T15:52:11.169421Z","shell.execute_reply.started":"2022-09-29T15:52:11.155958Z","shell.execute_reply":"2022-09-29T15:52:11.168153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Use Optuna to find the best hyperparameters","metadata":{}},{"cell_type":"code","source":"def collate_fn(list_items):\n    x = []\n    y = []\n    for x_, y_ in list_items:\n        #print(f'x_={x_}, y_={y_}')\n        x.append(x_)\n        y.append(y_)\n    return x, y","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:52:12.036761Z","iopub.execute_input":"2022-09-29T15:52:12.037887Z","iopub.status.idle":"2022-09-29T15:52:12.044656Z","shell.execute_reply.started":"2022-09-29T15:52:12.037835Z","shell.execute_reply":"2022-09-29T15:52:12.043536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Don't want notebook to crash when saving because of logs\noptuna.logging.set_verbosity(optuna.logging.WARNING) \nstudy = optuna.create_study()\nstudy.optimize(objective_cv, gc_after_trial=True,timeout=21600) # Search for 6 hours 21600\n","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:54:04.034468Z","iopub.execute_input":"2022-09-29T15:54:04.034852Z","iopub.status.idle":"2022-09-29T16:07:25.359053Z","shell.execute_reply.started":"2022-09-29T15:54:04.034816Z","shell.execute_reply":"2022-09-29T16:07:25.357866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optuna.visualization.plot_optimization_history(study)","metadata":{"execution":{"iopub.status.busy":"2022-09-29T16:07:38.560685Z","iopub.execute_input":"2022-09-29T16:07:38.561159Z","iopub.status.idle":"2022-09-29T16:07:38.669181Z","shell.execute_reply.started":"2022-09-29T16:07:38.561122Z","shell.execute_reply":"2022-09-29T16:07:38.668134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optuna.visualization.plot_slice(study)","metadata":{"execution":{"iopub.status.busy":"2022-09-29T16:07:38.981443Z","iopub.execute_input":"2022-09-29T16:07:38.982113Z","iopub.status.idle":"2022-09-29T16:07:39.208463Z","shell.execute_reply.started":"2022-09-29T16:07:38.982076Z","shell.execute_reply":"2022-09-29T16:07:39.207511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Best loss:\",study.best_value)","metadata":{"execution":{"iopub.status.busy":"2022-09-29T16:07:42.669828Z","iopub.execute_input":"2022-09-29T16:07:42.670431Z","iopub.status.idle":"2022-09-29T16:07:42.675982Z","shell.execute_reply.started":"2022-09-29T16:07:42.670394Z","shell.execute_reply":"2022-09-29T16:07:42.675002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(study.best_params)","metadata":{"execution":{"iopub.status.busy":"2022-09-29T16:07:43.050103Z","iopub.execute_input":"2022-09-29T16:07:43.050465Z","iopub.status.idle":"2022-09-29T16:07:43.059321Z","shell.execute_reply.started":"2022-09-29T16:07:43.050433Z","shell.execute_reply":"2022-09-29T16:07:43.058321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study.trials_dataframe().to_csv(\"study.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-09-29T16:07:44.005453Z","iopub.execute_input":"2022-09-29T16:07:44.005833Z","iopub.status.idle":"2022-09-29T16:07:44.021614Z","shell.execute_reply.started":"2022-09-29T16:07:44.005781Z","shell.execute_reply":"2022-09-29T16:07:44.020516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train a final model using the best parameters","metadata":{}},{"cell_type":"markdown","source":"##### Get best params\nbest_params = cfg[\"param_search\"]#study.best_params\n#del study\ngc.collect()\nnum_clusters = best_params[\"num_clusters\"][0]\nconv_layers = best_params[\"conv_layers\"][0]\nconv_dropout =best_params[\"conv_dropout\"][0]\npooling = \"max\"\n\n# Load csv\n## Embedded, transformed and over-sampled dataset\nsampled = pd.read_csv(\"../input/strip-ai-embeddings-training-data/embedded_train.csv\")\n## Embedded dataset\nunsampled = pd.read_csv(\"../input/strip-ai-effnet-embeddings-validation/embedded_val.csv\")\nX = unsampled\nX[\"patient_id\"] = X[\"patient_id\"].apply(lambda x: x.strip(\"('',)\"))\ny = unsampled.label\ntrain_data,val_data,_,_ = train_test_split(X,y,test_size=0.15,random_state=97,stratify=y)\n\n#Find patients in valset and remove them from train\n\n#val_patients = val_data['patient_id'].values\n#print(val_patients)\n#print(sampled.patient_id)\n#train_data = sampled[~sampled.patient_id.isin(val_patients)]\n# Create model using the sampled hyperparams\nmodel = DeepAttnMIL(num_clusters,pooling,conv_layers,conv_dropout).cuda()\nmodel_name = f\"DeepAttnMIL_{num_clusters}_{pooling}_{conv_layers}_{conv_dropout}\"\ncriterion = torch.nn.CrossEntropyLoss().cuda()\noptimizer = torch.optim.SGD(model.parameters(),lr=1e-3,momentum=0.9)\n#optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4,weight_decay=1e-5)\n#Create datasets and dataloaders\ntrainset = TrainDataset(train_data,\n                       embed_dir=\"../input/strip-ai-effnet-embeddings-validation/output/val\",\n                       num_clusters = num_clusters,\n                       embed_samples=256)\ntrain_loader = DataLoader(trainset,batch_size=8,collate_fn=collate_fn)\n#scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer,\n#                                                 max_lr=cfg[\"max_lr\"],\n#                                                 steps_per_epoch=len(train_loader), epochs=cfg[\"num_epochs\"])\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,patience=3, min_lr=1e-7,verbose=True)\n# lr_finder = LRFinder(model, optimizer, criterion, device=\"cuda\")\n# lr_finder.range_test(train_loader, end_lr=100, num_iter=100)\n# lr_finder.plot() # to inspect the loss-learning rate graph\n# lr_finder.reset() # to reset the model and optimizer to their initial state\n\n\nvalset = TestDataset(val_data,\n                    embed_dir=\"../input/strip-ai-effnet-embeddings-validation/output/val\",\n                    num_clusters = num_clusters,embed_samples=256)\nval_loader = DataLoader(valset,batch_size=8,collate_fn=collate_fn)\ndataloader_dict = {\"train\":train_loader,\"val\":val_loader}\n\nresults = train_model(model_name,model,dataloader_dict,criterion,\n                      optimizer,scheduler,cfg[\"num_epochs\"],save_model=True,use_es=False)\nprint(f\"Best val loss: {results['val_loss']}\")\nplt.figure(figsize=(10,5))\nplt.title(\"Training and Validation Loss\")\nplt.plot(results[\"val_losses\"],label=\"Val\")\nplt.plot(results[\"train_losses\"],label=\"Train\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-29T15:12:13.616602Z","iopub.execute_input":"2022-09-29T15:12:13.616984Z","iopub.status.idle":"2022-09-29T15:18:19.821589Z","shell.execute_reply.started":"2022-09-29T15:12:13.616949Z","shell.execute_reply":"2022-09-29T15:18:19.820661Z"}}},{"cell_type":"markdown","source":"### Calibrate the model using temperature scaling and validation set","metadata":{}},{"cell_type":"markdown","source":"### Taken from https://github.com/gpleiss/temperature_scaling/blob/master/temperature_scaling.py\nfrom torch import nn, optim\nfrom torch.nn import functional as F\nclass ModelWithTemperature(nn.Module):\n    \"\"\"\n    A thin decorator, which wraps a model with temperature scaling\n    model (nn.Module):\n        A classification neural network\n        NB: Output of the neural network should be the classification logits,\n            NOT the softmax (or log softmax)!\n    \"\"\"\n    def __init__(self, model):\n        super(ModelWithTemperature, self).__init__()\n        self.model = model\n        self.temperature = nn.Parameter(torch.ones(1) * 1.5)\n\n    def forward(self, input):\n        logits = self.model(input)\n        return self.temperature_scale(logits)\n\n    def temperature_scale(self, logits):\n        \"\"\"\n        Perform temperature scaling on logits\n        \"\"\"\n        # Expand temperature to match the size of logits\n        temperature = self.temperature.unsqueeze(1).expand(logits.size(0), logits.size(1))\n        return logits / temperature\n\n    # This function probably should live outside of this class, but whatever\n    def set_temperature(self, valid_loader):\n        \"\"\"\n        Tune the tempearature of the model (using the validation set).\n        We're going to set it to optimize NLL.\n        valid_loader (DataLoader): validation set loader\n        \"\"\"\n        self.cuda()\n        nll_criterion = nn.CrossEntropyLoss().cuda()\n        ece_criterion = _ECELoss().cuda()\n\n        # First: collect all the logits and labels for the validation set\n        logits_list = []\n        labels_list = []\n        with torch.no_grad():\n            for input, label in valid_loader:\n                label = torch.stack(label,dim=0).view(-1,2)\n                _,label = torch.max(label,1)\n                #print(label.shape)\n                logits = self.model(input)\n                #print(logits.shape)\n                logits_list.append(logits)\n                labels_list.append(label)\n            #print(logits_list)\n            logits = torch.cat(logits_list).cuda()\n            labels = torch.cat(labels_list).cuda()\n        # Calculate NLL and ECE before temperature scaling\n        before_temperature_nll = nll_criterion(logits, labels).item()\n        before_temperature_ece = ece_criterion(logits, labels).item()\n        print('Before temperature - NLL: %.3f, ECE: %.3f' % (before_temperature_nll, before_temperature_ece))\n\n        # Next: optimize the temperature w.r.t. NLL\n        optimizer = optim.LBFGS([self.temperature], lr=0.01, max_iter=50)\n\n        def eval():\n            optimizer.zero_grad()\n            loss = nll_criterion(self.temperature_scale(logits), labels)\n            loss.backward()\n            return loss\n        optimizer.step(eval)\n\n        # Calculate NLL and ECE after temperature scaling\n        after_temperature_nll = nll_criterion(self.temperature_scale(logits), labels).item()\n        after_temperature_ece = ece_criterion(self.temperature_scale(logits), labels).item()\n        print('Optimal temperature: %.3f' % self.temperature.item())\n        print('After temperature - NLL: %.3f, ECE: %.3f' % (after_temperature_nll, after_temperature_ece))\n\n        return self\n\n\nclass _ECELoss(nn.Module):\n    \"\"\"\n    Calculates the Expected Calibration Error of a model.\n    (This isn't necessary for temperature scaling, just a cool metric).\n    The input to this loss is the logits of a model, NOT the softmax scores.\n    This divides the confidence outputs into equally-sized interval bins.\n    In each bin, we compute the confidence gap:\n    bin_gap = | avg_confidence_in_bin - accuracy_in_bin |\n    We then return a weighted average of the gaps, based on the number\n    of samples in each bin\n    See: Naeini, Mahdi Pakdaman, Gregory F. Cooper, and Milos Hauskrecht.\n    \"Obtaining Well Calibrated Probabilities Using Bayesian Binning.\" AAAI.\n    2015.\n    \"\"\"\n    def __init__(self, n_bins=15):\n        \"\"\"\n        n_bins (int): number of confidence interval bins\n        \"\"\"\n        super(_ECELoss, self).__init__()\n        bin_boundaries = torch.linspace(0, 1, n_bins + 1)\n        self.bin_lowers = bin_boundaries[:-1]\n        self.bin_uppers = bin_boundaries[1:]\n\n    def forward(self, logits, labels):\n        softmaxes = F.softmax(logits, dim=1)\n        confidences, predictions = torch.max(softmaxes, 1)\n        #print(predictions)\n        accuracies = predictions.eq(labels)\n\n        ece = torch.zeros(1, device=logits.device)\n        for bin_lower, bin_upper in zip(self.bin_lowers, self.bin_uppers):\n            # Calculated |confidence - accuracy| in each bin\n            in_bin = confidences.gt(bin_lower.item()) * confidences.le(bin_upper.item())\n            prop_in_bin = in_bin.float().mean()\n            if prop_in_bin.item() > 0:\n                accuracy_in_bin = accuracies[in_bin].float().mean()\n                avg_confidence_in_bin = confidences[in_bin].mean()\n                ece += torch.abs(avg_confidence_in_bin - accuracy_in_bin) * prop_in_bin\n\n        return ece","metadata":{"execution":{"iopub.status.busy":"2022-09-27T07:03:27.874531Z","iopub.execute_input":"2022-09-27T07:03:27.874897Z","iopub.status.idle":"2022-09-27T07:03:27.896671Z","shell.execute_reply.started":"2022-09-27T07:03:27.874867Z","shell.execute_reply":"2022-09-27T07:03:27.895542Z"}}},{"cell_type":"markdown","source":"model = DeepAttnMIL(num_clusters,pooling,conv_layers,conv_dropout).cuda()\nmodel.load_state_dict(torch.load(f\"./{model_name}.pth\")) # create an uncalibrated model somehow\nscaled_model = ModelWithTemperature(model)\nscaled_model.set_temperature(val_loader)\n# Save final model\ntorch.save(scaled_model.state_dict(), f\"final_model.pth\")\n# Save hyperparams\npd.DataFrame([best_params]).to_csv(\"hyperparams.csv\",index=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-27T07:03:28.646858Z","iopub.execute_input":"2022-09-27T07:03:28.647228Z","iopub.status.idle":"2022-09-27T07:03:30.44582Z","shell.execute_reply.started":"2022-09-27T07:03:28.647197Z","shell.execute_reply":"2022-09-27T07:03:30.4448Z"}}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}