{"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":"markdown","source":"<div style=\"padding:20px;color:#2b2d42;margin:0;font-size:180%;text-align:center;display:fill;border-radius:5px;background-color:white;overflow:hidden;font-weight:600\">[Inference] Binary segmentation UNeXt50</div>\n\n<img src=\"https://drive.google.com/uc?id=1pbIvjTlhGywfhiMTqcsdOB5LSHlklM90\" style=\"border-radius:5px\">\n\n<h5 style=\"text-align: center; font-family: Verdana; font-size: 12px; font-style: normal; font-weight: bold; text-decoration: None; text-transform: none; letter-spacing: 1px; color: black; background-color: #ffffff;\">OUR TEAM: NGHI HUYNH, YUAN HONG, MATTEO CACCIOLA</h5>","metadata":{}},{"cell_type":"markdown","source":"# <div style=\"padding:20px;color:white;margin:0;font-size:100%;text-align:left;display:fill;border-radius:5px;background-color:#735d78;overflow:hidden\">1. Imports</div>\n","metadata":{"papermill":{"duration":0.00991,"end_time":"2021-03-12T06:33:14.88117","exception":false,"start_time":"2021-03-12T06:33:14.87126","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport rasterio\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\nimport tifffile as tiff\nfrom albumentations import *\nimport matplotlib.pyplot as plt\nfrom fastai.vision.all import *\nfrom rasterio.windows import Window\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport warnings; warnings.filterwarnings(\"ignore\")","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":3.066435,"end_time":"2021-03-12T06:33:17.956368","exception":false,"start_time":"2021-03-12T06:33:14.889933","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:43.688567Z","iopub.execute_input":"2022-09-23T01:50:43.689334Z","iopub.status.idle":"2022-09-23T01:50:43.699297Z","shell.execute_reply.started":"2022-09-23T01:50:43.689286Z","shell.execute_reply":"2022-09-23T01:50:43.698301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <div style=\"padding:20px;color:white;margin:0;font-size:100%;text-align:left;display:fill;border-radius:5px;background-color:#735d78;overflow:hidden\">2. Config</div>","metadata":{}},{"cell_type":"code","source":"bs = 16\nsz = 512   # the size of tiles\n#reduce = 2  # reduce the original images by 4 times\n#TH = 0.225  # threshold for positive predictions\nDATA = '../input/hubmap-organ-segmentation/test_images/'\nTRAIN = '../input/hubmap-organ-segmentation/train_images'\n#ensemble models?\nMODELS = [f'../input/unext50-wsi-models/run_5/run_5/model_{i}.pth' for i in range(2)]\\\n+ [f'../input/unext50-wsi-models/run_4/run_4/model_{i}.pth' for i in range(2)]\\\n+ [f'../input/unext50-wsi-models/run_4/run_4/model_3.pth']\\\n+ [f'../input/unext50-wsi-models/run_6/run_6/model_0.pth']\\\n+ [f'../input/unext50-wsi-models/run_7/run_7/model_{i}.pth' for i in range(2)]\\\n+ [f'../input/unext50-wsi-models/run_7/run_7/model_1.pth']\n\ndf_sample = pd.read_csv('../input/hubmap-organ-segmentation/sample_submission.csv') \ntest_df = pd.read_csv('../input/hubmap-organ-segmentation/test.csv').set_index('id')\ntrain_df = pd.read_csv('../input/hubmap-organ-segmentation/train.csv').set_index('id')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":0.024698,"end_time":"2021-03-12T06:33:17.991398","exception":false,"start_time":"2021-03-12T06:33:17.9667","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:45.768040Z","iopub.execute_input":"2022-09-23T01:50:45.768460Z","iopub.status.idle":"2022-09-23T01:50:46.189235Z","shell.execute_reply.started":"2022-09-23T01:50:45.768427Z","shell.execute_reply":"2022-09-23T01:50:46.188293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <div style=\"padding:20px;color:white;margin:0;font-size:100%;text-align:left;display:fill;border-radius:5px;background-color:#735d78;overflow:hidden\">3. Data</div>","metadata":{"papermill":{"duration":0.008314,"end_time":"2021-03-12T06:33:18.008555","exception":false,"start_time":"2021-03-12T06:33:18.000241","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# functions to convert encoding to mask and mask to encoding\ndef enc2mask(encs, shape):\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for m,enc in enumerate(encs):\n        if isinstance(enc,np.float) and np.isnan(enc): continue\n        s = enc.split()\n        for i in range(len(s)//2):\n            start = int(s[2*i]) - 1\n            length = int(s[2*i+1])\n            img[start:start+length] = 1 + m\n    return img.reshape(shape).T\n\ndef mask2enc(mask, n=1):\n    pixels = mask.T.flatten()\n\n    encs = []\n\n    for i in range(1,n+1):\n        p = (pixels == i).astype(np.int8)\n        if p.sum() == 0: encs.append(np.nan)\n        else:\n            p = np.concatenate([[0], p, [0]])\n            runs = np.where(p[1:] != p[:-1])[0] + 1\n            runs[1::2] -= runs[::2]\n            encs.append(' '.join(str(x) for x in runs))\n    return encs\n\n#https://www.kaggle.com/bguberfain/memory-aware-rle-encoding\n#with transposed mask\ndef rle_encode_less_memory(img):\n    #the image should be transposed\n    pixels = img.T.flatten()\n    \n    # This simplified method requires first and last pixel to be zero\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.026676,"end_time":"2021-03-12T06:33:18.043775","exception":false,"start_time":"2021-03-12T06:33:18.017099","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:47.725099Z","iopub.execute_input":"2022-09-23T01:50:47.726005Z","iopub.status.idle":"2022-09-23T01:50:47.741019Z","shell.execute_reply.started":"2022-09-23T01:50:47.725960Z","shell.execute_reply":"2022-09-23T01:50:47.739929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# creating tiles notebook \nmean = np.array([0.62140427, 0.69426792, 0.61241501])\nstd = np.array([0.26450964, 0.29292092, 0.28286994])\n\nidentity = rasterio.Affine(1, 0, 0, 0, 1, 0)\n\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\nclass HuBMAPDataset(Dataset):\n    def __init__(self, idx, sz=sz, reduce=reduce):\n        self.data = rasterio.open(os.path.join(DATA,idx+'.tiff'), transform = identity,\n                                 num_threads='all_cpus')\n        # some images have issues with their format \n        # and must be saved correctly before reading with rasterio\n        if self.data.count != 3:\n            subdatasets = self.data.subdatasets\n            self.layers = []\n            if len(subdatasets) > 0:\n                for i, subdataset in enumerate(subdatasets, 0):\n                    self.layers.append(rasterio.open(subdataset))\n        self.shape = self.data.shape\n        if self.shape[0] < sz:\n            self.ratio = 1\n        else:\n            self.ratio = math.ceil(self.shape[0] / sz)\n        #self.reduce = reduce\n        self.reduce = self.ratio\n        self.sz = self.reduce*sz\n        self.pad0 = (self.sz - self.shape[0]%self.sz)%self.sz\n        self.pad1 = (self.sz - self.shape[1]%self.sz)%self.sz\n        self.n0max = (self.shape[0] + self.pad0)//self.sz # //: floor division: rounded to the next smallest whole number\n        self.n1max = (self.shape[1] + self.pad1)//self.sz\n      \n        \n    def __len__(self):\n        return self.n0max*self.n1max\n    \n    def __getitem__(self, idx):\n  \n        n0,n1 = idx//self.n1max, idx%self.n1max\n        # x0,y0 - are the coordinates of the lower left corner of the tile in the image\n        # negative numbers correspond to padding (which must not be loaded)\n        x0,y0 = -self.pad0//2 + n0*self.sz, -self.pad1//2 + n1*self.sz\n        # make sure that the region to read is within the image\n        p00,p01 = max(0,x0), min(x0+self.sz,self.shape[0])\n        p10,p11 = max(0,y0), min(y0+self.sz,self.shape[1])\n        #print(f'idx: {idx} n0: {n0} n1: {n1} x0: {x0} y0: {y0} pad0: {self.pad0} pad1: {self.pad1} n0max: {self.n0max} p00: {p00} p01: {p01} p10: {p10} p11: {p11}' )\n        img = np.zeros((self.sz,self.sz,3),np.uint8)\n        # mapping the loade region to the tile\n        if self.data.count == 3:\n            # move axes of an array to new positions\n            # windowed reading and writing: a view onto a rectangular subset of a raster dataset\n            # described in rasterio by column and row offsets, and width and height in pixels\n            # Window(col_off, row_off, width, height)\n            # Window.from_slices((row_start, row_stop), (col_start, col_stop))\n            img[(p00-x0):(p01-x0),(p10-y0):(p11-y0)] = np.moveaxis(self.data.read([1,2,3],\n                window=Window.from_slices((p00,p01),(p10,p11))), 0, -1)\n        else:\n            for i,layer in enumerate(self.layers):\n                img[(p00-x0):(p01-x0),(p10-y0):(p11-y0),i] =\\\n                  layer.read(1,window=Window.from_slices((p00,p01),(p10,p11)))\n        \n        if self.reduce != 1:\n            img = cv2.resize(img,(self.sz//self.reduce,self.sz//self.reduce),\n                             interpolation = cv2.INTER_AREA)\n\n        return img2tensor((img/255.0 - mean)/std), idx","metadata":{"papermill":{"duration":0.037945,"end_time":"2021-03-12T06:33:18.090462","exception":false,"start_time":"2021-03-12T06:33:18.052517","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:50.022787Z","iopub.execute_input":"2022-09-23T01:50:50.023185Z","iopub.status.idle":"2022-09-23T01:50:50.041655Z","shell.execute_reply.started":"2022-09-23T01:50:50.023149Z","shell.execute_reply":"2022-09-23T01:50:50.040625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <div style=\"padding:20px;color:white;margin:0;font-size:100%;text-align:left;display:fill;border-radius:5px;background-color:#735d78;overflow:hidden\">4. Model prediction</div>","metadata":{}},{"cell_type":"code","source":"#iterator like wrapper that returns predicted masks\n\nclass Model_pred:\n    def __init__(self, models, dl, reduce, tta:bool=True, half:bool=False):\n        self.models = models\n        self.dl = dl\n        self.tta = tta\n        self.half = half\n        self.reduce = reduce\n        \n    def __iter__(self):\n        count=0\n        with torch.no_grad():\n            for x,y in iter(self.dl): #iterate through dataset, x: img, y:idx\n                if ((y>=0).sum() > 0): #exclude empty images\n                    x = x[y>=0].to(device) #images\n                    y = y[y>=0] #idx\n                    if self.half: x = x.half() # convert to half precision, convert model to fp16\n                    py = None\n                    for model in self.models: #iterate through all models\n                        p = model(x)\n                        p = torch.sigmoid(p).detach()\n                        if py is None: py = p\n                        else: py += p #accumulate prediction\n                    if self.tta: #test time augmentation\n                        #x,y,xy flips as TTA\n                        flips = [[-1],[-2],[-2,-1],[-1,-2]]\n                       \n                        for f in flips:\n                            xf = torch.flip(x,f)\n                            for model in self.models:\n                                p = model(xf)\n                                p = torch.flip(p,f)\n                                \n                                py += torch.sigmoid(p).detach()\n                        py /= (1+len(flips))     \n                    py /= len(self.models) #take the average prediction from all models\n\n                    py = F.upsample(py, scale_factor=self.reduce, mode=\"bilinear\")\n                    py = py.permute(0,2,3,1).float().cpu()\n                    \n                    batch_size = len(py)\n                    for i in range(batch_size):\n                        yield py[i],y[i]\n                        count += 1\n                    \n    def __len__(self):\n        return len(self.dl.dataset)","metadata":{"papermill":{"duration":0.030274,"end_time":"2021-03-12T06:33:18.129546","exception":false,"start_time":"2021-03-12T06:33:18.099272","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:51.935812Z","iopub.execute_input":"2022-09-23T01:50:51.936192Z","iopub.status.idle":"2022-09-23T01:50:51.949426Z","shell.execute_reply.started":"2022-09-23T01:50:51.936159Z","shell.execute_reply":"2022-09-23T01:50:51.948258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <div style=\"padding:20px;color:white;margin:0;font-size:100%;text-align:left;display:fill;border-radius:5px;background-color:#735d78;overflow:hidden\">5. Model architecture</div>","metadata":{"papermill":{"duration":0.008902,"end_time":"2021-03-12T06:33:18.153045","exception":false,"start_time":"2021-03-12T06:33:18.144143","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class FPN(nn.Module):\n    def __init__(self, input_channels:list, output_channels:list):\n        super().__init__()\n        self.convs = nn.ModuleList(\n            [nn.Sequential(nn.Conv2d(in_ch, out_ch*2, kernel_size=3, padding=1),\n             nn.ReLU(inplace=True), nn.BatchNorm2d(out_ch*2),\n             nn.Conv2d(out_ch*2, out_ch, kernel_size=3, padding=1))\n            for in_ch, out_ch in zip(input_channels, output_channels)])\n        \n    def forward(self, xs:list, last_layer):\n        hcs = [F.interpolate(c(x),scale_factor=2**(len(self.convs)-i),mode='bilinear') \n               for i,(c,x) in enumerate(zip(self.convs, xs))]\n        hcs.append(last_layer)\n        return torch.cat(hcs, dim=1)\n\nclass UnetBlock(Module):\n    def __init__(self, up_in_c:int, x_in_c:int, nf:int=None, blur:bool=False,\n                 self_attention:bool=False, **kwargs):\n        super().__init__()\n        self.shuf = PixelShuffle_ICNR(up_in_c, up_in_c//2, blur=blur, **kwargs)\n        self.bn = nn.BatchNorm2d(x_in_c)\n        ni = up_in_c//2 + x_in_c\n        nf = nf if nf is not None else max(up_in_c//2,32)\n        self.conv1 = ConvLayer(ni, nf, norm_type=None, **kwargs)\n        self.conv2 = ConvLayer(nf, nf, norm_type=None,\n            xtra=SelfAttention(nf) if self_attention else None, **kwargs)\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, up_in:Tensor, left_in:Tensor) -> Tensor:\n        s = left_in\n        up_out = self.shuf(up_in)\n        cat_x = self.relu(torch.cat([up_out, self.bn(s)], dim=1))\n        return self.conv2(self.conv1(cat_x))\n        \nclass _ASPPModule(nn.Module):\n    def __init__(self, inplanes, planes, kernel_size, padding, dilation, groups=1):\n        super().__init__()\n        self.atrous_conv = nn.Conv2d(inplanes, planes, kernel_size=kernel_size,\n                stride=1, padding=padding, dilation=dilation, bias=False, groups=groups)\n        self.bn = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU()\n\n        self._init_weight()\n\n    def forward(self, x):\n        x = self.atrous_conv(x)\n        x = self.bn(x)\n\n        return self.relu(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\nclass ASPP(nn.Module):\n    def __init__(self, inplanes=512, mid_c=256, dilations=[6, 12, 18, 24], out_c=None):\n        super().__init__()\n        self.aspps = [_ASPPModule(inplanes, mid_c, 1, padding=0, dilation=1)] + \\\n            [_ASPPModule(inplanes, mid_c, 3, padding=d, dilation=d,groups=4) for d in dilations]\n        self.aspps = nn.ModuleList(self.aspps)\n        self.global_pool = nn.Sequential(nn.AdaptiveMaxPool2d((1, 1)),\n                        nn.Conv2d(inplanes, mid_c, 1, stride=1, bias=False),\n                        nn.BatchNorm2d(mid_c), nn.ReLU())\n        out_c = out_c if out_c is not None else mid_c\n        self.out_conv = nn.Sequential(nn.Conv2d(mid_c*(2+len(dilations)), out_c, 1, bias=False),\n                                    nn.BatchNorm2d(out_c), nn.ReLU(inplace=True))\n        self.conv1 = nn.Conv2d(mid_c*(2+len(dilations)), out_c, 1, bias=False)\n        self._init_weight()\n\n    def forward(self, x):\n        x0 = self.global_pool(x)\n        xs = [aspp(x) for aspp in self.aspps]\n        x0 = F.interpolate(x0, size=xs[0].size()[2:], mode='bilinear', align_corners=True)\n        x = torch.cat([x0] + xs, dim=1)\n        return self.out_conv(x)\n    \n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.04647,"end_time":"2021-03-12T06:33:18.208385","exception":false,"start_time":"2021-03-12T06:33:18.161915","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:54.038012Z","iopub.execute_input":"2022-09-23T01:50:54.038879Z","iopub.status.idle":"2022-09-23T01:50:54.062450Z","shell.execute_reply.started":"2022-09-23T01:50:54.038845Z","shell.execute_reply":"2022-09-23T01:50:54.061325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.models.resnet import ResNet, Bottleneck\nclass UneXt50(nn.Module):\n    def __init__(self, stride=1, **kwargs):\n        super().__init__()\n        #encoder\n        m = ResNet(Bottleneck, [3, 4, 6, 3], groups=32, width_per_group=4) \n        \n        self.enc0 = nn.Sequential(m.conv1, m.bn1, nn.ReLU(inplace=True))\n        self.enc1 = nn.Sequential(nn.MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1),\n                            m.layer1) #256\n        self.enc2 = m.layer2 #512\n        self.enc3 = m.layer3 #1024\n        self.enc4 = m.layer4 #2048\n        #aspp with customized dilatations\n        self.aspp = ASPP(2048,256,out_c=512,dilations=[stride*1,stride*2,stride*3,stride*4])\n        self.drop_aspp = nn.Dropout2d(0.5)\n        #decoder\n        self.dec4 = UnetBlock(512,1024,256)\n        self.dec3 = UnetBlock(256,512,128)\n        self.dec2 = UnetBlock(128,256,64)\n        self.dec1 = UnetBlock(64,64,32)\n        self.fpn = FPN([512,256,128,64],[16]*4)\n        self.drop = nn.Dropout2d(0.1)\n        self.final_conv = ConvLayer(32+16*4, 1, ks=1, norm_type=None, act_cls=None)\n        \n    def forward(self, x):\n        enc0 = self.enc0(x)\n        enc1 = self.enc1(enc0)\n        enc2 = self.enc2(enc1)\n        enc3 = self.enc3(enc2)\n        enc4 = self.enc4(enc3)\n        enc5 = self.aspp(enc4)\n        dec3 = self.dec4(self.drop_aspp(enc5),enc3)\n        dec2 = self.dec3(dec3,enc2)\n        dec1 = self.dec2(dec2,enc1)\n        dec0 = self.dec1(dec1,enc0)\n        x = self.fpn([enc5, dec3, dec2, dec1], dec0)\n        x = self.final_conv(self.drop(x))\n        x = F.interpolate(x,scale_factor=2,mode='bilinear')\n        return x","metadata":{"papermill":{"duration":0.028491,"end_time":"2021-03-12T06:33:18.245935","exception":false,"start_time":"2021-03-12T06:33:18.217444","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:58.148004Z","iopub.execute_input":"2022-09-23T01:50:58.148717Z","iopub.status.idle":"2022-09-23T01:50:58.160991Z","shell.execute_reply.started":"2022-09-23T01:50:58.148676Z","shell.execute_reply":"2022-09-23T01:50:58.159637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor path in MODELS:\n    state_dict = torch.load(path,map_location=torch.device('cpu'))\n    model = UneXt50()\n    model.load_state_dict(state_dict)\n    model.float()\n    model.eval()\n    model.to(device)\n    models.append(model)\n\ndel state_dict","metadata":{"papermill":{"duration":13.838863,"end_time":"2021-03-12T06:33:32.093854","exception":false,"start_time":"2021-03-12T06:33:18.254991","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:50:58.813123Z","iopub.execute_input":"2022-09-23T01:50:58.813714Z","iopub.status.idle":"2022-09-23T01:51:21.995809Z","shell.execute_reply.started":"2022-09-23T01:50:58.813679Z","shell.execute_reply":"2022-09-23T01:51:21.994465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <div style=\"padding:20px;color:white;margin:0;font-size:100%;text-align:left;display:fill;border-radius:5px;background-color:#735d78;overflow:hidden\">6. Submission</div>","metadata":{"papermill":{"duration":0.009141,"end_time":"2021-03-12T06:33:32.112738","exception":false,"start_time":"2021-03-12T06:33:32.103597","status":"completed"},"tags":[]}},{"cell_type":"code","source":"data_source = ['Hubmap']\norgans = ['kidney', 'prostate', 'largeintestine', 'spleen'] # without lung","metadata":{"execution":{"iopub.status.busy":"2022-09-23T01:51:22.001731Z","iopub.execute_input":"2022-09-23T01:51:22.002701Z","iopub.status.idle":"2022-09-23T01:51:22.008315Z","shell.execute_reply.started":"2022-09-23T01:51:22.002663Z","shell.execute_reply":"2022-09-23T01:51:22.006950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"organ_thr_hubmap = {\n                    \"kidney\" : 0.425,\n                    \"prostate\":0.425,\n                    \"largeintestine\":0.425,\n                    \"spleen\":0.425,\n                    \"lung\":0.250,\n                    }\n","metadata":{"execution":{"iopub.status.busy":"2022-09-23T01:51:22.013165Z","iopub.execute_input":"2022-09-23T01:51:22.013560Z","iopub.status.idle":"2022-09-23T01:51:22.023050Z","shell.execute_reply.started":"2022-09-23T01:51:22.013526Z","shell.execute_reply":"2022-09-23T01:51:22.022044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"other_organs_hpa = 0.525","metadata":{"execution":{"iopub.status.busy":"2022-09-23T01:51:22.026409Z","iopub.execute_input":"2022-09-23T01:51:22.028412Z","iopub.status.idle":"2022-09-23T01:51:22.040862Z","shell.execute_reply.started":"2022-09-23T01:51:22.028369Z","shell.execute_reply":"2022-09-23T01:51:22.039570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"names,preds = [],[]\nfor idx,row in tqdm(test_df.iterrows(),total=len(test_df)):\n    # set different threshold for different organs\n    threshold = 0.345\n    if row['data_source'] == \"HPA\":\n        threshold = other_organs_hpa\n    elif row['data_source'] == 'Hubmap':\n        threshold = organ_thr_hubmap[row['organ']]\n    \n#     print(threshold)\n    \n    idx = str(idx)\n    ds = HuBMAPDataset(idx)\n\n    #rasterio cannot be used with multiple workers\n\n    dl = DataLoader(ds,bs,num_workers=0,shuffle=False,pin_memory=True)\n    mp = Model_pred(models,dl, ds.reduce)\n    \n    #generate masks\n    mask = torch.zeros(len(ds),ds.sz,ds.sz,dtype=torch.int8)\n    \n    for p,i in iter(mp):        \n        mask[i.item()] = p.squeeze(-1) > threshold\n        \n\n    #reshape tiled masks into a single mask and crop padding\n    mask = mask.view(ds.n0max,ds.n1max,ds.sz,ds.sz).\\\n        permute(0,2,1,3).reshape(ds.n0max*ds.sz,ds.n1max*ds.sz)\n    mask = mask[ds.pad0//2:-(ds.pad0-ds.pad0//2) if ds.pad0 > 0 else ds.n0max*ds.sz,\n        ds.pad1//2:-(ds.pad1-ds.pad1//2) if ds.pad1 > 0 else ds.n1max*ds.sz]\n\n    #convert to rle\n    #https://www.kaggle.com/bguberfain/memory-aware-rle-encoding\n    rle = rle_encode_less_memory(mask.numpy())\n    names.append(idx)\n    preds.append(rle)\n    del mask, ds, dl\n\n            \n    gc.collect()","metadata":{"_kg_hide-output":true,"papermill":{"duration":638.710533,"end_time":"2021-03-12T06:44:10.832427","exception":false,"start_time":"2021-03-12T06:33:32.121894","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:51:22.042372Z","iopub.execute_input":"2022-09-23T01:51:22.042760Z","iopub.status.idle":"2022-09-23T01:51:30.050505Z","shell.execute_reply.started":"2022-09-23T01:51:22.042697Z","shell.execute_reply":"2022-09-23T01:51:30.049539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame({'id':names,'rle':preds})\ndf.to_csv('submission.csv',index=False)","metadata":{"papermill":{"duration":0.419953,"end_time":"2021-03-12T06:44:11.262501","exception":false,"start_time":"2021-03-12T06:44:10.842548","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T01:51:30.053740Z","iopub.execute_input":"2022-09-23T01:51:30.056726Z","iopub.status.idle":"2022-09-23T01:51:30.068327Z","shell.execute_reply.started":"2022-09-23T01:51:30.056687Z","shell.execute_reply":"2022-09-23T01:51:30.067225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}