{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# HW_7_Segmentation (baseline)","metadata":{}},{"cell_type":"markdown","source":"**Case Description**: Neurological disorders, including neurodegenerative diseases such as Alzheimer's and brain tumors, are a leading cause of death and disability across the globe. However, it is hard to quantify how well these deadly disorders respond to treatment. One accepted method is to review neuronal cells via light microscopy, which is both accessible and non-invasive. Unfortunately, segmenting individual neuronal cells in microscopic images can be challenging and time-intensive. Accurate instance segmentation of these cells—with the help of computer vision—could lead to new and effective drug discoveries to treat the millions of people with these disorders.\n\n**Objective**: detecting masks of different cell objects in phase contrast microscopy images.\n\n**Metrics**: IoU","metadata":{}},{"cell_type":"markdown","source":"There are 606 images, 73585 annotations in training set, and there are roughly 240 images in hidden test set. Average annotations per image is 121.42 in training set and same ratio is expected in the hidden test set. In addition to that, there are 1972 images without annotations in *train_semi_supervised* directory. Their metadata isn't listed in train.csv file.\n\nThere are 9 columns in image metadata file.\n\n* id - Unique ID of the image\n* annotation - Run length encoded segmentation masks\n* width - Width of the image\n* height - Height of the image\n* cell_type - Type of the cell line\n* plate_time - Plate creation time\n* sample_date - Timestamp of the sample\n* sample_id - Unique ID of the sample\n* elapsed_timedelta - Time since first image taken of sample","metadata":{}},{"cell_type":"markdown","source":"***","metadata":{}},{"cell_type":"markdown","source":"# 0. Istall & Import","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport random\nimport colorsys\nfrom tqdm.auto import tqdm\n\nimport cv2\nimport albumentations as A\n\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nimport torchvision\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.transforms import functional as F\n\nfrom sklearn.model_selection import train_test_split","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-10T18:14:32.492005Z","iopub.execute_input":"2024-05-10T18:14:32.492630Z","iopub.status.idle":"2024-05-10T18:14:39.170476Z","shell.execute_reply.started":"2024-05-10T18:14:32.492598Z","shell.execute_reply":"2024-05-10T18:14:39.169627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Data exploration","metadata":{}},{"cell_type":"code","source":"!ls /kaggle/input/sartorius-cell-instance-segmentation","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:39.172348Z","iopub.execute_input":"2024-05-10T18:14:39.173076Z","iopub.status.idle":"2024-05-10T18:14:40.152873Z","shell.execute_reply.started":"2024-05-10T18:14:39.173044Z","shell.execute_reply":"2024-05-10T18:14:40.151725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT = \"/kaggle/input/sartorius-cell-instance-segmentation\"","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:40.154619Z","iopub.execute_input":"2024-05-10T18:14:40.155551Z","iopub.status.idle":"2024-05-10T18:14:40.160089Z","shell.execute_reply.started":"2024-05-10T18:14:40.155512Z","shell.execute_reply":"2024-05-10T18:14:40.159181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(ROOT + \"/train.csv\")\ndisplay(train_df.head(2))\nprint(f'Train data shape: {train_df.shape}')","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:40.162446Z","iopub.execute_input":"2024-05-10T18:14:40.162732Z","iopub.status.idle":"2024-05-10T18:14:40.849926Z","shell.execute_reply.started":"2024-05-10T18:14:40.162701Z","shell.execute_reply":"2024-05-10T18:14:40.848932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Unique id: {len(set(train_df.id))}')","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:40.851352Z","iopub.execute_input":"2024-05-10T18:14:40.851729Z","iopub.status.idle":"2024-05-10T18:14:40.866222Z","shell.execute_reply.started":"2024-05-10T18:14:40.851695Z","shell.execute_reply":"2024-05-10T18:14:40.865305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Let's take the first image as an example\ntrain_df[\"id\"].iloc[0]","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:40.867447Z","iopub.execute_input":"2024-05-10T18:14:40.867729Z","iopub.status.idle":"2024-05-10T18:14:40.876413Z","shell.execute_reply.started":"2024-05-10T18:14:40.867706Z","shell.execute_reply":"2024-05-10T18:14:40.875451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_id = train_df.iloc[0][\"id\"]\nimg = cv2.imread(f'{ROOT}/train/{img_id}.png')","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:40.877680Z","iopub.execute_input":"2024-05-10T18:14:40.877950Z","iopub.status.idle":"2024-05-10T18:14:40.910799Z","shell.execute_reply.started":"2024-05-10T18:14:40.877927Z","shell.execute_reply":"2024-05-10T18:14:40.910054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(img)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:40.911851Z","iopub.execute_input":"2024-05-10T18:14:40.912110Z","iopub.status.idle":"2024-05-10T18:14:41.372042Z","shell.execute_reply.started":"2024-05-10T18:14:40.912088Z","shell.execute_reply":"2024-05-10T18:14:41.371077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"After image loading let's get target values.","metadata":{}},{"cell_type":"code","source":"cls_map = {value:idx for idx,value in enumerate(train_df[\"cell_type\"].unique())}\ncls_map","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.373321Z","iopub.execute_input":"2024-05-10T18:14:41.373642Z","iopub.status.idle":"2024-05-10T18:14:41.388239Z","shell.execute_reply.started":"2024-05-10T18:14:41.373617Z","shell.execute_reply":"2024-05-10T18:14:41.387162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_targets_mask(df, img_id):\n    \"\"\"\n    Function to get target masks\n    \n    rles contains mask's description\n    \"\"\"\n    targets = df[df[\"id\"] == img_id][\"cell_type\"].apply(lambda x: cls_map[x]).values \n    rles = df[df[\"id\"] == img_id][\"annotation\"].values\n    \n    return targets, rles","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.392183Z","iopub.execute_input":"2024-05-10T18:14:41.392493Z","iopub.status.idle":"2024-05-10T18:14:41.398549Z","shell.execute_reply.started":"2024-05-10T18:14:41.392469Z","shell.execute_reply":"2024-05-10T18:14:41.397554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Look at the number of labels, rles\nlabels, rles = get_targets_mask(train_df, img_id)\nlen(rles)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.400005Z","iopub.execute_input":"2024-05-10T18:14:41.400716Z","iopub.status.idle":"2024-05-10T18:14:41.438546Z","shell.execute_reply.started":"2024-05-10T18:14:41.400687Z","shell.execute_reply":"2024-05-10T18:14:41.437527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Honeslty, it's a lot of. Probably it relates to the specifics of the task.\n\nAt the next step it's needed to write function to decode mask and encode result at back.","metadata":{}},{"cell_type":"code","source":"def random_colors(N, bright=True):\n    \"\"\"\n    Generate random colors.\n    To get visually distinct colors, generate them in HSV space then\n    convert to RGB.\n    \"\"\"\n    brightness = 1.0 if bright else 0.7\n    hsv = [(i / N, 1, brightness) for i in range(N)]\n    colors = list(map(lambda c: colorsys.hsv_to_rgb(*c), hsv))\n    random.shuffle(colors)\n    \n    return colors","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.440083Z","iopub.execute_input":"2024-05-10T18:14:41.440870Z","iopub.status.idle":"2024-05-10T18:14:41.448178Z","shell.execute_reply.started":"2024-05-10T18:14:41.440833Z","shell.execute_reply":"2024-05-10T18:14:41.447209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_rle_mask(rle_mask, shape=(520, 704)):\n\n    \"\"\"\n    Decode run-length encoded segmentation mask string into 2d array\n\n    Parameters\n    ----------\n    rle_mask (str): Run-length encoded segmentation mask string\n    rle_mask - it is a string with start and length coordinate values\n    \n    shape (tuple): Height and width of the mask\n\n    Returns\n    -------\n    mask [numpy.ndarray of shape (height, width)]: Decoded 2d segmentation mask\n    \"\"\"\n\n    rle_mask = rle_mask.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (rle_mask[0:][::2], rle_mask[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n\n    mask = np.zeros((shape[0] * shape[1]), dtype=np.uint8)\n    \n    for start, end in zip(starts, ends):\n        mask[start:end] = 1\n\n    mask = mask.reshape(shape[0], shape[1])\n    # Transform mask value (necessary step before Image augmentation)\n    mask = np.uint8(mask)\n    \n    return mask\n\n\ndef encode_rle_mask(mask, shape=(520, 704)):\n    \"\"\"\n    (Used for create submission file)\n    Parameters\n    ----------\n    mask with a given shape\n    \"\"\"\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    rle = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    rle[1::2] -= rle[::2]\n    \n    return rle.tolist()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.449351Z","iopub.execute_input":"2024-05-10T18:14:41.449688Z","iopub.status.idle":"2024-05-10T18:14:41.460474Z","shell.execute_reply.started":"2024-05-10T18:14:41.449663Z","shell.execute_reply":"2024-05-10T18:14:41.459456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***","metadata":{}},{"cell_type":"code","source":"\"\"\"\nHow does work encode_rle_mask() function?\n\nThere is a mask of one object at the image\n\"\"\"\nmask_ex = decode_rle_mask(rles[0])\nplt.imshow(mask_ex)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.461730Z","iopub.execute_input":"2024-05-10T18:14:41.462051Z","iopub.status.idle":"2024-05-10T18:14:41.792418Z","shell.execute_reply.started":"2024-05-10T18:14:41.462025Z","shell.execute_reply":"2024-05-10T18:14:41.791478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# At the next step transform this mask to flatten():\nprint(mask_ex.flatten().shape)\n\n# Add lenght value:\n# the first values are 0\n# the next is rle and rle[1::2]\nprint(mask_ex.flatten()[1::2].shape)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.793431Z","iopub.execute_input":"2024-05-10T18:14:41.793706Z","iopub.status.idle":"2024-05-10T18:14:41.799442Z","shell.execute_reply.started":"2024-05-10T18:14:41.793683Z","shell.execute_reply":"2024-05-10T18:14:41.798412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# encode_rle_mask() function result\nstr(encode_rle_mask(mask_ex))","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.800649Z","iopub.execute_input":"2024-05-10T18:14:41.800932Z","iopub.status.idle":"2024-05-10T18:14:41.810162Z","shell.execute_reply.started":"2024-05-10T18:14:41.800907Z","shell.execute_reply":"2024-05-10T18:14:41.809272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# After this step we wait to get original image\nmask_encode_decode = decode_rle_mask(str(encode_rle_mask(mask_ex)).replace('[', '').replace(']', '').replace(',', ''))\nplt.imshow(mask_encode_decode)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:41.811502Z","iopub.execute_input":"2024-05-10T18:14:41.811773Z","iopub.status.idle":"2024-05-10T18:14:42.133213Z","shell.execute_reply.started":"2024-05-10T18:14:41.811749Z","shell.execute_reply":"2024-05-10T18:14:42.132240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It's work well!","metadata":{}},{"cell_type":"code","source":"#  Let's create masks\nmasks = []\n\nfor mask in train_df.loc[train_df['id'] == img_id, 'annotation'].values:\n    decoded_mask = decode_rle_mask(rle_mask=mask, shape=img.shape)\n    \n    masks.append(decoded_mask)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:42.136415Z","iopub.execute_input":"2024-05-10T18:14:42.136810Z","iopub.status.idle":"2024-05-10T18:14:42.223077Z","shell.execute_reply.started":"2024-05-10T18:14:42.136767Z","shell.execute_reply":"2024-05-10T18:14:42.222184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(masks)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:42.385154Z","iopub.execute_input":"2024-05-10T18:14:42.385999Z","iopub.status.idle":"2024-05-10T18:14:42.392334Z","shell.execute_reply.started":"2024-05-10T18:14:42.385958Z","shell.execute_reply":"2024-05-10T18:14:42.391290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"masks_stack = np.stack(masks)\nmasks_stack.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:42.720679Z","iopub.execute_input":"2024-05-10T18:14:42.721067Z","iopub.status.idle":"2024-05-10T18:14:42.794567Z","shell.execute_reply.started":"2024-05-10T18:14:42.721039Z","shell.execute_reply":"2024-05-10T18:14:42.793511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"masks_stack","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:43.050667Z","iopub.execute_input":"2024-05-10T18:14:43.051421Z","iopub.status.idle":"2024-05-10T18:14:43.059503Z","shell.execute_reply.started":"2024-05-10T18:14:43.051390Z","shell.execute_reply":"2024-05-10T18:14:43.058279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_bboxes_from_mask(masks):\n    \"\"\"\n    Function to get bboxes in coco format\n    \"\"\"\n    coco_boxes = []\n    \n    for mask in masks:\n        pos = np.nonzero(mask)\n        xmin = np.min(pos[1])\n        xmax = np.max(pos[1])\n        ymin = np.min(pos[0])\n        ymax = np.max(pos[0])\n        coco_boxes.append([xmin, ymin, xmax, ymax])\n    \n    # Преобразуем в формат np для ДатаЛоадер  \n    coco_boxes = np.asarray(coco_boxes)\n    \n    return coco_boxes","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:43.337837Z","iopub.execute_input":"2024-05-10T18:14:43.338745Z","iopub.status.idle":"2024-05-10T18:14:43.345119Z","shell.execute_reply.started":"2024-05-10T18:14:43.338711Z","shell.execute_reply":"2024-05-10T18:14:43.344129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Let's check result of get_bboxes_from_mask() function\nbboxes = get_bboxes_from_mask(masks_stack)\nbboxes","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:43.622681Z","iopub.execute_input":"2024-05-10T18:14:43.623031Z","iopub.status.idle":"2024-05-10T18:14:44.139983Z","shell.execute_reply.started":"2024-05-10T18:14:43.623003Z","shell.execute_reply":"2024-05-10T18:14:44.139018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def apply_mask(image, mask, color, alpha=0.5):\n    \"\"\"\n    Apply the given mask to the image\n    \"\"\"\n    for c in range(3):\n        image[:, :, c] = np.where(mask == 1,\n                                  image[:, :, c] *\n                                  (1 - alpha) + alpha * color[c] * 255,\n                                  image[:, :, c])\n    \n    return image\n\n# def plot_image_annotations(image, masks, bboxes, aug=None):\n#     for box in bboxes:\n#         image = cv2.rectangle(image, (box[2], box[3]), (box[0], box[1]), (255, 255, 0), thickness=2)\n        \n#     plt.figure(figsize=(15, 15))\n#     plt.imshow(image)\n#     plt.show()\n\ndef plot_image_annotations(image, masks, bboxes, labels, aug=None):\n    \"\"\"\n    Function to plot masks and bboxes with augmentation\n    \"\"\"\n    image = image.copy()\n    \n    colors = {unique_lbl:random_colors(1, True)[0] for unique_lbl in np.unique(labels)}\n    \n    if aug is not None:\n        augmented = aug(image=image, masks=masks, bboxes=bboxes,\n                        labels=labels)\n        \n        image = augmented['image']\n        masks = augmented['masks']\n        bboxes = augmented['bboxes']\n    \n    bboxes = np.stack(bboxes).astype(int)\n        \n    for idx, box in enumerate(bboxes):\n        color = tuple([int(value*255) for value in colors[labels[idx]]])\n        image = cv2.rectangle(image, (box[2], box[3]), (box[0], box[1]), color=color, thickness=2)\n    \n    for idx, mask in enumerate(masks):\n        color = colors[labels[idx]]\n        image = apply_mask(image, mask, color)\n    \n    plt.figure(figsize=(15, 15))\n    plt.imshow(image)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:44.141384Z","iopub.execute_input":"2024-05-10T18:14:44.141668Z","iopub.status.idle":"2024-05-10T18:14:44.153106Z","shell.execute_reply.started":"2024-05-10T18:14:44.141644Z","shell.execute_reply":"2024-05-10T18:14:44.152176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_image_annotations(image=img, masks=masks_stack, bboxes=bboxes, labels=labels)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:44.238214Z","iopub.execute_input":"2024-05-10T18:14:44.238838Z","iopub.status.idle":"2024-05-10T18:14:47.030739Z","shell.execute_reply.started":"2024-05-10T18:14:44.238811Z","shell.execute_reply":"2024-05-10T18:14:47.029736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_augmentations = A.Compose([\n    A.Resize(640, 640),\n#     A.RandomResizedCrop(640, 640, scale=(0.8, 1.0), ratio=(0.9, 1.3)),\n    A.HorizontalFlip(),\n    A.VerticalFlip(),\n    A.RandomRotate90()\n], bbox_params={\"format\": \"pascal_voc\",\n                \"min_area\": 0,\n                \"min_visibility\": 0,\n                'label_fields': ['labels']\n               })","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:47.032466Z","iopub.execute_input":"2024-05-10T18:14:47.032761Z","iopub.status.idle":"2024-05-10T18:14:47.039075Z","shell.execute_reply.started":"2024-05-10T18:14:47.032736Z","shell.execute_reply":"2024-05-10T18:14:47.037919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_image_annotations(image=img, masks=masks_stack, bboxes=bboxes, labels=labels, aug=train_augmentations)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:47.040671Z","iopub.execute_input":"2024-05-10T18:14:47.041462Z","iopub.status.idle":"2024-05-10T18:14:50.307110Z","shell.execute_reply.started":"2024-05-10T18:14:47.041428Z","shell.execute_reply":"2024-05-10T18:14:50.306127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Model: MaskRCNN","metadata":{}},{"cell_type":"markdown","source":"# 2.1. Preparing data for model: Dataset, DataLoader","metadata":{}},{"cell_type":"code","source":"\"\"\"\nCreate class CellSegData to get train Dataset\n\"\"\"\nclass CellSegData(torch.utils.data.Dataset):\n    \n    def __init__(self, root, df, split='train', aug=None, cls_map=None):\n        self.augmentations = aug\n        self.cls_map = cls_map\n        \n        train, test = train_test_split(df['id'].unique(), train_size=0.9)\n        \n        if split=='train':\n            self.dataset = train\n        else:\n            self.dataset = test\n        \n        # let's create pointers to splitting pictures so as not to search through all the data\n        self.dict_df = {img_id:df[df['id']==img_id] for img_id in tqdm(self.dataset)}\n        self.root = root\n        \n    def __len__(self):\n        return len(self.dataset)\n    \n    def __getitem__(self, index):\n        img_id = self.dataset[index]\n        image = cv2.imread(os.path.join(self.root, img_id+'.png'))\n        \n        info = self.dict_df[img_id]\n        n_objects = len(info['annotation'])\n                \n        # Replace code <labels, rles = get_targets_mask(info, img_id)> to:\n        labels = info[\"cell_type\"].apply(lambda x: self.cls_map[x]).values\n        rles = info[\"annotation\"].values\n        \n        masks = []\n        for mask in rles:\n            decoded_mask = decode_rle_mask(rle_mask=mask, shape=img.shape)\n            masks.append(decoded_mask)\n            \n        bboxes = get_bboxes_from_mask(masks)\n        \n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image, masks=masks, bboxes=bboxes,\n                            labels=labels)\n            image = augmented['image']\n            masks = augmented['masks']\n            bboxes = augmented['bboxes']\n            bboxes = np.stack(bboxes).astype(int)\n        \n        # Transform to np\n        masks = np.asarray(masks)\n        \n        bboxes = torch.as_tensor(bboxes, dtype=torch.int64)\n#         labels = torch.as_tensor(labels, dtype=torch.int64)\n#         masks = torch.as_tensor(masks, dtype=torch.uint8)\n\n        is_bad_labels = False\n        degenerate_boxes = bboxes[:, 2:] <= bboxes[:, :2]\n        \n        if degenerate_boxes.any():\n            is_bad_labels = True\n            # print the first degenerate box\n            bb_idxs = torch.where(degenerate_boxes.any(dim=1))[0].numpy()\n\n        labels = torch.as_tensor([1 for i in range(len(bboxes))], dtype=torch.int64)\n        masks = torch.as_tensor(masks, dtype=torch.uint8)\n\n        if is_bad_labels:\n            bboxes = bboxes[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n            labels = labels[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n            masks = masks[[i for i in range(len(bboxes)) if i not in bb_idxs]]\n\n        image_id = torch.tensor([index])\n        area = (bboxes[:, 3] - bboxes[:, 1]) * (bboxes[:, 2] - bboxes[:, 0])\n        iscrowd = torch.zeros((n_objects,), dtype=torch.int64)\n\n        # This is the required target for the Mask R-CNN\n        target = {\n            'boxes': bboxes,\n            'labels': labels,\n            'masks': masks,\n            'image_id': image_id,\n            'area': area,\n            'iscrowd': iscrowd\n        }\n        \n        image = image.transpose((2, 0, 1))\n        \n        return torch.Tensor(image), target #torch.Tensor(image)/255, target","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:50.309663Z","iopub.execute_input":"2024-05-10T18:14:50.310019Z","iopub.status.idle":"2024-05-10T18:14:50.332393Z","shell.execute_reply.started":"2024-05-10T18:14:50.309990Z","shell.execute_reply":"2024-05-10T18:14:50.331424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use our dataset and defined transformations\ndataset = CellSegData(f'{ROOT}/train', train_df, 'train', train_augmentations, cls_map)\ndataset_test = CellSegData(f'{ROOT}/train', train_df, 'test', train_augmentations, cls_map)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:50.333469Z","iopub.execute_input":"2024-05-10T18:14:50.333742Z","iopub.status.idle":"2024-05-10T18:14:58.530056Z","shell.execute_reply.started":"2024-05-10T18:14:50.333715Z","shell.execute_reply":"2024-05-10T18:14:58.529009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:14:58.531361Z","iopub.execute_input":"2024-05-10T18:14:58.531667Z","iopub.status.idle":"2024-05-10T18:14:58.584818Z","shell.execute_reply.started":"2024-05-10T18:14:58.531640Z","shell.execute_reply":"2024-05-10T18:14:58.583747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2.2 Model training","metadata":{}},{"cell_type":"code","source":"!git clone https://github.com/pytorch/vision.git","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:15:11.852220Z","iopub.execute_input":"2024-05-10T18:15:11.852870Z","iopub.status.idle":"2024-05-10T18:16:17.830192Z","shell.execute_reply.started":"2024-05-10T18:15:11.852840Z","shell.execute_reply":"2024-05-10T18:16:17.829060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp vision/references/detection/engine.py /kaggle/working\n!cp vision/references/detection/utils.py /kaggle/working\n!cp vision/references/detection/transforms.py /kaggle/working\n!cp vision/references/detection/coco_eval.py /kaggle/working\n!cp vision/references/detection/coco_utils.py /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:17.832673Z","iopub.execute_input":"2024-05-10T18:16:17.833444Z","iopub.status.idle":"2024-05-10T18:16:22.841710Z","shell.execute_reply.started":"2024-05-10T18:16:17.833404Z","shell.execute_reply":"2024-05-10T18:16:22.840375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!pip install pycocotools","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:22.843244Z","iopub.execute_input":"2024-05-10T18:16:22.843580Z","iopub.status.idle":"2024-05-10T18:16:36.729001Z","shell.execute_reply.started":"2024-05-10T18:16:22.843550Z","shell.execute_reply":"2024-05-10T18:16:36.727714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from engine import train_one_epoch, evaluate\nimport utils","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:36.731931Z","iopub.execute_input":"2024-05-10T18:16:36.732704Z","iopub.status.idle":"2024-05-10T18:16:36.741729Z","shell.execute_reply.started":"2024-05-10T18:16:36.732666Z","shell.execute_reply":"2024-05-10T18:16:36.740729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define training and validation data loaders\ndata_loader = torch.utils.data.DataLoader(\n    dataset, batch_size=2, shuffle=True, num_workers=2, prefetch_factor=2, collate_fn=utils.collate_fn)\n\ndata_loader_test = torch.utils.data.DataLoader(\n    dataset_test, batch_size=1, shuffle=False, num_workers=4, collate_fn=utils.collate_fn) ","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:36.742882Z","iopub.execute_input":"2024-05-10T18:16:36.743195Z","iopub.status.idle":"2024-05-10T18:16:36.749904Z","shell.execute_reply.started":"2024-05-10T18:16:36.743171Z","shell.execute_reply":"2024-05-10T18:16:36.748783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***","metadata":{}},{"cell_type":"code","source":"NUM_CLASSES = 2\n\nmodel = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=False,\n                                                           box_detections_per_img=600)\n\n# get the number of input features for the classifier\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n# replace the pre-trained head with a new one\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, NUM_CLASSES)\n\n# now get the number of input features for the mask classifier\nin_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\nhidden_layer = 256\n# and replace the mask predictor with a new one\nmodel.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, NUM_CLASSES)\n# weights = torch.load('weights/checkpoint.pth')\n# move model to the right device\n# model.state_dict(weights['model'])\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:36.751254Z","iopub.execute_input":"2024-05-10T18:16:36.751608Z","iopub.status.idle":"2024-05-10T18:16:38.663671Z","shell.execute_reply.started":"2024-05-10T18:16:36.751556Z","shell.execute_reply":"2024-05-10T18:16:38.662677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# construct an optimizer\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.AdamW(params, lr=0.0001)\n\n# and a learning rate scheduler\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer,\n                                               step_size=3,\n                                               gamma=0.1)\n\noutput_dir = 'weights'\nos.makedirs(output_dir, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:38.664908Z","iopub.execute_input":"2024-05-10T18:16:38.665243Z","iopub.status.idle":"2024-05-10T18:16:38.672711Z","shell.execute_reply.started":"2024-05-10T18:16:38.665217Z","shell.execute_reply":"2024-05-10T18:16:38.671550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def analyze_train_sample(model, ds_train, sample_index):\n    img, targets = ds_train[sample_index]\n    \n    plt.imshow(img.numpy().astype(np.uint8).transpose((1, 2, 0)))\n    plt.title(\"Image\")\n    plt.show()\n\n    masks = np.zeros((640, 640))\n    for mask in targets['masks']:\n        masks = np.logical_or(masks, mask)\n        \n    plt.imshow(img.numpy().transpose((1, 2, 0)))\n    plt.imshow(masks, alpha=0.3)\n    plt.title(\"Ground truth\")\n    plt.show()\n\n    model.eval()\n    \n    with torch.no_grad():\n        preds = model([img.cuda()])[0]\n\n    plt.imshow(img.cpu().numpy().transpose((1, 2, 0)))\n    all_preds_masks = np.zeros((640, 640))\n    for mask in preds['masks'].cpu().detach().numpy():\n        \n        all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n        \n    plt.imshow(all_preds_masks, alpha=0.4)\n    plt.title(\"Predictions\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:38.674289Z","iopub.execute_input":"2024-05-10T18:16:38.674747Z","iopub.status.idle":"2024-05-10T18:16:38.686140Z","shell.execute_reply.started":"2024-05-10T18:16:38.674713Z","shell.execute_reply":"2024-05-10T18:16:38.685164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# let's train it for 10 epochs\nnum_epochs = 10\n\nfor epoch in tqdm(range(num_epochs)):\n    # train for one epoch, printing every 10 iterations\n    model.train()\n        \n    metric_logger = utils.MetricLogger(delimiter=\"  \")\n    metric_logger.add_meter(\"lr\", utils.SmoothedValue(window_size=1, fmt=\"{value:.6f}\"))\n    header = f\"Epoch: [{epoch}]\"\n\n    lr_scheduler = None\n        \n    if epoch == 0:\n        warmup_factor = 1.0 / 1000\n        warmup_iters = min(1000, len(data_loader) - 1)\n\n        lr_scheduler = torch.optim.lr_scheduler.LinearLR(\n            optimizer, start_factor=warmup_factor, total_iters=warmup_iters\n        )\n\n    for images, targets in metric_logger.log_every(data_loader, 10, header):\n        images = list(image.to(device) for image in images)\n        targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n\n        loss_dict = model(images, targets)\n        losses = sum(loss for loss in loss_dict.values())\n\n        # reduce losses over all GPUs for logging purposes\n        loss_dict_reduced = utils.reduce_dict(loss_dict)\n        losses_reduced = sum(loss for loss in loss_dict_reduced.values())\n\n        loss_value = losses_reduced.item()\n\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n        if lr_scheduler is not None:\n            lr_scheduler.step()\n\n        metric_logger.update(loss=losses_reduced, **loss_dict_reduced)\n        metric_logger.update(lr=optimizer.param_groups[0][\"lr\"])\n\n    analyze_train_sample(model, dataset, 20)\n    # evaluate on the test dataset\n    # evaluate(model, data_loader_test, device=device)\n    if output_dir:\n        checkpoint = {\n            \"model\": model.state_dict(),\n            \"optimizer\": optimizer.state_dict(),\n            \"epoch\": epoch,\n        }\n        utils.save_on_master(checkpoint, os.path.join(output_dir, f\"model_{epoch}.pth\"))\n        utils.save_on_master(checkpoint, os.path.join(output_dir, \"checkpoint.pth\"))","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:16:46.912836Z","iopub.execute_input":"2024-05-10T18:16:46.913590Z","iopub.status.idle":"2024-05-10T18:47:21.428089Z","shell.execute_reply.started":"2024-05-10T18:16:46.913556Z","shell.execute_reply":"2024-05-10T18:47:21.426989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Prediction","metadata":{}},{"cell_type":"code","source":"\"\"\"\nCreate Dataset to read test data\n\"\"\"\nclass CellTestDataset(torch.utils.data.Dataset):\n    def __init__(self, image_dir, transforms=None):\n        self.transforms = transforms\n        self.image_dir = image_dir\n        self.image_ids = [f[:-4]for f in os.listdir(self.image_dir)]\n        \n    \n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        image_path = os.path.join(self.image_dir, image_id + \".png\")\n        image = Image.open(image_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            image, _ = self.transforms(image=image, target=None)\n            \n        return {'image': image, 'image_id': image_id}\n\n    \n    def __len__(self):\n        return len(self.image_ids)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:49:14.073253Z","iopub.execute_input":"2024-05-10T18:49:14.074251Z","iopub.status.idle":"2024-05-10T18:49:14.082344Z","shell.execute_reply.started":"2024-05-10T18:49:14.074218Z","shell.execute_reply":"2024-05-10T18:49:14.081408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine function to transform data\n\"\"\"\n\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nclass Compose:\n    def __init__(self, transforms):\n        self.transforms = transforms\n\n    def __call__(self, image, target):\n        for t in self.transforms:\n            image, target = t(image, target)\n        return image, target\n    \n\nclass Normalize:\n    def __call__(self, image, target):\n        image = F.normalize(image, RESNET_MEAN, RESNET_STD)\n        return image, target\n    \n    \nclass ToTensor:\n    def __call__(self, image, target):\n        image = F.to_tensor(image)\n        \n        return image, target\n    \n    \ndef get_transform(train):\n    transforms = [ToTensor()]\n    if True:\n        transforms.append(Normalize())\n\n    return Compose(transforms)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:49:18.916566Z","iopub.execute_input":"2024-05-10T18:49:18.917190Z","iopub.status.idle":"2024-05-10T18:49:18.925664Z","shell.execute_reply.started":"2024-05-10T18:49:18.917156Z","shell.execute_reply":"2024-05-10T18:49:18.924565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_test = CellTestDataset(f'{ROOT}/test', transforms=get_transform(train=False))\nds_test[0]","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:49:20.952627Z","iopub.execute_input":"2024-05-10T18:49:20.953321Z","iopub.status.idle":"2024-05-10T18:49:20.991353Z","shell.execute_reply.started":"2024-05-10T18:49:20.953290Z","shell.execute_reply":"2024-05-10T18:49:20.990308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nRewrite function to encode rle masks\n\"\"\"\ndef rle_encoding(x):\n    dots = np.where(x.flatten() == 1)[0]\n    run_lengths = []\n    prev = -2\n    \n    for b in dots:\n        if (b>prev+1): run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n        \n    return ' '.join(map(str, run_lengths))\n\n\ndef remove_overlapping_pixels(mask, other_masks):\n    for other_mask in other_masks:\n        if np.sum(np.logical_and(mask, other_mask)) > 0:\n            mask[np.logical_and(mask, other_mask)] = 0\n    return mask","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:49:25.513069Z","iopub.execute_input":"2024-05-10T18:49:25.513475Z","iopub.status.idle":"2024-05-10T18:49:25.520593Z","shell.execute_reply.started":"2024-05-10T18:49:25.513435Z","shell.execute_reply":"2024-05-10T18:49:25.519588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nRun predictions\n\"\"\"\nmodel.eval();\n\nsubmission = []\n\nfor sample in ds_test:\n    img = sample['image']\n    image_id = sample['image_id']\n    \n    with torch.no_grad():\n        result = model([img.to(device)])[0]\n    \n    previous_masks = []\n    \n    for i, mask in enumerate(result[\"masks\"]):\n        # Filter-out low-scoring results. Not tried yet.\n        score = result[\"scores\"][i].cpu().item()\n        if score < 0.3:\n            continue\n        \n        mask = mask.cpu().numpy()\n        # Keep only highly likely pixels\n        binary_mask = mask > 0.5\n        binary_mask = remove_overlapping_pixels(binary_mask, previous_masks)\n        previous_masks.append(binary_mask)\n        rle = rle_encoding(binary_mask)\n        \n    submission.append((image_id, rle))\n    print(submission)\n    \n    plt.figure(figsize=(12,12))\n    ax1 = plt.subplot(121)\n    ax1.imshow(img.numpy().transpose((1,2,0)))\n    all_preds_masks = np.zeros((520, 704))\n    \n    for mask in result['masks'].cpu().detach().numpy():\n        all_preds_masks = np.logical_or(all_preds_masks, mask[0] > 0.5)\n        \n    ax2 = plt.subplot(122)\n    ax2.imshow(img.numpy().transpose((1,2,0)))\n    ax2.imshow(all_preds_masks, alpha=0.3)\n    plt.tight_layout()\n    plt.show()\n    \n#     # Add empty prediction if no RLE was generated for this image\n#     all_images_ids = [image_id for image_id, rle in submission]\n#     if image_id not in all_images_ids:\n#         submission.append((image_id, \"\"))","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:50:04.454799Z","iopub.execute_input":"2024-05-10T18:50:04.455255Z","iopub.status.idle":"2024-05-10T18:52:18.997700Z","shell.execute_reply.started":"2024-05-10T18:50:04.455227Z","shell.execute_reply":"2024-05-10T18:52:18.996718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Submission","metadata":{}},{"cell_type":"code","source":"\"\"\"\nSubmission file\n\"\"\"\ndf_sub = pd.DataFrame(submission, columns=['id', 'predicted'])\ndf_sub.to_csv(\"submission.csv\", index=False)\ndf_sub.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T18:52:37.649976Z","iopub.execute_input":"2024-05-10T18:52:37.650792Z","iopub.status.idle":"2024-05-10T18:52:37.666723Z","shell.execute_reply.started":"2024-05-10T18:52:37.650759Z","shell.execute_reply":"2024-05-10T18:52:37.665676Z"},"trusted":true},"execution_count":null,"outputs":[]}]}