{"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":"!cd ../input/hubmap-downloads && \\\npip install -q efficientnet_pytorch-0.6.3.tar.gz pretrainedmodels-0.7.4.tar.gz timm-0.4.12-py3-none-any.whl  segmentation_models_pytorch-0.2.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:32:28.558176Z","iopub.execute_input":"2022-07-26T10:32:28.559417Z","iopub.status.idle":"2022-07-26T10:33:05.201771Z","shell.execute_reply.started":"2022-07-26T10:32:28.559286Z","shell.execute_reply":"2022-07-26T10:33:05.200470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport torch\nimport rasterio\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\nimport tifffile as tiff\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:37:40.767399Z","iopub.execute_input":"2022-07-26T10:37:40.768480Z","iopub.status.idle":"2022-07-26T10:37:40.778067Z","shell.execute_reply.started":"2022-07-26T10:37:40.768422Z","shell.execute_reply":"2022-07-26T10:37:40.777036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    data_dir = '../input/hubmap-organ-segmentation/test_images/'\n    df_path = '../input/hubmap-organ-segmentation/test.csv'\n    size = 512\n    threshold = {\n        'kidney': 0.5,\n        'largeintestine': 0.5,\n        'prostate': 0.5,\n        'spleen': 0.5,\n        'lung': 0.2\n    }\n    tta = 3\n    backbone = 'tu-resnet101d'\n    model_paths = [\n        '../input/hubmap-test-weights/unet_tu-resnet101d_e141_0.6783_fold-0_07-26_13-07.pt',\n        '../input/hubmap-test-weights/unet_tu-resnet101d_e149_0.7874_fold-1_07-26_14-29.pt',\n        '../input/hubmap-test-weights/unet_tu-resnet101d_e158_0.7478_fold-2_07-26_17-43.pt',\n        '../input/hubmap-test-weights/unet_tu-resnet101d_e140_0.7373_fold-3_07-26_18-43.pt',\n        '../input/hubmap-test-weights/unet_tu-resnet101d_e139_0.7662_fold-4_07-26_19-48.pt'\n    ]\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.326166Z","iopub.execute_input":"2022-07-26T10:38:22.327175Z","iopub.status.idle":"2022-07-26T10:38:22.335197Z","shell.execute_reply.started":"2022-07-26T10:38:22.327137Z","shell.execute_reply":"2022-07-26T10:38:22.334133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transform = A.Compose([\n    ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.410954Z","iopub.execute_input":"2022-07-26T10:38:22.411312Z","iopub.status.idle":"2022-07-26T10:38:22.416412Z","shell.execute_reply.started":"2022-07-26T10:38:22.411280Z","shell.execute_reply":"2022-07-26T10:38:22.415408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    encs = []\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":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.491259Z","iopub.execute_input":"2022-07-26T10:38:22.491640Z","iopub.status.idle":"2022-07-26T10:38:22.505572Z","shell.execute_reply.started":"2022-07-26T10:38:22.491608Z","shell.execute_reply":"2022-07-26T10:38:22.504503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HuBMAPDataset(Dataset):\n    def __init__(self, df, data_dir, transform=None):\n        self.data_dir = data_dir\n        self.ids = df['id'].values.tolist()\n        self.labels = df['organ'].values.tolist()\n        self.pixel_sizes = df['pixel_size'].values.tolist()\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, idx):\n        path = self.data_dir + str(self.ids[idx]) + '.tiff'\n        img = tiff.imread(path)\n        s = self.pixel_sizes[idx] / 0.4 * CFG.size / 3000\n        img = cv2.resize(img, dsize=None, fx=s, fy=s, interpolation=cv2.INTER_LINEAR)\n        img = img.astype(np.float32) / 255.0\n        data = self.transform(image=img)\n        label = self.labels[idx]\n        return data['image'], s, label","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.551254Z","iopub.execute_input":"2022-07-26T10:38:22.552146Z","iopub.status.idle":"2022-07-26T10:38:22.562461Z","shell.execute_reply.started":"2022-07-26T10:38:22.552109Z","shell.execute_reply":"2022-07-26T10:38:22.561446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\ndef build_model():\n    model = smp.Unet(\n        encoder_name=CFG.backbone,      # choose encoder, e.g. mobilenet_v2 or efficientnet-b7\n        encoder_weights=None,     # use `imagenet` pre-trained weights for encoder initialization\n        decoder_channels=(320, 256, 128, 64, 32),\n        in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n        classes=1,                      # model output channels (number of classes in your dataset)\n        activation=None,\n    )\n    model.to(CFG.device)\n    return model\n\ndef load_model(path):\n    model = build_model()\n    model.load_state_dict(torch.load(path))\n    model.eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.633987Z","iopub.execute_input":"2022-07-26T10:38:22.634631Z","iopub.status.idle":"2022-07-26T10:38:22.642029Z","shell.execute_reply.started":"2022-07-26T10:38:22.634596Z","shell.execute_reply":"2022-07-26T10:38:22.640612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pad(img):\n    h, w = img.size()[-2:]\n    pad_h = (h//32+1)*32 - h if h % 32 != 0 else 0\n    pad_w = (w//32+1)*32 - w if w % 32 != 0 else 0\n    img = F.pad(img, (pad_w//2, pad_w-pad_w//2, pad_h//2, pad_h-pad_h//2), 'constant', 0)\n    return img, pad_h, pad_w\n\ndef predict(models, img):\n    pred = None\n    for model in models:\n        if pred is None:\n            pred = model(img).sigmoid()\n        else:\n            pred += model(img).sigmoid()\n    pred /= len(models)\n    return pred\n\n@torch.no_grad()\ndef inference(models, loader):\n    pbar = tqdm(loader, total=len(loader), desc='Inference')\n    rles = []\n    for img, s, label in pbar:\n        img, pad_h, pad_w = pad(img)\n        img, s = img.to(CFG.device), s.item()\n        pred = predict(models, img)\n        if CFG.tta>=3:\n            for dim in [[-1], [-2], [-1, -2]]:\n                fliped_img = torch.flip(img, dim)\n                fliped_pred = predict(models, fliped_img)\n                pred += torch.flip(fliped_pred, dim)\n            pred /= 4\n        if CFG.tta>=6:\n            rot90_img = img.permute(0, 1, 3, 2)\n            pred += predict(models, rot90_img).permute(0, 1, 3, 2) / 4\n            for dim in [[-1], [-2], [-1, -2]]:\n                fliped_img = torch.flip(rot90_img, dim)\n                fliped_pred = predict(models, fliped_img)\n                pred += torch.flip(fliped_pred, dim).permute(0, 1, 3, 2) / 4\n            pred /= 2\n        if pad_h != 0:\n            pred = pred[:, :, pad_h//2:-(pad_h-pad_h//2), :]\n        if pad_w != 0:\n            pred = pred[:, :, :, pad_w//2:-(pad_w-pad_w//2)]\n        pred = F.upsample(pred, scale_factor=1/s, mode='bilinear').squeeze(1).squeeze(0)\n        pred = (pred > CFG.threshold[label[0]]).cpu().numpy()\n        rle = rle_encode_less_memory(pred)\n        rles.append(rle)\n    return rles","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.694905Z","iopub.execute_input":"2022-07-26T10:38:22.695650Z","iopub.status.idle":"2022-07-26T10:38:22.715052Z","shell.execute_reply.started":"2022-07-26T10:38:22.695614Z","shell.execute_reply":"2022-07-26T10:38:22.714039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(CFG.df_path)\ndataset = HuBMAPDataset(df, data_dir=CFG.data_dir, transform=data_transform)\nloader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=4)\nmodels = []\nfor path in CFG.model_paths:\n    model = load_model(path)\n    models.append(model)\nrles = inference(models, loader)\nids = df['id'].values.tolist()\nsub_df = pd.DataFrame({'id':ids, 'rle':rles})\nif len(sub_df)<=1:\n    row = sub_df.iloc[0]\n    img = tiff.imread(CFG.data_dir+str(row['id'])+'.tiff')\n    rle = row['rle']\n    mask = enc2mask([rle], (img.shape[0], img.shape[1]))\n    plt.figure(figsize=(15, 15))\n    plt.imshow(img)\n    plt.imshow(mask, cmap='coolwarm', alpha=0.5)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:22.736891Z","iopub.execute_input":"2022-07-26T10:38:22.738757Z","iopub.status.idle":"2022-07-26T10:38:29.095374Z","shell.execute_reply.started":"2022-07-26T10:38:22.738708Z","shell.execute_reply":"2022-07-26T10:38:29.094315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T10:38:29.097402Z","iopub.execute_input":"2022-07-26T10:38:29.098225Z","iopub.status.idle":"2022-07-26T10:38:29.106058Z","shell.execute_reply.started":"2022-07-26T10:38:29.098182Z","shell.execute_reply":"2022-07-26T10:38:29.104841Z"},"trusted":true},"execution_count":null,"outputs":[]}]}