{"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":"# **Training Notebook**","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/code/vexxingbanana/hubmap-unet-semantic-approach-train","metadata":{}},{"cell_type":"markdown","source":"# **Install segmentation_models_pytorch**","metadata":{}},{"cell_type":"code","source":"!cp -r ../input/pytorch-segmentation-models-lib/ ./","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:03:57.010620Z","iopub.execute_input":"2022-08-30T14:03:57.011210Z","iopub.status.idle":"2022-08-30T14:03:58.255049Z","shell.execute_reply.started":"2022-08-30T14:03:57.011115Z","shell.execute_reply":"2022-08-30T14:03:58.253246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip config set global.disable-pip-version-check true","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:03:58.262530Z","iopub.execute_input":"2022-08-30T14:03:58.265210Z","iopub.status.idle":"2022-08-30T14:04:00.185711Z","shell.execute_reply.started":"2022-08-30T14:03:58.265165Z","shell.execute_reply":"2022-08-30T14:04:00.184358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q ./pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ./pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ./pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl\n!pip install -q ./pytorch-segmentation-models-lib/segmentation_models_pytorch-0.2.0-py3-none-any.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-30T14:04:00.188909Z","iopub.execute_input":"2022-08-30T14:04:00.189520Z","iopub.status.idle":"2022-08-30T14:04:41.895853Z","shell.execute_reply.started":"2022-08-30T14:04:00.189476Z","shell.execute_reply":"2022-08-30T14:04:41.894622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Import Libraries**","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport time\nimport matplotlib.pyplot as plt\nimport cv2\nimport glob\nimport os\nimport shutil\nimport timm\nimport random\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.cuda import amp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport transformers\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\nimport multiprocessing as mp\nimport segmentation_models_pytorch as smp\nimport copy\nfrom collections import defaultdict\nimport gc\nfrom tqdm import tqdm\nimport tifffile\nfrom colorama import Fore, Back, Style\nfrom sklearn.utils import shuffle  ","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:04:41.899950Z","iopub.execute_input":"2022-08-30T14:04:41.900301Z","iopub.status.idle":"2022-08-30T14:04:51.931929Z","shell.execute_reply.started":"2022-08-30T14:04:41.900246Z","shell.execute_reply":"2022-08-30T14:04:51.930935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Config**","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed = 0\n    batch_size = 16\n    head = \"UNet\"\n    backbone = \"efficientnet-b3\"\n    img_size = [512, 512]\n    lr = 1e-3\n    scheduler = 'CosineAnnealingLR' #['CosineAnnealingLR']\n    epochs = 20\n    warmup_epochs = 2\n    n_folds = 5\n    folds_to_run = [0]\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    base_path = '../input/hubmap-organ-segmentation'\n    num_workers = mp.cpu_count()\n    num_classes = 1\n    n_accumulate = max(1, 16//batch_size)\n    loss = 'Dice'\n    optimizer = 'Adam'\n    weight_decay = 1e-6\n    ckpt_path = '../input/hubmap-my-dataset-256/best_epoch-00.bin' #Checkpoint path\n    ckpt_paths = ['../input/other-unet/best_epoch-00.bin', '../input/other-unet/best_epoch-01.bin', '../input/other-unet/best_epoch-02.bin', '../input/other-unet/best_epoch-03.bin', '../input/other-unet/best_epoch-04.bin']\n#     ckpt_paths = ['../input/hubmap-my-dataset-256/best_epoch-00.bin', '../input/hubmap-my-dataset-256/best_epoch-01.bin', '../input/hubmap-my-dataset-256/best_epoch-02.bin', '../input/hubmap-my-dataset-256/best_epoch-03.bin', '../input/hubmap-my-dataset-256/best_epoch-04.bin']\n    \n    threshold = 0.5\n    img_total = 5","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:04:51.933599Z","iopub.execute_input":"2022-08-30T14:04:51.934253Z","iopub.status.idle":"2022-08-30T14:04:52.001908Z","shell.execute_reply.started":"2022-08-30T14:04:51.934214Z","shell.execute_reply":"2022-08-30T14:04:51.999275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Helper Functions**","metadata":{}},{"cell_type":"code","source":"# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)  # Needed to align to RLE direction\n\n\n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n#ref: https://www.kaggle.com/code/bguberfain/memory-aware-rle-encoding/notebook\ndef rle_encode_less_memory(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    This simplified method requires first and last pixel to be zero\n    '''\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-08-30T14:04:52.004178Z","iopub.execute_input":"2022-08-30T14:04:52.005290Z","iopub.status.idle":"2022-08-30T14:04:52.017306Z","shell.execute_reply.started":"2022-08-30T14:04:52.005231Z","shell.execute_reply":"2022-08-30T14:04:52.016163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_tiff(path, scale=None, verbose=0): #Modified from https://www.kaggle.com/code/abhinand05/hubmap-extensive-eda-what-are-we-hacking\n    image = tifffile.imread(path)\n    if len(image.shape) == 5:\n        image = image.squeeze().transpose(1, 2, 0)\n    \n    if verbose:\n        print(f\"[{path}] Image shape: {image.shape}\")\n    \n    if scale:\n        new_size = (image.shape[1] // scale, image.shape[0] // scale)\n        image = cv2.resize(image, new_size)\n        \n        if verbose:\n            print(f\"[{path}] Resized Image shape: {image.shape}\")\n        \n    mx = np.max(image)\n    image = image.astype(np.float32)\n    if mx:\n        image /= mx # scale image to [0, 1]\n    return image","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:04:52.018971Z","iopub.execute_input":"2022-08-30T14:04:52.019300Z","iopub.status.idle":"2022-08-30T14:04:52.032086Z","shell.execute_reply.started":"2022-08-30T14:04:52.019249Z","shell.execute_reply":"2022-08-30T14:04:52.031047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Grab Metadata**","metadata":{}},{"cell_type":"code","source":"# 获取器官名单\nDATA_DIR = '../input/hubmap-organ-segmentation'\ndf_tr = pd.read_csv(os.path.join(DATA_DIR, 'train.csv'))\nor_index = df_tr.organ.unique()\n\n#获取.tiff文件路径\nTRAIN_IMAGES_DIR = \"/kaggle/input/hubmap-organ-segmentation/train_images\"\nall_train_images = glob.glob(os.path.join(TRAIN_IMAGES_DIR, \"*.tiff\"), recursive=True)\nall_train_images = shuffle(all_train_images)\n# print(all_train_images)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:04:52.033392Z","iopub.execute_input":"2022-08-30T14:04:52.033856Z","iopub.status.idle":"2022-08-30T14:04:52.377808Z","shell.execute_reply.started":"2022-08-30T14:04:52.033819Z","shell.execute_reply":"2022-08-30T14:04:52.376835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 加载.csv\nTR_CSV   = os.path.join(DATA_DIR, 'train.csv')\ntr_df = pd.read_csv(TR_CSV)[['id', 'rle', 'img_height', 'img_width', 'organ']] \ntr_df = shuffle(tr_df)\ntr_df['image_path'] = tr_df['id'].apply(lambda x: os.path.join(CFG.base_path, 'train_images', str(x) + '.tiff'))\n# tr_df.head()\n#print(tr_df.shape)\ntr_df.loc[tr_df.organ == 'prostate'].iloc[:CFG.img_total].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:04:52.379426Z","iopub.execute_input":"2022-08-30T14:04:52.379807Z","iopub.status.idle":"2022-08-30T14:04:52.545095Z","shell.execute_reply.started":"2022-08-30T14:04:52.379767Z","shell.execute_reply":"2022-08-30T14:04:52.544037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 创建文件夹\nbasedir_path = \"/kaggle/working/hubmap_compare\"\n    \nif not os.path.exists(basedir_path):\n    print(\"新建文件夹:\", basedir_path)\n    !mkdir /kaggle/working/hubmap_compare\n    !mkdir /kaggle/working/hubmap_compare/kidney\n    !mkdir /kaggle/working/hubmap_compare/kidney/train\n    !mkdir /kaggle/working/hubmap_compare/kidney/masks\n    !mkdir /kaggle/working/hubmap_compare/kidney/pred\n    !mkdir /kaggle/working/hubmap_compare/kidney/overlay\n    !mkdir /kaggle/working/hubmap_compare/prostate\n    !mkdir /kaggle/working/hubmap_compare/prostate/train\n    !mkdir /kaggle/working/hubmap_compare/prostate/masks\n    !mkdir /kaggle/working/hubmap_compare/prostate/pred\n    !mkdir /kaggle/working/hubmap_compare/prostate/overlay\n    !mkdir /kaggle/working/hubmap_compare/largeintestine\n    !mkdir /kaggle/working/hubmap_compare/largeintestine/train\n    !mkdir /kaggle/working/hubmap_compare/largeintestine/masks\n    !mkdir /kaggle/working/hubmap_compare/largeintestine/pred\n    !mkdir /kaggle/working/hubmap_compare/largeintestine/overlay\n    !mkdir /kaggle/working/hubmap_compare/spleen\n    !mkdir /kaggle/working/hubmap_compare/spleen/train\n    !mkdir /kaggle/working/hubmap_compare/spleen/masks\n    !mkdir /kaggle/working/hubmap_compare/spleen/pred\n    !mkdir /kaggle/working/hubmap_compare/spleen/overlay\n    !mkdir /kaggle/working/hubmap_compare/lung\n    !mkdir /kaggle/working/hubmap_compare/lung/train\n    !mkdir /kaggle/working/hubmap_compare/lung/masks\n    !mkdir /kaggle/working/hubmap_compare/lung/pred\n    !mkdir /kaggle/working/hubmap_compare/lung/overlay\nelse:\n    print(\"已经存在该文件夹:\", basedir_path)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:04:52.550237Z","iopub.execute_input":"2022-08-30T14:04:52.550541Z","iopub.status.idle":"2022-08-30T14:05:18.930680Z","shell.execute_reply.started":"2022-08-30T14:04:52.550514Z","shell.execute_reply":"2022-08-30T14:05:18.929335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"../input/hubmap-organ-segmentation/test.csv\")\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:18.932870Z","iopub.execute_input":"2022-08-30T14:05:18.933303Z","iopub.status.idle":"2022-08-30T14:05:18.956848Z","shell.execute_reply.started":"2022-08-30T14:05:18.933246Z","shell.execute_reply":"2022-08-30T14:05:18.955756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Processing**","metadata":{}},{"cell_type":"code","source":"df['image_path'] = df['id'].apply(lambda x: os.path.join(CFG.base_path, 'test_images', str(x) + '.tiff'))","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:18.958422Z","iopub.execute_input":"2022-08-30T14:05:18.959096Z","iopub.status.idle":"2022-08-30T14:05:18.966040Z","shell.execute_reply.started":"2022-08-30T14:05:18.959059Z","shell.execute_reply":"2022-08-30T14:05:18.965144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataset**","metadata":{}},{"cell_type":"code","source":"class HuBMAP_Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, labeled=True, transforms=None):\n        self.df = df\n        self.labeled = labeled\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.df.loc[index, 'image_path']\n        img_height = self.df.loc[index, 'img_height']\n        img_width = self.df.loc[index, 'img_width']\n        id_ = self.df.loc[index, 'id']\n        img = read_tiff(img_path)\n        \n        if self.labeled:\n            rle_mask = self.df.loc[index, 'rle']\n            mask = rle_decode(rle_mask, (img_height, img_width))\n            \n            if self.transforms:\n                data = self.transforms(image=img, mask=mask)\n                img  = data['image']\n                mask  = data['mask']\n            \n            mask = np.expand_dims(mask, axis=0)\n            img = np.transpose(img, (2, 0, 1))\n#             mask = np.transpose(mask, (2, 0, 1))\n            \n            return torch.tensor(img), torch.tensor(mask)\n        \n        else:\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n                \n            img = np.transpose(img, (2, 0, 1))\n            \n            return torch.tensor(img), img_height, img_width, id_","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:18.969236Z","iopub.execute_input":"2022-08-30T14:05:18.970100Z","iopub.status.idle":"2022-08-30T14:05:18.981743Z","shell.execute_reply.started":"2022-08-30T14:05:18.970073Z","shell.execute_reply":"2022-08-30T14:05:18.980666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Augmentations**","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"inference\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n}","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:18.983225Z","iopub.execute_input":"2022-08-30T14:05:18.983625Z","iopub.status.idle":"2022-08-30T14:05:18.992625Z","shell.execute_reply.started":"2022-08-30T14:05:18.983590Z","shell.execute_reply":"2022-08-30T14:05:18.991551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Models**","metadata":{}},{"cell_type":"code","source":"\ndef build_model():\n    model = smp.Unet(\n        encoder_name=CFG.backbone,      \n        encoder_weights=None,     \n        in_channels=3,                  \n        classes=CFG.num_classes,\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\n","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:18.995761Z","iopub.execute_input":"2022-08-30T14:05:18.996756Z","iopub.status.idle":"2022-08-30T14:05:19.003352Z","shell.execute_reply.started":"2022-08-30T14:05:18.996719Z","shell.execute_reply":"2022-08-30T14:05:19.002310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor path in CFG.ckpt_paths:\n    models.append(load_model(path))","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:19.004685Z","iopub.execute_input":"2022-08-30T14:05:19.005496Z","iopub.status.idle":"2022-08-30T14:05:28.744202Z","shell.execute_reply.started":"2022-08-30T14:05:19.005451Z","shell.execute_reply":"2022-08-30T14:05:28.743205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def deal_model(x):\n    py = None\n    for model in models:\n        p = model(x)\n        p = torch.sigmoid(p).detach()\n        if py is None: py = p\n        else: py += p\n    if 1 == 1:\n        #x,y,xy flips as TTA\n        flips = [[-1],[-2],[-2,-1]]\n        for f in flips:\n            xf = torch.flip(x,f)\n            for model in models:\n                p = model(xf)\n                p = torch.flip(p,f)\n                py += torch.sigmoid(p).detach()\n        py /= (1+len(flips))        \n    py /= len(models)\n    return py","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:28.747935Z","iopub.execute_input":"2022-08-30T14:05:28.748419Z","iopub.status.idle":"2022-08-30T14:05:28.755520Z","shell.execute_reply.started":"2022-08-30T14:05:28.748385Z","shell.execute_reply":"2022-08-30T14:05:28.754525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataloader**","metadata":{}},{"cell_type":"code","source":"def prepare_loaders():\n\n    infer_dataset = HuBMAP_Dataset(df, labeled=False, transforms=data_transforms['inference'])\n\n    infer_loader = torch.utils.data.DataLoader(infer_dataset, batch_size=CFG.batch_size,\n                              num_workers=CFG.num_workers, shuffle=False, pin_memory=True, drop_last=False)\n    \n    return infer_loader","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:28.757086Z","iopub.execute_input":"2022-08-30T14:05:28.757757Z","iopub.status.idle":"2022-08-30T14:05:28.766205Z","shell.execute_reply.started":"2022-08-30T14:05:28.757719Z","shell.execute_reply":"2022-08-30T14:05:28.765243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Inference**","metadata":{}},{"cell_type":"code","source":"# df","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:28.767916Z","iopub.execute_input":"2022-08-30T14:05:28.769119Z","iopub.status.idle":"2022-08-30T14:05:28.774654Z","shell.execute_reply.started":"2022-08-30T14:05:28.769079Z","shell.execute_reply":"2022-08-30T14:05:28.773728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# maks_df\n","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:05:28.775921Z","iopub.execute_input":"2022-08-30T14:05:28.776344Z","iopub.status.idle":"2022-08-30T14:05:28.784353Z","shell.execute_reply.started":"2022-08-30T14:05:28.776310Z","shell.execute_reply":"2022-08-30T14:05:28.783423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#存mask和png\nfor _, item in enumerate(or_index):\n    #存mask\n    mask_path = os.path.join(basedir_path, item, \"masks\")\n    train_path = os.path.join(basedir_path, item, \"train\")\n    pred_path = os.path.join(basedir_path, item, \"pred\")\n    maks_df = tr_df.loc[tr_df.organ == item].iloc[:CFG.img_total].reset_index(drop=True)  # 只放5个数\n    \n    print(item)\n    \n    infer_dataset = HuBMAP_Dataset(maks_df, labeled=False, transforms=data_transforms['inference'])\n    infer_loader = torch.utils.data.DataLoader(infer_dataset, batch_size=CFG.batch_size,\n                              num_workers=CFG.num_workers, shuffle=False, pin_memory=True, drop_last=False)    \n    \n    if os.path.exists(mask_path) and len(os.listdir(mask_path)) == 0:\n        print(\"come in ...\")\n        with tqdm(total = CFG.img_total) as pbar:\n            for index, row in maks_df.iterrows():\n                ##mask\n                #.csv转.png\n                img = (np.stack([rle_decode(row['rle'], shape=(row['img_width'], row['img_height']))]*3, axis=-1)).astype(np.float32)\n                #resize(512X512)\n                img = cv2.resize(img, (512, 512), interpolation=cv2.INTER_LINEAR)\n                cv2.imwrite(mask_path + \"/\" + str(row['id']) + '.png', img)\n                \n#                 if index == 1:\n#                     plt.figure(figsize=(16, 16))\n#                     plt.subplot(111)\n#                     plt.title(\"img_mask\")\n#                     plt.imshow(img)\n\n                ##png\n                for (i, train_img_path) in enumerate(all_train_images):\n                    idx = train_img_path[:-5].rsplit(\"/\", 1)[-1]\n                    if idx == str(row['id']):\n                        \n                        img_or = read_tiff(train_img_path).astype(np.float32)\n                        img_or = cv2.resize(img_or, (512, 512), interpolation=cv2.INTER_LINEAR)\n                        cv2.imwrite(train_path + \"/\" + idx + \".png\", img_or) \n                        \n#                         if index == 1:\n#                             plt.figure(figsize=(16, 16))\n#                             plt.subplot(111)\n#                             plt.title(\"img_origin\")\n#                             plt.imshow(img_or)\n                        break\n                pbar.update(1)\n\n        ##pred\n        for (images, heights, widths, ids) in infer_loader:\n            print('=' * 60)\n            print(len(images))\n            print(images.shape)\n            print('=' * 60)\n            images = images.to(CFG.device)\n    #         output = model(images)\n            output = deal_model(images)\n            print(\"images.shape\", images.shape, \"output.shape\", output.shape)\n            output = nn.Sigmoid()(output)\n            msks = (output.permute((0,2,3,1))>CFG.threshold).to(torch.uint8).cpu().detach().numpy()  # [NCHW] -> [NHWC]\n\n            for idx in range(msks.shape[0]):\n                height = heights[idx].item()\n                width = widths[idx].item()\n                id_ = ids[idx].item()\n                print(\"01msk.shape: \", msks[idx].shape)\n#                         msk = cv2.resize(msks[idx].squeeze(), \n#                                          dsize=(width, height), \n#                                          interpolation=cv2.INTER_NEAREST)\n\n                msk = cv2.resize(msks[idx], (512, 512), interpolation=cv2.INTER_NEAREST)\n\n                print(\"save: \", pred_path + \"/\" + str(id_) + \".png\")\n#                 if idx == 1:\n#                     plt.figure(figsize=(16, 16))\n#                     plt.subplot(111)\n#                     plt.title(\"pred\")\n#                     plt.imshow(msk)\n                cv2.imwrite(pred_path + \"/\" + str(id_) + \".png\", msk) \n                print(\"msk.shape: \", msk.shape)\n#                         rle = rle_encode_less_memory(msk)\n#                         pred_rles.append(rle)\n#                         pred_ids.append(id_)\n\n            gc.collect()\n            torch.cuda.empty_cache()\n\n                         ","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:13:03.510413Z","iopub.execute_input":"2022-08-30T14:13:03.510969Z","iopub.status.idle":"2022-08-30T14:13:03.548370Z","shell.execute_reply.started":"2022-08-30T14:13:03.510935Z","shell.execute_reply":"2022-08-30T14:13:03.547411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# infer_loader = prepare_loaders()\n# model = load_model(CFG.ckpt_path)\n\n# pred_ids = []\n# pred_rles = []\n# with torch.no_grad():\n#     for (images, heights, widths, ids) in infer_loader:\n#         images = images.to(CFG.device)\n# #         output = model(images)\n#         output = deal_model(images)\n#         print(output.shape)\n#         output = nn.Sigmoid()(output)\n#         msks = (output.permute((0,2,3,1))>CFG.threshold).to(torch.uint8).cpu().detach().numpy()  # [NCHW] -> [NHWC]\n\n#         for idx in range(msks.shape[0]):\n#             height = heights[idx].item()\n#             width = widths[idx].item()\n#             id_ = ids[idx].item()\n#             msk = cv2.resize(msks[idx].squeeze(), \n#                              dsize=(width, height), \n#                              interpolation=cv2.INTER_NEAREST)\n#             rle = rle_encode_less_memory(msk)\n#             pred_rles.append(rle)\n#             pred_ids.append(id_)\n\n#         gc.collect()\n#         torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:06:06.940674Z","iopub.execute_input":"2022-08-30T14:06:06.941060Z","iopub.status.idle":"2022-08-30T14:06:06.947176Z","shell.execute_reply.started":"2022-08-30T14:06:06.941022Z","shell.execute_reply":"2022-08-30T14:06:06.946295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# len(pred_rles)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:06:06.948782Z","iopub.execute_input":"2022-08-30T14:06:06.949450Z","iopub.status.idle":"2022-08-30T14:06:06.960442Z","shell.execute_reply.started":"2022-08-30T14:06:06.949409Z","shell.execute_reply":"2022-08-30T14:06:06.959561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pred_df = pd.DataFrame({\n#     \"id\":pred_ids,\n#     \"rle\":pred_rles\n# })\n# pred_df.to_csv('submission.csv',index=False)\n# display(pred_df.head(5))","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:06:06.964144Z","iopub.execute_input":"2022-08-30T14:06:06.964949Z","iopub.status.idle":"2022-08-30T14:06:06.969780Z","shell.execute_reply.started":"2022-08-30T14:06:06.964922Z","shell.execute_reply":"2022-08-30T14:06:06.968807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show Effect\n","metadata":{}},{"cell_type":"code","source":"official_color = [155, 25, 245]\nmy_color = [[230, 0, 73], [11, 180, 255], [80, 233, 145], [230, 216, 0], [132, 236, 76]]","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:06:06.971237Z","iopub.execute_input":"2022-08-30T14:06:06.971846Z","iopub.status.idle":"2022-08-30T14:06:06.979702Z","shell.execute_reply.started":"2022-08-30T14:06:06.971812Z","shell.execute_reply":"2022-08-30T14:06:06.978815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for index, organ in enumerate(or_index):\n    organ_path = os.path.join(basedir_path, organ)  # 定义路径\n\n    all_images_ids = glob.glob(os.path.join(organ_path, 'masks', \"*.png\"), recursive=True)\n    for img_id in all_images_ids:\n        img_id = img_id[:-4].rsplit(\"/\", 1)[-1]\n        print(img_id)\n        print(os.path.join(organ_path, 'train', img_id + '.png'))\n        off_msk = cv2.imread(os.path.join(organ_path, 'masks', img_id + '.png'))\n        prd_msk = cv2.imread(os.path.join(organ_path, 'pred', img_id + '.png'))\n        ori_img = cv2.imread(os.path.join(organ_path, 'train', img_id + '.png'))\n        plt.figure(figsize=(14, 14))\n        print(glob.glob(os.path.join(organ_path, 'pred', \"*.png\"), recursive=True))\n        off_msk[:, :, 0] *= official_color[0]; off_msk[:, :, 1] *= official_color[1]; off_msk[:, :, 2] *= official_color[2]  # 官方mask上色\n        prd_msk[:, :, 0] *= my_color[index][0]; prd_msk[:, :, 1] *= my_color[index][1]; prd_msk[:, :, 2] *= my_color[index][2]  # 自制mask上色\n        plt.subplot(131)\n        plt.title(\"off_msk\")\n        plt.imshow(off_msk)\n        plt.subplot(132)\n        plt.title(\"prd_msk\")\n        plt.imshow(prd_msk)\n        plt.subplot(133)\n        plt.title(\"ori_img\")\n        plt.imshow(ori_img)\n        print(off_msk.shape, prd_msk.shape, ori_img.shape)\n        off_overlay = cv2.addWeighted(src1=ori_img, alpha=0.9999, src2=off_msk, beta=0.35, gamma=0)\n        prd_overlay = cv2.addWeighted(src1=ori_img, alpha=0.9999, src2=prd_msk, beta=0.35, gamma=0)\n#         plt.figure(figsize=(14, 14))\n        plt.subplot(231)\n        plt.imshow(off_overlay)\n        plt.title(\"off_overlay\")\n        plt.subplot(232)\n        plt.imshow(prd_overlay)\n        plt.title(\"prd_overlay\")\n        cv2.imwrite(os.path.join(organ_path, 'masks', img_id + '_off.png'), off_overlay)\n        cv2.imwrite(os.path.join(organ_path, 'masks', img_id + '_prd.png'), prd_overlay)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-30T14:14:03.798316Z","iopub.execute_input":"2022-08-30T14:14:03.799056Z","iopub.status.idle":"2022-08-30T14:14:03.845577Z","shell.execute_reply.started":"2022-08-30T14:14:03.799020Z","shell.execute_reply":"2022-08-30T14:14:03.844333Z"},"trusted":true},"execution_count":null,"outputs":[]}]}