{"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 segmentation_models_pytorch -q","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-11T01:33:36.535362Z","iopub.execute_input":"2022-07-11T01:33:36.536021Z","iopub.status.idle":"2022-07-11T01:33:46.961973Z","shell.execute_reply.started":"2022-07-11T01:33:36.535980Z","shell.execute_reply":"2022-07-11T01:33:46.960723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ['CUDA_VISIBLE_DEVICES'] = '0'\n\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nimport tifffile as tiff \nimport rasterio\nfrom rasterio.windows import Window\n","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:46.965702Z","iopub.execute_input":"2022-07-11T01:33:46.966015Z","iopub.status.idle":"2022-07-11T01:33:46.972293Z","shell.execute_reply.started":"2022-07-11T01:33:46.965985Z","shell.execute_reply":"2022-07-11T01:33:46.971218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport pdb\nimport time\nimport warnings\nimport random\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm_notebook as tqdm\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom sklearn.model_selection import KFold\nimport torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\nimport torch.backends.cudnn as cudnn\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom matplotlib import pyplot as plt\nfrom albumentations import (HorizontalFlip, VerticalFlip, ShiftScaleRotate, Normalize, Resize, Compose, GaussNoise)\nfrom albumentations.pytorch import ToTensorV2\nwarnings.filterwarnings(\"ignore\")\n\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:46.974092Z","iopub.execute_input":"2022-07-11T01:33:46.974834Z","iopub.status.idle":"2022-07-11T01:33:46.984581Z","shell.execute_reply.started":"2022-07-11T01:33:46.974795Z","shell.execute_reply":"2022-07-11T01:33:46.983611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SAMPLE_SUBMISSION  = '../input/sartorius-cell-instance-segmentation/sample_submission.csv'\nTRAIN_CSV = \"../input/hubmap-organ-segmentation/train.csv\"\n# TRAIN_PATH =  \"../input/hubmap-organ-segmentation/train_images/\"\nTRAIN_PATH =  \"../temp/images/\"\nLABELS_PATH =  \"../temp/masks/\"\n\nTEST_PATH = \"../input/hubmap-organ-segmentation/test_images\"\ndf_sample = pd.read_csv('../input/hubmap-organ-segmentation/sample_submission.csv')\n\n\n\n# (336, 336)\nIMAGE_RESIZE = (512, 512)\nLEARNING_RATE = 5e-4\nEPOCHS = 50","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:46.987502Z","iopub.execute_input":"2022-07-11T01:33:46.987997Z","iopub.status.idle":"2022-07-11T01:33:47.011280Z","shell.execute_reply.started":"2022-07-11T01:33:46.987959Z","shell.execute_reply":"2022-07-11T01:33:47.010343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(TRAIN_CSV)\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:47.012803Z","iopub.execute_input":"2022-07-11T01:33:47.013174Z","iopub.status.idle":"2022-07-11T01:33:47.355314Z","shell.execute_reply.started":"2022-07-11T01:33:47.013138Z","shell.execute_reply":"2022-07-11T01:33:47.354121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# tiled images","metadata":{}},{"cell_type":"code","source":"!mkdir -p /kaggle/temp/images\n!mkdir -p /kaggle/temp/masks","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:47.356984Z","iopub.execute_input":"2022-07-11T01:33:47.357340Z","iopub.status.idle":"2022-07-11T01:33:48.850092Z","shell.execute_reply.started":"2022-07-11T01:33:47.357304Z","shell.execute_reply":"2022-07-11T01:33:48.848782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\ndef tile_image(p_img, folder, size: int = 1024):\n    w = h = size\n    im = np.array(Image.open(p_img))\n    # https://stackoverflow.com/a/47581978/4521646\n    tiles = [im[i:(i + h), j:(j + w), ...] for i in range(0, im.shape[0], h) for j in range(0, im.shape[1], w)]\n    idxs = [(i, (i + h), j, (j + w)) for i in range(0, im.shape[0], h) for j in range(0, im.shape[1], w)]\n    name, _ = os.path.splitext(os.path.basename(p_img))\n    files = []\n    for k, tile in enumerate(tiles):\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile = np.zeros_like(tiles[0])\n            tile[:tile_.shape[0], :tile_.shape[1], ...] = tile_\n        p_img = os.path.join(folder, f\"{name}_{k:02}.png\")\n        Image.fromarray(tile).save(p_img)\n        files.append(p_img)\n    return files, idxs","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:48.852348Z","iopub.execute_input":"2022-07-11T01:33:48.852750Z","iopub.status.idle":"2022-07-11T01:33:48.864985Z","shell.execute_reply.started":"2022-07-11T01:33:48.852710Z","shell.execute_reply":"2022-07-11T01:33:48.863817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tile_image_2(im, size: int = 1024):\n    w = h = size\n    # https://stackoverflow.com/a/47581978/4521646\n    tiles = [im[i:(i + h), j:(j + w), ...] for i in range(0, im.shape[0], h) for j in range(0, im.shape[1], w)]\n    idxs = [(i, (i + h), j, (j + w)) for i in range(0, im.shape[0], h) for j in range(0, im.shape[1], w)]\n    files = []\n    for k, tile in enumerate(tiles):\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile = np.zeros_like(tiles[0])\n            tile[:tile_.shape[0], :tile_.shape[1], ...] = tile_\n        files.append(tile)\n    return files, idxs","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:33:48.868154Z","iopub.execute_input":"2022-07-11T01:33:48.868533Z","iopub.status.idle":"2022-07-11T01:33:48.880383Z","shell.execute_reply.started":"2022-07-11T01:33:48.868478Z","shell.execute_reply":"2022-07-11T01:33:48.879300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tiles_img, _ = tile_image(\"../input/hubmap-organ-segmentation/train_images/10044.tiff\", \"/kaggle/temp/images\", size=1024)\ntiles_seg, idxs = tile_image(\"../input/hacking-the-human-body-annotation-masks/train_binary_masks/10044.png\", \"/kaggle/temp/masks\", size=1024)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:09.626433Z","iopub.execute_input":"2022-07-11T01:34:09.627049Z","iopub.status.idle":"2022-07-11T01:34:14.185468Z","shell.execute_reply.started":"2022-07-11T01:34:09.627008Z","shell.execute_reply":"2022-07-11T01:34:14.184351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idxs","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:14.189818Z","iopub.execute_input":"2022-07-11T01:34:14.190321Z","iopub.status.idle":"2022-07-11T01:34:14.201263Z","shell.execute_reply.started":"2022-07-11T01:34:14.190286Z","shell.execute_reply":"2022-07-11T01:34:14.200192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls -lh /kaggle/temp/images\n!ls -lh /kaggle/temp/masks","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:36:32.854857Z","iopub.status.idle":"2022-07-11T01:36:32.855809Z","shell.execute_reply.started":"2022-07-11T01:36:32.855494Z","shell.execute_reply":"2022-07-11T01:36:32.855522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from skimage import color\n\nfig, axes = plt.subplots(nrows=3, ncols=3, figsize=(9, 9))\nfor i, (p_img, p_seg) in enumerate(zip(tiles_img, tiles_seg)):\n    img = plt.imread(p_img)\n    mask = np.array(Image.open(p_seg))\n    axes[i // 3, i % 3].imshow(color.label2rgb(mask, img, bg_label=0, bg_color=(1.,1.,1.), alpha=0.25))\n    axes[i // 3, i % 3].set_axis_off()\nfig.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:18.138392Z","iopub.execute_input":"2022-07-11T01:34:18.139120Z","iopub.status.idle":"2022-07-11T01:34:23.679505Z","shell.execute_reply.started":"2022-07-11T01:34:18.139078Z","shell.execute_reply":"2022-07-11T01:34:23.678620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def un_tile_image(tiles_seg, idxs, folder_image):\n    tiles = [np.array(Image.open(p_seg)) for p_seg in tiles_seg]\n    im = plt.imread(folder_image)\n    seg = np.zeros(im.shape[:2], dtype=np.uint8)\n    for tile, (i1, i2, j1, j2) in zip(tiles, idxs):\n        i2 = min(i2, im.shape[0])\n        j2 = min(j2, im.shape[1])\n        seg[i1:i2, j1:j2] = tile[:(i2 - i1), :(j2 - j1)]\n    return seg\n    ","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:23.681411Z","iopub.execute_input":"2022-07-11T01:34:23.682507Z","iopub.status.idle":"2022-07-11T01:34:23.692640Z","shell.execute_reply.started":"2022-07-11T01:34:23.682466Z","shell.execute_reply":"2022-07-11T01:34:23.690622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tiles_seg, idxs = tile_image(\"../input/hacking-the-human-body-annotation-masks/train_binary_masks/12233.png\", \"/kaggle/temp/masks\", size=1024)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:23.695078Z","iopub.execute_input":"2022-07-11T01:34:23.696411Z","iopub.status.idle":"2022-07-11T01:34:23.880410Z","shell.execute_reply.started":"2022-07-11T01:34:23.696371Z","shell.execute_reply":"2022-07-11T01:34:23.879417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folder= \"../input/hubmap-organ-segmentation/train_images/12233.tiff\"\ntiles_seg, idxs = tile_image(\"../input/hacking-the-human-body-annotation-masks/train_binary_masks/12233.png\", \"/kaggle/temp/masks\", size=1024)\n\nreconstruted_seg = un_tile_image(tiles_seg, idxs, folder)\n\nplt.figure(figsize=(10, 10))    \nplt.imshow(reconstruted_seg)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:23.884493Z","iopub.execute_input":"2022-07-11T01:34:23.884782Z","iopub.status.idle":"2022-07-11T01:34:25.662430Z","shell.execute_reply.started":"2022-07-11T01:34:23.884757Z","shell.execute_reply":"2022-07-11T01:34:25.661498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"original_mask = plt.imread(\"../input/hacking-the-human-body-annotation-masks/train_binary_masks/12233.png\")\nplt.figure(figsize=(10, 10))    \nplt.imshow(original_mask)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:25.663746Z","iopub.execute_input":"2022-07-11T01:34:25.664689Z","iopub.status.idle":"2022-07-11T01:34:26.659739Z","shell.execute_reply.started":"2022-07-11T01:34:25.664650Z","shell.execute_reply":"2022-07-11T01:34:26.658631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATASET_FOLDER = \"/kaggle/input/hubmap-organ-segmentation\"\nANNOT_DATASET = \"/kaggle/input/hacking-the-human-body-annotation-masks\"","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:34:26.661893Z","iopub.execute_input":"2022-07-11T01:34:26.662375Z","iopub.status.idle":"2022-07-11T01:34:26.667935Z","shell.execute_reply.started":"2022-07-11T01:34:26.662334Z","shell.execute_reply":"2022-07-11T01:34:26.666954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from joblib import Parallel, delayed\nimport os, glob\n\n\nTILE_SIZE = 1024\n\nfor dir_source, dir_target in [(os.path.join(DATASET_FOLDER, 'train_images'), \"/kaggle/temp/images\")]:\n    ls = glob.glob(os.path.join(dir_source, '*'))\n    _= Parallel(n_jobs=3)(delayed(tile_image)(p_img, dir_target, size=TILE_SIZE) for p_img in tqdm(ls))","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:36:38.941478Z","iopub.execute_input":"2022-07-11T01:36:38.942135Z","iopub.status.idle":"2022-07-11T01:36:38.949094Z","shell.execute_reply.started":"2022-07-11T01:36:38.942096Z","shell.execute_reply":"2022-07-11T01:36:38.947990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for dir_source, dir_target in [(os.path.join(ANNOT_DATASET, 'train_binary_masks'), \"/kaggle/temp/masks\")]:\n    ls = glob.glob(os.path.join(dir_source, '*'))\n    _= Parallel(n_jobs=3)(delayed(tile_image)(p_img, dir_target, size=TILE_SIZE) for p_img in tqdm(ls))","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:36:32.788780Z","iopub.status.idle":"2022-07-11T01:36:32.789145Z","shell.execute_reply.started":"2022-07-11T01:36:32.788972Z","shell.execute_reply":"2022-07-11T01:36:32.788989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls = glob.glob(os.path.join(TRAIN_PATH, '*'))","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:36:50.864002Z","iopub.execute_input":"2022-07-11T01:36:50.864355Z","iopub.status.idle":"2022-07-11T01:36:50.882472Z","shell.execute_reply.started":"2022-07-11T01:36:50.864325Z","shell.execute_reply":"2022-07-11T01:36:50.881566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_imagen = []\nfor image in ls:\n    list_imagen.append(image[15:-4])\n    \ndf_list = pd.DataFrame(list_imagen)\ndf_list.columns = [\"id\"]\ndf_list","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:36:57.288404Z","iopub.execute_input":"2022-07-11T01:36:57.289110Z","iopub.status.idle":"2022-07-11T01:36:57.306196Z","shell.execute_reply.started":"2022-07-11T01:36:57.289072Z","shell.execute_reply":"2022-07-11T01:36:57.304923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf = KFold(n_splits=5, shuffle=True, random_state=42, )\nfor fold, (train_idx, val_idx) in enumerate(kf.split(df_list)):\n    df_list.loc[val_idx, 'fold'] = fold\ndf_list.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:09.253646Z","iopub.execute_input":"2022-07-11T01:37:09.254648Z","iopub.status.idle":"2022-07-11T01:37:09.278003Z","shell.execute_reply.started":"2022-07-11T01:37:09.254599Z","shell.execute_reply":"2022-07-11T01:37:09.276843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_list.fold.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:12.748367Z","iopub.execute_input":"2022-07-11T01:37:12.748748Z","iopub.status.idle":"2022-07-11T01:37:12.761826Z","shell.execute_reply.started":"2022-07-11T01:37:12.748717Z","shell.execute_reply":"2022-07-11T01:37:12.760621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rleToMask(rleString,height,width):\n  rows,cols = height,width\n  rleNumbers = [int(numstring) for numstring in rleString.split(' ')]\n  rlePairs = np.array(rleNumbers).reshape(-1,2)\n  img = np.zeros(rows*cols,dtype=np.uint8)\n  for index,length in rlePairs:\n    index -= 1\n    img[index:index+length] = 255\n  img = img.reshape(cols,rows)\n  img = img.T\n  return img","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:13.566202Z","iopub.execute_input":"2022-07-11T01:37:13.566598Z","iopub.status.idle":"2022-07-11T01:37:13.573968Z","shell.execute_reply.started":"2022-07-11T01:37:13.566546Z","shell.execute_reply":"2022-07-11T01:37:13.572835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# helper function for data visualization\ndef visualize(**images):\n    \"\"\"PLot images in one row.\"\"\"\n    n = len(images)\n    plt.figure(figsize=(10, 10))\n    for i, (name, image) in enumerate(images.items()):\n        plt.subplot(1, n, i + 1)\n        plt.xticks([])\n        plt.yticks([])\n        plt.title(' '.join(name.split('_')).title())\n        plt.imshow(image)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:14.747692Z","iopub.execute_input":"2022-07-11T01:37:14.748859Z","iopub.status.idle":"2022-07-11T01:37:14.756753Z","shell.execute_reply.started":"2022-07-11T01:37:14.748810Z","shell.execute_reply":"2022-07-11T01:37:14.755699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\n\nclass CellDatasetTiles(Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.base_path = TRAIN_PATH\n        self.labels_path = LABELS_PATH\n\n        self.transforms = Compose([Resize(IMAGE_RESIZE[0], IMAGE_RESIZE[1]), \n#                                    Normalize(mean=RESNET_MEAN, std=RESNET_STD, p=1), \n                                   HorizontalFlip(p=0.5),\n                                   VerticalFlip(p=0.5),\n                                   ToTensorV2()])\n        self.gb = self.df.groupby('id')\n        self.image_ids = df.id.unique().tolist()\n\n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        df = self.gb.get_group(image_id)\n\n        image_path = os.path.join(self.base_path, str(image_id) + \".png\")\n        image =  cv2.imread(image_path).astype('double')\n        \n        label_path = os.path.join(self.labels_path, str(image_id) + \".png\")\n        mask =  cv2.imread(label_path)\n        mask = (mask >= 1).astype('double')\n        \n        augmented = self.transforms(image=image, mask=mask)\n        image = augmented['image']\n        mask = augmented['mask']\n        return (((image/255.0).permute(1,2,0)-mean)/std).permute(2,0,1).float() , mask[:,:,0].reshape((1, IMAGE_RESIZE[0], IMAGE_RESIZE[1])).float()\n#         return image, mask\n\n    def __len__(self):\n        return len(self.image_ids)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:15.343823Z","iopub.execute_input":"2022-07-11T01:37:15.344453Z","iopub.status.idle":"2022-07-11T01:37:15.356640Z","shell.execute_reply.started":"2022-07-11T01:37:15.344411Z","shell.execute_reply":"2022-07-11T01:37:15.355684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = CellDatasetTiles(df_list[df_list.fold != 0])\nds_valid = CellDatasetTiles(df_list[df_list.fold == 0])\n\nimage, mask = ds_train[15]\n\nimage.shape, mask.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:18.849969Z","iopub.execute_input":"2022-07-11T01:37:18.850341Z","iopub.status.idle":"2022-07-11T01:37:18.948626Z","shell.execute_reply.started":"2022-07-11T01:37:18.850311Z","shell.execute_reply":"2022-07-11T01:37:18.947581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize(\n            image=image.permute(2,1,0),\n            mask=mask.permute(2,1,0),\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:19.499281Z","iopub.execute_input":"2022-07-11T01:37:19.499661Z","iopub.status.idle":"2022-07-11T01:37:19.780537Z","shell.execute_reply.started":"2022-07-11T01:37:19.499630Z","shell.execute_reply":"2022-07-11T01:37:19.779595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl_train = DataLoader(ds_train, batch_size=8, num_workers=4, pin_memory=True, shuffle=False)\ndl_valid = DataLoader(ds_valid, batch_size=8, num_workers=4, pin_memory=True, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:24.539495Z","iopub.execute_input":"2022-07-11T01:37:24.540459Z","iopub.status.idle":"2022-07-11T01:37:24.548082Z","shell.execute_reply.started":"2022-07-11T01:37:24.540420Z","shell.execute_reply":"2022-07-11T01:37:24.546859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dl_train)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:24.979929Z","iopub.execute_input":"2022-07-11T01:37:24.980311Z","iopub.status.idle":"2022-07-11T01:37:24.987705Z","shell.execute_reply.started":"2022-07-11T01:37:24.980281Z","shell.execute_reply":"2022-07-11T01:37:24.986494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dl_valid)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:25.333617Z","iopub.execute_input":"2022-07-11T01:37:25.336466Z","iopub.status.idle":"2022-07-11T01:37:25.342669Z","shell.execute_reply.started":"2022-07-11T01:37:25.336419Z","shell.execute_reply":"2022-07-11T01:37:25.341492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get a batch from the dataloader\nbatch = next(iter(dl_train))\nimages, masks = batch","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:29.315017Z","iopub.execute_input":"2022-07-11T01:37:29.317467Z","iopub.status.idle":"2022-07-11T01:37:32.651508Z","shell.execute_reply.started":"2022-07-11T01:37:29.317429Z","shell.execute_reply":"2022-07-11T01:37:32.650191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx=1\nplt.figure(figsize=(10,10))\nplt.imshow(images[3][0].permute(1,0), cmap='bone')\nplt.show()\nplt.figure(figsize=(10,10))\nplt.imshow(masks[3].permute(2,1,0), alpha=0.3)\nplt.show()\nplt.figure(figsize=(10,10))\nplt.imshow(images[3][0].permute(1,0), cmap='bone')\nplt.imshow(masks[3].permute(2,1,0), alpha=0.3)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:47.489630Z","iopub.execute_input":"2022-07-11T01:37:47.490468Z","iopub.status.idle":"2022-07-11T01:37:48.465255Z","shell.execute_reply.started":"2022-07-11T01:37:47.490422Z","shell.execute_reply":"2022-07-11T01:37:48.461039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_loss(input, target):\n    input = torch.sigmoid(input)\n    smooth = 1.0\n    iflat = input.view(-1)\n    tflat = target.view(-1)\n    intersection = (iflat * tflat).sum()\n    return ((2.0 * intersection + smooth) / (iflat.sum() + tflat.sum() + smooth))\n\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma):\n        super().__init__()\n        self.gamma = gamma\n\n    def forward(self, input, target):\n        if not (target.size() == input.size()):\n            raise ValueError(\"Target size ({}) must be the same as input size ({})\"\n                             .format(target.size(), input.size()))\n        max_val = (-input).clamp(min=0)\n        loss = input - input * target + max_val + \\\n            ((-max_val).exp() + (-input - max_val).exp()).log()\n        invprobs = F.logsigmoid(-input * (target * 2.0 - 1.0))\n        loss = (invprobs * self.gamma).exp() * loss\n        return loss.mean()\n\n\nclass MixedLoss(nn.Module):\n    def __init__(self, alpha, gamma):\n        super().__init__()\n        self.alpha = alpha\n        self.focal = FocalLoss(gamma)\n\n    def forward(self, input, target):\n        loss = self.alpha*self.focal(input, target) - torch.log(dice_loss(input, target))\n        return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:53.029914Z","iopub.execute_input":"2022-07-11T01:37:53.030282Z","iopub.status.idle":"2022-07-11T01:37:53.043549Z","shell.execute_reply.started":"2022-07-11T01:37:53.030251Z","shell.execute_reply":"2022-07-11T01:37:53.042479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create model and train\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport collections.abc as container_abcs\ntorch._six.container_abcs = container_abcs\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:57.561059Z","iopub.execute_input":"2022-07-11T01:37:57.562004Z","iopub.status.idle":"2022-07-11T01:37:57.567335Z","shell.execute_reply.started":"2022-07-11T01:37:57.561968Z","shell.execute_reply":"2022-07-11T01:37:57.566237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ssl\nssl._create_default_https_context = ssl._create_unverified_context","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:37:59.812162Z","iopub.execute_input":"2022-07-11T01:37:59.813137Z","iopub.status.idle":"2022-07-11T01:37:59.817962Z","shell.execute_reply.started":"2022-07-11T01:37:59.813096Z","shell.execute_reply":"2022-07-11T01:37:59.816601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ENCODER = 'se_resnet101'\nENCODER_WEIGHTS = 'imagenet'\nCLASSES = ['cell']\nACTIVATION = 'sigmoid' # could be None for logits or 'softmax2d' for multiclass segmentation\nDEVICE = 'cuda'\n\n# create segmentation model with pretrained encoder\nmodel = smp.FPN(\n    encoder_name=ENCODER, \n    encoder_weights=ENCODER_WEIGHTS, \n    classes=len(CLASSES), \n    activation=ACTIVATION,\n)\n\npreprocessing_fn = smp.encoders.get_preprocessing_fn(ENCODER, ENCODER_WEIGHTS)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:36:32.839400Z","iopub.status.idle":"2022-07-11T01:36:32.840194Z","shell.execute_reply.started":"2022-07-11T01:36:32.839924Z","shell.execute_reply":"2022-07-11T01:36:32.839959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.load('../input/model-10-07-2022/best_model.pth', map_location ='cuda')","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:38:04.765452Z","iopub.execute_input":"2022-07-11T01:38:04.766061Z","iopub.status.idle":"2022-07-11T01:38:07.227393Z","shell.execute_reply.started":"2022-07-11T01:38:04.766011Z","shell.execute_reply":"2022-07-11T01:38:07.226353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss = smp.utils.losses.DiceLoss()\nmetrics = [\n    smp.utils.metrics.IoU(threshold=0.5),\n]\n\noptimizer = torch.optim.Adam([ \n    dict(params=model.parameters(), lr=0.0001),\n])","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:38:11.778701Z","iopub.execute_input":"2022-07-11T01:38:11.779435Z","iopub.status.idle":"2022-07-11T01:38:11.788820Z","shell.execute_reply.started":"2022-07-11T01:38:11.779393Z","shell.execute_reply":"2022-07-11T01:38:11.787434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create epoch runners \n# it is a simple loop of iterating over dataloader`s samples\ntrain_epoch = smp.utils.train.TrainEpoch(\n    model, \n    loss=loss, \n    metrics=metrics, \n    optimizer=optimizer,\n    device='cuda',\n    verbose=True,\n)\n\nvalid_epoch = smp.utils.train.ValidEpoch(\n    model, \n    loss=loss, \n    metrics=metrics, \n    device=DEVICE,\n    verbose=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:38:19.076368Z","iopub.execute_input":"2022-07-11T01:38:19.076752Z","iopub.status.idle":"2022-07-11T01:38:19.105444Z","shell.execute_reply.started":"2022-07-11T01:38:19.076719Z","shell.execute_reply":"2022-07-11T01:38:19.104512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train model for 40 epochs\n\nmax_score = 0\n\nfor i in range(0, 40):\n    \n    print('\\nEpoch: {}'.format(i))\n    train_logs = train_epoch.run(dl_train)\n    valid_logs = valid_epoch.run(dl_valid)\n    \n#     do something (save model, change lr, etc.)\n    if max_score < valid_logs['iou_score']:\n        max_score = valid_logs['iou_score']\n        torch.save(model, './best_model.pth')\n        print('Model saved!')\n    torch.save(model, './best_model.pth')\n\n    if i == 25:\n        optimizer.param_groups[0]['lr'] = 1e-5\n        print('Decrease decoder learning rate to 1e-5!')","metadata":{"execution":{"iopub.status.busy":"2022-07-10T22:59:41.830554Z","iopub.execute_input":"2022-07-10T22:59:41.831347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize predictions","metadata":{}},{"cell_type":"code","source":"# get a batch from the dataloader\nbatch = next(iter(dl_train))\nimages, masks = batch","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:38:34.666756Z","iopub.execute_input":"2022-07-11T01:38:34.667138Z","iopub.status.idle":"2022-07-11T01:38:37.332667Z","shell.execute_reply.started":"2022-07-11T01:38:34.667106Z","shell.execute_reply":"2022-07-11T01:38:37.331302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize(\n                images=images[3].permute(2,1,0),\n                mask=masks[3].permute(2,1,0),\n\n            )","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:39:22.154471Z","iopub.execute_input":"2022-07-11T01:39:22.155424Z","iopub.status.idle":"2022-07-11T01:39:22.415375Z","shell.execute_reply.started":"2022-07-11T01:39:22.155384Z","shell.execute_reply":"2022-07-11T01:39:22.414509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(8):\n    \n    \n    image = images[i]\n    gt_mask = masks[i]\n#     image_vis = image.astype('uint8')\n#     image, gt_mask = test_dataset[n]\n    \n    gt_mask = gt_mask.squeeze()\n    \n    x_tensor = image.to(DEVICE).unsqueeze(0)\n    pr_mask = model.predict(x_tensor)\n    pr_mask = (pr_mask.squeeze().cpu().numpy().round())\n        \n    visualize(\n        ground_truth_mask=gt_mask, \n        predicted_mask=pr_mask\n    )","metadata":{"execution":{"iopub.status.busy":"2022-07-11T01:39:37.970639Z","iopub.execute_input":"2022-07-11T01:39:37.971409Z","iopub.status.idle":"2022-07-11T01:39:40.233720Z","shell.execute_reply.started":"2022-07-11T01:39:37.971369Z","shell.execute_reply":"2022-07-11T01:39:40.232689Z"},"trusted":true},"execution_count":null,"outputs":[]}]}