{"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    \"max_size\":20000,\n    \"tile_size\" : 224,\n    \"fg_perc\" : 0.15, # Assume x percentage of image is foreground\n    \"batch_size\":64,\n}","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:32:59.228379Z","iopub.execute_input":"2022-09-26T20:32:59.228827Z","iopub.status.idle":"2022-09-26T20:32:59.234189Z","shell.execute_reply.started":"2022-09-26T20:32:59.228792Z","shell.execute_reply":"2022-09-26T20:32:59.233062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!conda install ../input/strip-ai-library-fetcher/*.tar.bz2","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:27:08.723296Z","iopub.execute_input":"2022-09-26T20:27:08.72365Z","iopub.status.idle":"2022-09-26T20:29:52.967767Z","shell.execute_reply.started":"2022-09-26T20:27:08.723621Z","shell.execute_reply":"2022-09-26T20:29:52.966493Z"},"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-26T20:29:52.970305Z","iopub.execute_input":"2022-09-26T20:29:52.971062Z","iopub.status.idle":"2022-09-26T20:30:07.285373Z","shell.execute_reply.started":"2022-09-26T20:29:52.971016Z","shell.execute_reply":"2022-09-26T20:30:07.284131Z"},"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 math\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\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\nimport openslide\nimport pyvips\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom torch import nn, optim\nfrom torch.nn import functional as F\nimport torch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-26T20:30:07.288825Z","iopub.execute_input":"2022-09-26T20:30:07.289235Z","iopub.status.idle":"2022-09-26T20:30:23.445573Z","shell.execute_reply.started":"2022-09-26T20:30:07.289189Z","shell.execute_reply":"2022-09-26T20:30:23.444497Z"},"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-26T20:56:32.472476Z","iopub.execute_input":"2022-09-26T20:56:32.472845Z","iopub.status.idle":"2022-09-26T20:56:32.480208Z","shell.execute_reply.started":"2022-09-26T20:56:32.472812Z","shell.execute_reply":"2022-09-26T20:56:32.478982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"format_to_dtype = {\n       'uchar': np.uint8,\n       'char': np.int8,\n       'ushort': np.uint16,\n       'short': np.int16,\n       'uint': np.uint32,\n       'int': np.int32,\n       'float': np.float32,\n       'double': np.float64,\n       'complex': np.complex64,\n       'dpcomplex': np.complex128,\n    }\ndef vips2numpy(vi):\n\n        return np.ndarray(\n            buffer=vi.write_to_memory(),\n            dtype=format_to_dtype[vi.format],\n            shape=[vi.height, vi.width, vi.bands])\n# Can open at least 35kx35kx3 images\ndef load_image(img_path, max_size):\n    slide = openslide.OpenSlide(img_path)\n    x_dim,y_dim = slide.dimensions\n\n    del slide\n    gc.collect()\n    size = x_dim * y_dim\n    size_limit = max_size*max_size\n    if size > size_limit:\n        resize_factor = math.sqrt(size / size_limit)\n        width = int(x_dim / resize_factor)\n        height= int(y_dim / resize_factor)\n \n        img = pyvips.Image.thumbnail(img_path,width,height=height) \n    else:\n        img = pyvips.Image.new_from_file(img_path)\n    #img = vips2numpy(img)\n    return img.numpy()\ndef make_tiles(img, tile_size=224, fg_perc = 0.2):\n    '''\n    img: np.ndarray with dtype np.uint8 and shape (width, height, channel)\n    '''\n\n    w, h, ch = img.shape\n    num_tiles = int((w * h / (tile_size*tile_size)) * fg_perc) \n    pad0, pad1 = (tile_size - w%tile_size) % tile_size, (tile_size - h%tile_size) % tile_size\n    padding = [[pad0//2, pad0-pad0//2], [pad1//2, pad1-pad1//2], [0, 0]]\n    img = np.pad(img, padding, mode='constant', constant_values=255)\n    img = img.reshape(img.shape[0]//tile_size, tile_size, img.shape[1]//tile_size, tile_size, ch)\n    img = img.transpose(0, 2, 1, 3, 4).reshape(-1, tile_size, tile_size, ch)\n    if len(img) < num_tiles: # pad images so that the output shape be the same\n        padding = [[0, num_tiles-len(img)], [0, 0], [0, 0], [0, 0]]\n        img = np.pad(img, padding, mode='constant', constant_values=255)\n    sort = np.argsort(img.reshape(img.shape[0], -1).sum(-1))\n    idxs = sort[:num_tiles]\n    img = img[idxs]\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:30:23.493411Z","iopub.execute_input":"2022-09-26T20:30:23.493796Z","iopub.status.idle":"2022-09-26T20:30:23.50971Z","shell.execute_reply.started":"2022-09-26T20:30:23.493759Z","shell.execute_reply":"2022-09-26T20:30:23.508819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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_dropout ==\"false\":\n        if conv_layers == 1:\n            return nn.Sequential(nn.Conv2d(1280,32,1),\n                                         nn.ReLU())\n        elif conv_layers == 2:\n            return nn.Sequential(nn.Conv2d(40,2048,1),\n                                nn.ReLU(),\n                                 nn.Conv2d(2048,64,1),\n                                nn.ReLU())\n        elif conv_layers == 3:\n            return nn.Sequential(nn.Conv2d(4096,2048,1),\n                                nn.ReLU(),\n                                 nn.Conv2d(2048,1024,1),\n                                nn.ReLU(),\n                                nn.Conv2d(1024,64,1),\n                                nn.ReLU())\n    else:\n        if conv_layers == 1:\n            return nn.Sequential(nn.Conv2d(1280,32,1),\n                                         nn.ReLU(),\n                                 nn.Dropout2d(0.3))\n        elif conv_layers == 2:\n            return nn.Sequential(nn.Conv2d(4096,2048,1),\n                                nn.ReLU(),\n                                 nn.Dropout2d(0.2),\n                                 nn.Conv2d(2048,64,1),\n                                nn.ReLU())\n        elif conv_layers == 3:\n            return nn.Sequential(nn.Conv2d(4096,2048,1),\n                                nn.ReLU(),\n                                 nn.Dropout2d(0.2),\n                                 nn.Conv2d(2048,1024,1),\n                                nn.ReLU(),\n                                nn.Conv2d(1024,64,1),\n                                nn.ReLU())","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:30:23.511478Z","iopub.execute_input":"2022-09-26T20:30:23.511909Z","iopub.status.idle":"2022-09-26T20:30:23.53096Z","shell.execute_reply.started":"2022-09-26T20:30:23.51186Z","shell.execute_reply":"2022-09-26T20:30:23.528974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = A.Compose([\n                      A.Normalize(\n                        mean=[0.485, 0.456, 0.406],\n                        std=[0.229, 0.224, 0.225]),\n                        ToTensorV2()\n                      ])","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:56:41.458632Z","iopub.execute_input":"2022-09-26T20:56:41.459019Z","iopub.status.idle":"2022-09-26T20:56:41.464125Z","shell.execute_reply.started":"2022-09-26T20:56:41.458988Z","shell.execute_reply":"2022-09-26T20:56:41.463111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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.Linear(32,2)\n        )\n\n    def forward(self,x):\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        #print(\"batched\", batched.device)\n        #print(\"M:\",M.shape)\n        #print(M.shape)\n        Y_pred = self.fc(M)\n        return Y_pred\n            ","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:56:41.743252Z","iopub.execute_input":"2022-09-26T20:56:41.744185Z","iopub.status.idle":"2022-09-26T20:56:41.757135Z","shell.execute_reply.started":"2022-09-26T20:56:41.744141Z","shell.execute_reply":"2022-09-26T20:56:41.75599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\nclass EmbeddedDataset(Dataset):\n    def __init__(self, df, encoder,transform,num_clusters):\n        self.patient_ids = df['patient_id'].values\n        self.image_ids = df[\"image_id\"].values\n        self.image_dir = df[\"dir\"].values\n        self.encoder = nn.Sequential(encoder.stem,\n                                     encoder.layers,\n                                     encoder.features,\n                                     nn.AdaptiveAvgPool2d(output_size=1),\n                                     Flatten(dim=1),\n                                    )\n        self.transform = transform\n        self.num_clusters= num_clusters\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n    def __len__(self):\n        return len(self.patient_ids)\n    def __getitem__(self,idx):\n        start = time.time()\n        patient_id = self.patient_ids[idx]\n        try:\n            embeds = []\n            batch = []\n            batch_size = cfg[\"batch_size\"]\n            batch_counter = 0\n            for i,image_id in enumerate(tqdm_notebook(self.image_ids[idx],leave=False)):\n                img = load_image(f\"{self.image_dir[idx]}/{image_id}.tif\",cfg[\"max_size\"])\n                print(\"Loaded image\")\n                #img = torch.zeros(20000,20000,3)\n                tiles = make_tiles(img,cfg[\"tile_size\"],cfg[\"fg_perc\"])\n                print(\"Created tiles\")\n                del img\n                length = len(tiles)\n                for tile in tqdm_notebook(tiles):\n                    tile = self.transform(image=tile)[\"image\"]\n                    batch_counter+=1\n                    batch.append(tile)\n                    if batch_counter == batch_size or i+1 == length:\n                        batch = torch.stack(batch,dim=0)\n                        embed = self.encoder(batch.cuda())\n                        embeds.append(embed)\n                        batch = []\n                        batch_counter=0\n                del tiles\n            patient_embeds = torch.cat(embeds,dim=0)\n            #print(patient_embeds.shape)\n            patient_embeds = torch.flatten(patient_embeds,start_dim=1)\n            #print(patient_embeds.shape)\n            clusters,_ = KMeans(patient_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 = patient_embeds[idxs]\n                #print(embeds_cluster.shape)\n                cluster_embeds.append(torch.permute(embeds_cluster,(2,0,1))) # mxdx1\n            end = time.time()\n            print(f\"Processed image in  {end-start} seconds..\")\n            return patient_id,cluster_embeds\n        except:\n            return patient_id,None","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:56:42.622775Z","iopub.execute_input":"2022-09-26T20:56:42.623149Z","iopub.status.idle":"2022-09-26T20:56:42.639947Z","shell.execute_reply.started":"2022-09-26T20:56:42.623111Z","shell.execute_reply":"2022-09-26T20:56:42.638796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nencoder = torch.jit.load('../input/effnet-b0-offline/nvidia_effnet_b0.pt').cuda().eval()\n# Read hyperparams\nparams = pd.read_csv(\"../input/strip-ai-cluster-attention-training/hyperparams.csv\")\nnum_clusters= int(params[\"num_clusters\"][0].strip(\"[]\"))\npooling = \"max\"\nconv_layers = int(params[\"conv_layers\"][0].strip(\"[]\"))\nconv_dropout = params[\"conv_dropout\"][0].strip(\"[]\")\n# Read csv\ntest = pd.read_csv('../input/mayo-clinic-strip-ai/test.csv')\ntest = test.groupby(\"patient_id\").agg(list).reset_index()\ntest[\"dir\"] = \"../input/mayo-clinic-strip-ai/test\"\n#Load models\nmodel = DeepAttnMIL(num_clusters,pooling,conv_layers,conv_dropout).cuda()\nmodel.load_state_dict(torch.load(\"../input/final-model/DeepAttnMIL_6_max_1_true.pth\"))\n#scaled_model = ModelWithTemperature(model).cuda()\n\ndataset = EmbeddedDataset(test,encoder,transform,num_clusters)\nloader = DataLoader(dataset,batch_size=1, shuffle=False)\n\ndef process_probs(x):\n    return max(min(x,1-1e-15),1e-15)\nfunc = np.vectorize(process_probs)\n\npredictions = []\nfor patient_id,cluster_embeds in loader:\n    try:\n        if cluster_embeds is None:\n            predictions.append({\"patient_id\":patient_id,\"CE\":0.5,\"LAA\":0.5})\n            continue\n        patient_id = patient_id[0]\n        output = model(cluster_embeds)\n        soft_output = torch.softmax(output,dim=1)\n        prob_out = soft_output.detach().cpu().numpy()\n        prob_out = func(prob_out)\n        predictions.append({\"patient_id\":patient_id,\"CE\":prob_out[0,0],\"LAA\":prob_out[0,1]})\n    except:\n        predictions.append({\"patient_id\":patient_id,\"CE\":0.5,\"LAA\":0.5})\npd.DataFrame(predictions).to_csv(\"/kaggle/working/submission.csv\",float_format='%.6f',index=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-26T20:56:44.320573Z","iopub.execute_input":"2022-09-26T20:56:44.320944Z","iopub.status.idle":"2022-09-26T20:58:25.691836Z","shell.execute_reply.started":"2022-09-26T20:56:44.320897Z","shell.execute_reply":"2022-09-26T20:58:25.690967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}