{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install /kaggle/input/einops-041-wheel/einops-0.4.1-py3-none-any.whl\n\n\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nsys.path.append('../input/subweights2')\nimport timm\n\nimport os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:64\"\n\nfrom os import path, makedirs, listdir\n\nimport numpy as np\nimport random\n\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.backends import cudnn\ncudnn.benchmark = True\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\n\n\nimport pandas as pd\nfrom tqdm import tqdm\nimport timeit\nimport cv2\nt0 = timeit.default_timer()\n\nimport gc\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = '.'\ndata_dir = '../input/hubmap-organ-segmentation/'\n# models_folder = 'weights'\nmodels_folder = '../input/subweights3/'\nmodels_folder1 = '../input/subweights4/'\nmodels_folder2 = '../input/subweights2/'\n\ndf = pd.read_csv(path.join(data_dir, 'test.csv'))\n\norgans = ['prostate', 'spleen', 'lung', 'kidney', 'largeintestine']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode_less_memory(img):\n    pixels = img.T.flatten()\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    return ' '.join(str(x) for x in runs)\n\n\ndef preprocess_inputs(x):\n    x = np.asarray(x, dtype='float32')\n    x /= 127\n    x -= 1\n    return x\n\n\nclass TestDataset(Dataset):\n    def __init__(self, df, data_dir='test_images', new_size=None):\n        super().__init__()\n        self.df = df\n        self.data_dir = data_dir\n        self.new_size = new_size\n\n    def __len__(self):\n        return len(self.df)\n\n\n    def __getitem__(self, idx):\n        r = self.df.iloc[idx]\n\n        img0 = cv2.imread(path.join(self.data_dir, '{}.tiff'.format(r['id'])), cv2.IMREAD_UNCHANGED)\n\n        orig_shape = img0.shape\n\n        sample = {'id': r['id'], 'organ': r['organ'], 'data_source': r['data_source'], 'orig_h': orig_shape[0], 'orig_w': orig_shape[1]}\n\n        for i in range(len(self.new_size)):\n\n            img = cv2.resize(img0, self.new_size[i])\n\n            img = preprocess_inputs(img)\n            img = torch.from_numpy(img.transpose((2, 0, 1)).copy()).float()\n\n            sample['img{}'.format(i)] = img\n\n        return sample\n\n\nclass ConvSilu(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3):\n        super(ConvSilu, self).__init__()\n        self.layer = nn.Sequential(\n            nn.Conv2d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, padding=1),\n            nn.SiLU(inplace=True)\n        )\n    def forward(self, x):\n        return self.layer(x)\n\n\n\nfrom coat import *\n\n\nclass Timm_Unet(nn.Module):\n    def __init__(self, name='resnet34', pretrained=True, inp_size=3, otp_size=1, decoder_filters=[32, 48, 64, 96, 128], **kwargs):\n        super(Timm_Unet, self).__init__()\n\n        if name.startswith('coat'):\n            encoder = coat_lite_medium()\n\n            if pretrained:\n                checkpoint = './weights/coat_lite_medium_384x384_f9129688.pth'\n                checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)\n                state_dict = checkpoint['model']\n                encoder.load_state_dict(state_dict,strict=False)\n        \n            encoder_filters = encoder.embed_dims\n        else:\n            encoder = timm.create_model(name, features_only=True, pretrained=pretrained, in_chans=inp_size)\n\n            encoder_filters = [f['num_chs'] for f in encoder.feature_info]\n\n        decoder_filters = decoder_filters\n\n        self.conv6 = ConvSilu(encoder_filters[-1], decoder_filters[-1])\n        self.conv6_2 = ConvSilu(decoder_filters[-1] + encoder_filters[-2], decoder_filters[-1])\n        self.conv7 = ConvSilu(decoder_filters[-1], decoder_filters[-2])\n        self.conv7_2 = ConvSilu(decoder_filters[-2] + encoder_filters[-3], decoder_filters[-2])\n        self.conv8 = ConvSilu(decoder_filters[-2], decoder_filters[-3])\n        self.conv8_2 = ConvSilu(decoder_filters[-3] + encoder_filters[-4], decoder_filters[-3])\n        self.conv9 = ConvSilu(decoder_filters[-3], decoder_filters[-4])\n\n        if len(encoder_filters) == 4:\n            self.conv9_2 = None\n        else:\n            self.conv9_2 = ConvSilu(decoder_filters[-4] + encoder_filters[-5], decoder_filters[-4])\n        \n        self.conv10 = ConvSilu(decoder_filters[-4], decoder_filters[-5])\n        \n        self.res = nn.Conv2d(decoder_filters[-5], otp_size, 1, stride=1, padding=0)\n\n        self.cls =  nn.Linear(encoder_filters[-1] * 2, 5)\n        self.pix_sz =  nn.Linear(encoder_filters[-1] * 2, 1)\n\n        self._initialize_weights()\n\n        self.encoder = encoder\n\n\n    def forward(self, x):\n        batch_size, C, H, W = x.shape\n\n        if self.conv9_2 is None:\n            enc2, enc3, enc4, enc5 = self.encoder(x)\n        else:\n            enc1, enc2, enc3, enc4, enc5 = self.encoder(x)\n\n        dec6 = self.conv6(F.interpolate(enc5, scale_factor=2))\n        dec6 = self.conv6_2(torch.cat([dec6, enc4\n                ], 1))\n\n        dec7 = self.conv7(F.interpolate(dec6, scale_factor=2))\n        dec7 = self.conv7_2(torch.cat([dec7, enc3\n                ], 1))\n        \n        dec8 = self.conv8(F.interpolate(dec7, scale_factor=2))\n        dec8 = self.conv8_2(torch.cat([dec8, enc2\n                ], 1))\n\n        dec9 = self.conv9(F.interpolate(dec8, scale_factor=2))\n\n        if self.conv9_2 is not None:\n            dec9 = self.conv9_2(torch.cat([dec9, \n                    enc1\n                    ], 1))\n        \n        dec10 = self.conv10(dec9) # F.interpolate(dec9, scale_factor=2))\n\n        x1 = torch.cat([F.adaptive_avg_pool2d(enc5, output_size=1).view(batch_size, -1), \n                        F.adaptive_max_pool2d(enc5, output_size=1).view(batch_size, -1)], 1)\n\n        # x1 = F.dropout(x1, p=0.3, training=self.training)\n        organ_cls = self.cls(x1)\n        pixel_size = self.pix_sz(x1)\n\n        return self.res(dec10), organ_cls, pixel_size\n\n\n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d) or isinstance(m, nn.ConvTranspose2d) or isinstance(m, nn.Linear):\n                m.weight.data = nn.init.kaiming_normal_(m.weight.data)\n                if m.bias is not None:\n                    m.bias.data.zero_()\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_batch_size = 1\n\n# amp_autocast = suppress\namp_autocast = torch.cuda.amp.autocast\n\nhalf_size = True\n\nhubmap_only = False #True #False\n\n\norgan_threshold = {\n    'Hubmap': {\n        'kidney'        : 90,\n        'prostate'      : 100,\n        'largeintestine': 80,\n        'spleen'        : 100,\n        'lung'          : 15,\n    },\n    'HPA': {\n        'kidney'        : 127,\n        'prostate'      : 127,\n        'largeintestine': 127,\n        'spleen'        : 127,\n        'lung'          : 25,\n    },\n}\n\n\nparams = [\n    {'size': (672, 672), 'models': [\n                                    ('coat_lite_medium', 'coat_lite_medium_672_e49_{}_best', models_folder2, 1),\n                                   ],\n                         'pred_dir': 'test_pred_672', 'weight': 0.1},\n    {'size': (768, 768), 'models': [\n                                    ('tf_efficientnet_b7_ns', 'tf_efficientnet_b7_ns_768_e34_{}_best', models_folder, 1), \n                                    ('convnext_large_384_in22ft1k', 'convnext_large_384_in22ft1k_768_e37_{}_best', models_folder, 1),\n                                    ('tf_efficientnetv2_l_in21ft1k', 'tf_efficientnetv2_l_in21ft1k_768_e36_{}_best', models_folder, 1), \n                                    ('coat_lite_medium', 'coat_lite_medium_768_e40_{}_best', models_folder2, 3),\n                                   ],\n                         'pred_dir': 'test_pred_768', 'weight': 0.15},\n    {'size': (896, 896), 'models': [\n                                    ('coat_lite_medium', 'coat_lite_medium_896_e48_{}_best', models_folder2, 1),\n                                   ],\n                         'pred_dir': 'test_pred_896', 'weight': 0.1},\n    {'size': (1024, 1024), 'models': [\n                                      ('convnext_large_384_in22ft1k', 'convnext_large_384_in22ft1k_1024_e32_{}_best', models_folder2, 1), \n                                      ('tf_efficientnet_b7_ns', 'tf_efficientnet_b7_ns_1024_e33_{}_best', models_folder, 1),\n                                      ('tf_efficientnetv2_l_in21ft1k', 'tf_efficientnetv2_l_in21ft1k_1024_e38_{}_best', models_folder, 1),\n                                    ('coat_lite_medium', 'coat_lite_medium_1024_e41_{}_best', models_folder, 3),\n                                   ],\n                         'pred_dir': 'test_pred_1024', 'weight': 0.2},\n    {'size': (1248, 1248), 'models': [\n                                    ('coat_lite_medium', 'coat_lite_medium_1248_e47_{}_best', models_folder2, 1),\n                                   ],\n                         'pred_dir': 'test_pred_1248', 'weight': 0.1},\n    {'size': (1472, 1472), 'models': [\n                                    ('tf_efficientnet_b7_ns', 'tf_efficientnet_b7_ns_1472_e35_{}_best', models_folder, 1),\n                                    ('tf_efficientnetv2_l_in21ft1k', 'tf_efficientnetv2_l_in21ft1k_1472_e39_{}_best', models_folder, 1),\n                                    ('coat_lite_medium', 'coat_lite_medium_1472_e42_{}_best', models_folder2, 3),\n                                   ],\n                         'pred_dir': 'test_pred_1472', 'weight': 0.35},\n]\n\n\n\ndef predict_models(param):\n    print(param)\n\n    makedirs(param['pred_dir'], exist_ok=True)\n\n    models = []\n\n    test_data = TestDataset(df, path.join(data_dir, 'test_images'), new_size=[param['size']])\n\n    test_data_loader = DataLoader(test_data, batch_size=test_batch_size, num_workers=1, shuffle=False)\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    for model_name, checkpoint_name, checkpoint_dir, model_weight in param['models']:\n        for fold in range(5):\n            model = Timm_Unet(name=model_name, pretrained=None)\n            snap_to_load = checkpoint_name.format(fold)\n            print(\"=> loading checkpoint '{}'\".format(snap_to_load))\n            checkpoint = torch.load(path.join(checkpoint_dir, snap_to_load), map_location='cpu')\n            loaded_dict = checkpoint['state_dict']\n            sd = model.state_dict()\n            for k in model.state_dict():\n                if k in loaded_dict:\n                    sd[k] = loaded_dict[k]\n            loaded_dict = sd\n            model.load_state_dict(loaded_dict)\n            print(\"loaded checkpoint '{}' (epoch {}, best_score {})\".format(snap_to_load, \n                checkpoint['epoch'], checkpoint['best_score']))\n            model = model.eval().cuda()\n\n            models.append((model, model_weight))\n\n\n    torch.cuda.empty_cache()\n    with torch.no_grad():\n        for sample in tqdm(test_data_loader):\n            \n            ids = sample[\"id\"].cpu().numpy()\n            orig_w = sample[\"orig_w\"].cpu().numpy()\n            orig_h = sample[\"orig_h\"].cpu().numpy()\n            # pixel_size = sample[\"pixel_size\"].cpu().numpy()\n            organ = sample[\"organ\"]\n            data_source = sample[\"data_source\"]\n\n            \n            if hubmap_only and (data_source[0] != 'Hubmap'):\n                continue\n\n\n            msk_preds = []\n            for i in range(0, len(ids), 1):\n                msk_preds.append(np.zeros((orig_h[i], orig_w[i]), dtype='float32'))\n\n            cnt = 0\n\n            imgs = sample[\"img0\"].cpu().numpy()\n\n            with amp_autocast():\n                for _tta in range(3): #8\n                    _i = _tta // 2\n                    _flip = False\n                    if _tta % 2 == 1:\n                        _flip = True\n\n                    if _i == 0:\n                        inp = imgs.copy()\n                    elif _i == 1:\n                        inp = np.rot90(imgs, k=1, axes=(2,3)).copy()\n                    elif _i == 2:\n                        inp = np.rot90(imgs, k=2, axes=(2,3)).copy()\n                    elif _i == 3:\n                        inp = np.rot90(imgs, k=3, axes=(2,3)).copy()\n\n                    if _flip:\n                        inp = inp[:, :, :, ::-1].copy()\n\n                    inp = torch.from_numpy(inp).float().cuda()                   \n                    \n                    torch.cuda.empty_cache()\n                    \n                    for model, model_weight in models:\n                        out, res_cls, res_pix = model(inp)\n                        msk_pred = torch.sigmoid(out).cpu().numpy()\n                        \n                        res_cls = torch.softmax(res_cls, dim=1).cpu().numpy()\n                        res_pix = res_pix.cpu().numpy()\n                        \n                        if _flip:\n                            msk_pred = msk_pred[:, :, :, ::-1].copy()\n\n                        if _i == 1:\n                            msk_pred = np.rot90(msk_pred, k=4-1, axes=(2,3)).copy()\n                        elif _i == 2:\n                            msk_pred = np.rot90(msk_pred, k=4-2, axes=(2,3)).copy()\n                        elif _i == 3:\n                            msk_pred = np.rot90(msk_pred, k=4-3, axes=(2,3)).copy()\n\n                        cnt += model_weight\n\n                        for i in range(len(ids)):\n                            msk_preds[i] += model_weight * cv2.resize(msk_pred[i, 0].astype('float32'), (orig_w[i], orig_h[i]))\n\n                    del inp\n                    torch.cuda.empty_cache()\n\n\n            for i in range(len(ids)):\n                msk_pred = msk_preds[i] / cnt\n                msk_pred = (msk_pred * 255).astype('uint8')\n\n                print(ids[i], organ[i], res_cls[i], res_pix[i]) #pixel_size[i]\n\n                cv2.imwrite(path.join(param['pred_dir'] , '{}.png'.format(ids[i])), msk_pred, [cv2.IMWRITE_PNG_COMPRESSION, 4])\n\n    del models\n    torch.cuda.empty_cache()\n    gc.collect()\n\n\n\nfor param in params:\n    predict_models(param)\n\n\n\nres_df = []\n\nfor _, r in df.iterrows():\n    preds = []\n\n    if hubmap_only and (r['data_source'] != 'Hubmap'):\n        res_df.append({'id': r['id'], 'rle': ''})\n        continue\n\n    for param in params:\n        pred = cv2.imread(path.join(param['pred_dir'], '{}.png'.format(r['id'])), cv2.IMREAD_GRAYSCALE)\n        preds.append(pred * param['weight'])\n\n    _thr = organ_threshold[r['data_source']][r['organ']]\n\n    pred = np.asarray(preds).sum(axis=0)\n\n    res_df.append({'id': r['id'], 'rle': rle_encode_less_memory(pred > _thr)})\n\n    # cv2.imwrite(path.join('.', '{}.png'.format(r['id'])), pred.astype('uint8'), [cv2.IMWRITE_PNG_COMPRESSION, 4])\n\nres_df = pd.DataFrame(res_df)\nres_df.to_csv(\"submission.csv\", index=False)\n\nelapsed = timeit.default_timer() - t0\nprint('Time: {:.3f} min'.format(elapsed / 60))","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}