{"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":"# **[HuBMAP 2023] - Torch Mask R-CNN**","metadata":{}},{"cell_type":"markdown","source":"This notebook is inspired from:\n* https://www.kaggle.com/code/julian3833/sartorius-starter-torch-mask-r-cnn-lb-0-273\n* https://www.kaggle.com/code/rluethy/sartorius-torch-mask-r-cnn\n\nThanks [Julián Peller](https://www.kaggle.com/julian3833) AND [ROLAND LUETHY](https://www.kaggle.com/rluethy)","metadata":{}},{"cell_type":"code","source":"!cp -r /kaggle/input/pycocotools/ /kaggle/working/pycocotools\n!pip install /kaggle/working/pycocotools/pycocotools-2.0.6  --no-index --find-links=/kaggle/working/pycocotools/","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-25T10:22:34.345338Z","iopub.execute_input":"2023-06-25T10:22:34.345673Z","iopub.status.idle":"2023-06-25T10:23:10.964456Z","shell.execute_reply.started":"2023-06-25T10:22:34.345644Z","shell.execute_reply":"2023-06-25T10:23:10.963164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports","metadata":{"id":"1NZ-_x8E0zRW","papermill":{"duration":0.028124,"end_time":"2021-11-08T00:11:44.476803","exception":false,"start_time":"2021-11-08T00:11:44.448679","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport random\nimport time\nimport collections\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport cv2\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.notebook import tqdm\nimport torch\nimport torchvision\nfrom torchvision.transforms import ToPILImage\nfrom torchvision.transforms import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","id":"oQYSI0Y00zRX","papermill":{"duration":3.361808,"end_time":"2021-11-08T00:11:47.868521","exception":false,"start_time":"2021-11-08T00:11:44.506713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:10.970165Z","iopub.execute_input":"2023-06-25T10:23:10.972369Z","iopub.status.idle":"2023-06-25T10:23:16.045128Z","shell.execute_reply.started":"2023-06-25T10:23:10.972332Z","shell.execute_reply":"2023-06-25T10:23:16.044133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fix_all_seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n    \nfix_all_seeds(2021)","metadata":{"id":"Y7fwE02H0zRY","papermill":{"duration":0.086836,"end_time":"2021-11-08T00:11:47.987278","exception":false,"start_time":"2021-11-08T00:11:47.900442","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.046861Z","iopub.execute_input":"2023-06-25T10:23:16.047484Z","iopub.status.idle":"2023-06-25T10:23:16.130047Z","shell.execute_reply.started":"2023-06-25T10:23:16.047448Z","shell.execute_reply":"2023-06-25T10:23:16.129039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configuration","metadata":{"id":"NqZ4eNVK0zRZ","papermill":{"duration":0.027823,"end_time":"2021-11-08T00:11:48.042919","exception":false,"start_time":"2021-11-08T00:11:48.015096","status":"completed"},"tags":[]}},{"cell_type":"code","source":"TEST = False\nif os.path.exists(\"/kaggle/input/hubmap-hacking-the-human-vasculature\"):\n    # running on kaggle\n    data_directory = '/kaggle/input/hubmap-hacking-the-human-vasculature'\n    DEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n    BATCH_SIZE = 4\n    NUM_EPOCHS = 20 \n\nelif 'google.colab' in str(get_ipython()):\n    # running on CoLab\n    from google.colab import drive\n    drive.mount('/content/drive')\n    data_directory = '/content/drive/MyDrive/kaggle/[HuBMAP-2023]'\n    DEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n    BATCH_SIZE = 16\n    NUM_EPOCHS = 20\n    \nelse:\n    data_directory = 'input'\n    DEVICE = torch.device('cpu')\n    BATCH_SIZE = 2\n    NUM_EPOCHS = 1\n    TEST = True\n\nTRAIN_CSV = \"/kaggle/input/hubmap-dataframe/df_train.csv\"\n\nTRAIN_PATH = f\"{data_directory}/train\"\nTEST_PATH = f\"{data_directory}/test\"\n\nWIDTH = 512\nHEIGHT = 512\n\nresize_factor = False # 0.5\n\n# Normalize to resnet mean and std if True.\nNORMALIZE = False\nRESNET_MEAN = (0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\n# No changes tried with the optimizer yet.\nMOMENTUM = 0.9\nLEARNING_RATE = 0.001\nWEIGHT_DECAY = 0.0005\n\n# cell type specific thresholds\ncell_type_dict = {'blood_vessel': 1, 'glomerulus': 2, 'unsure': 3}\nmask_threshold_dict = {1: 0.40, 2: 0.80, 3:  0.80}\nmin_score_dict = {1: 0.40, 2: 0.80, 3: 0.80}\n\n# Use a StepLR scheduler if True. \nUSE_SCHEDULER = False\nTEST_SIZE=0.20\nBOX_DETECTIONS_PER_IMG = 500","metadata":{"id":"VSPe6quz0zRZ","lines_to_next_cell":1,"outputId":"e2cca0e2-0ada-471b-e688-33ee16049407","papermill":{"duration":0.042722,"end_time":"2021-11-08T00:11:48.113166","exception":false,"start_time":"2021-11-08T00:11:48.070444","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.133811Z","iopub.execute_input":"2023-06-25T10:23:16.134710Z","iopub.status.idle":"2023-06-25T10:23:16.148481Z","shell.execute_reply.started":"2023-06-25T10:23:16.134675Z","shell.execute_reply":"2023-06-25T10:23:16.147645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utilities","metadata":{"id":"iIpvad7y0zRb","papermill":{"duration":0.029286,"end_time":"2021-11-08T00:11:48.170772","exception":false,"start_time":"2021-11-08T00:11:48.141486","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ref: https://www.kaggle.com/inversion/run-length-decoding-quick-start\ndef rle_decode(mask_rle, shape, color=1):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height, width, channels) of array to return\n    color: color for the mask\n    Returns numpy array (mask)\n\n    '''\n    s = mask_rle.split()\n\n    starts = list(map(lambda x: int(x) - 1, s[0::2]))\n    lengths = list(map(int, s[1::2]))\n    ends = [x + y for x, y in zip(starts, lengths)]\n    if len(shape)==3:\n        img = np.zeros((shape[0] * shape[1], shape[2]), dtype=np.float32)\n    else:\n        img = np.zeros(shape[0] * shape[1], dtype=np.float32)\n    for start, end in zip(starts, ends):\n        img[start : end] = color\n\n    return img.reshape(shape)\n\n\ndef rle_encoding(x):\n    dots = np.where(x.flatten() == 1)[0]\n    run_lengths = []\n    prev = -2\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    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\n\ndef combine_masks(masks, mask_threshold):\n    \"\"\"\n    combine masks into one image\n    \"\"\"\n    maskimg = np.zeros((HEIGHT, WIDTH))\n    # print(len(masks.shape), masks.shape)\n    for m, mask in enumerate(masks,1):\n        maskimg[mask>mask_threshold] = m\n    return maskimg\n\n\ndef get_filtered_masks(pred):\n    \"\"\"\n    filter masks using MIN_SCORE for mask and MAX_THRESHOLD for pixels\n    \"\"\"\n    use_masks = []\n    use_labels = []\n    for i, mask in enumerate(pred[\"masks\"]):\n        # Filter-out low-scoring results. Not tried yet.\n        scr = pred[\"scores\"][i].cpu().item()\n        label = pred[\"labels\"][i].cpu().item()\n        if scr > min_score_dict[label]:\n            mask = mask.cpu().numpy().squeeze()\n            # Keep only highly likely pixels\n            binary_mask = mask > mask_threshold_dict[label]\n            binary_mask = remove_overlapping_pixels(binary_mask, use_masks)\n            use_masks.append(binary_mask)\n            use_labels.append(label)\n\n    return use_masks,use_labels\n","metadata":{"papermill":{"duration":0.051897,"end_time":"2021-11-08T00:11:48.251907","exception":false,"start_time":"2021-11-08T00:11:48.200010","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.150237Z","iopub.execute_input":"2023-06-25T10:23:16.151078Z","iopub.status.idle":"2023-06-25T10:23:16.169841Z","shell.execute_reply.started":"2023-06-25T10:23:16.151041Z","shell.execute_reply":"2023-06-25T10:23:16.168640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Metric: mAP IoU threshold 0.6","metadata":{"papermill":{"duration":0.027951,"end_time":"2021-11-08T00:11:48.308157","exception":false,"start_time":"2021-11-08T00:11:48.280206","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def compute_iou(labels, y_pred, verbose=0):\n    \"\"\"\n    Computes the IoU for instance labels and predictions.\n\n    Args:\n        labels (np array): Labels.\n        y_pred (np array): predictions\n\n    Returns:\n        np array: IoU matrix, of size true_objects x pred_objects.\n    \"\"\"\n\n    true_objects = len(np.unique(labels))\n    pred_objects = len(np.unique(y_pred))\n\n    if verbose:\n        print(\"Number of true objects: {}\".format(true_objects))\n        print(\"Number of predicted objects: {}\".format(pred_objects))\n\n    # Compute intersection between all objects\n    intersection = np.histogram2d(\n        labels.flatten(), y_pred.flatten(), bins=(true_objects, pred_objects)\n    )[0]\n\n    # Compute areas (needed for finding the union between all objects)\n    area_true = np.histogram(labels, bins=true_objects)[0]\n    area_pred = np.histogram(y_pred, bins=pred_objects)[0]\n    area_true = np.expand_dims(area_true, -1)\n    area_pred = np.expand_dims(area_pred, 0)\n\n    # Compute union\n    union = area_true + area_pred - intersection\n    intersection = intersection[1:, 1:] # exclude background\n    union = union[1:, 1:]\n    union[union == 0] = 1e-9\n    iou = intersection / union\n    \n    return iou  \n\ndef precision_at(threshold, iou):\n    \"\"\"\n    Computes the precision at a given threshold.\n\n    Args:\n        threshold (float): Threshold.\n        iou (np array): IoU matrix.\n\n    Returns:\n        int: Number of true positives,\n        int: Number of false positives,\n        int: Number of false negatives.\n    \"\"\"\n    matches = iou > threshold\n    true_positives = np.sum(matches, axis=1) == 1  # Correct objects\n    false_positives = np.sum(matches, axis=0) == 0  # Missed objects\n    false_negatives = np.sum(matches, axis=1) == 0  # Extra objects\n    tp, fp, fn = (\n        np.sum(true_positives),\n        np.sum(false_positives),\n        np.sum(false_negatives),\n    )\n    return tp, fp, fn\n\ndef iou_map(truths, preds, verbose=0):\n    \"\"\"\n    Computes the metric for the competition.\n    Masks contain the segmented pixels where each object has one value associated,\n    and 0 is the background.\n\n    Args:\n        truths (list of masks): Ground truths.\n        preds (list of masks): Predictions.\n        verbose (int, optional): Whether to print infos. Defaults to 0.\n\n    Returns:\n        float: mAP.\n    \"\"\"\n    ious = [compute_iou(truth, pred, verbose) for truth, pred in zip(truths, preds)]\n\n    if verbose:\n        print(\"Thresh\\tTP\\tFP\\tFN\\tPrec.\")\n\n    prec = []\n    for t in np.arange(0.6, 6.5, 0.05):\n        tps, fps, fns = 0, 0, 0\n        for iou in ious:\n            tp, fp, fn = precision_at(t, iou)\n            tps += tp\n            fps += fp\n            fns += fn\n\n        p = tps / (tps + fps + fns)\n        prec.append(p)\n\n        if verbose:\n            print(\"{:1.3f}\\t{}\\t{}\\t{}\\t{:1.3f}\".format(t, tps, fps, fns, p))\n\n    if verbose:\n        print(\"AP\\t-\\t-\\t-\\t{:1.3f}\".format(np.mean(prec)))\n\n    return np.mean(prec)\n\n\ndef get_score(ds, mdl):\n    \"\"\"\n    Get average IOU mAP score for a dataset\n    \"\"\"\n    mdl.eval()\n    iouscore = 0\n    for i in tqdm(range(len(ds))):\n        img, targets = ds[i]\n        with torch.no_grad():\n            result = mdl([img.to(DEVICE)])[0]\n            \n        masks = combine_masks(targets['masks'], 0.5)\n        labels = pd.Series(result['labels'].cpu().numpy()).value_counts()\n\n        mask_threshold = mask_threshold_dict[labels.sort_values().index[-1]]\n        masks_p,labels_p=get_filtered_masks(result)\n        pred_masks = combine_masks(masks_p, mask_threshold)\n        iouscore += iou_map([masks],[pred_masks])\n    return iouscore / len(ds)\n","metadata":{"papermill":{"duration":0.053394,"end_time":"2021-11-08T00:11:48.389782","exception":false,"start_time":"2021-11-08T00:11:48.336388","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.173097Z","iopub.execute_input":"2023-06-25T10:23:16.173939Z","iopub.status.idle":"2023-06-25T10:23:16.197648Z","shell.execute_reply.started":"2023-06-25T10:23:16.173904Z","shell.execute_reply":"2023-06-25T10:23:16.196670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Transformations\nJust Horizontal and Vertical Flip for now.\n\nNormalization to Resnet's mean and std can be performed using the parameter `NORMALIZE` in the top cell.\n\nThe first 3 transformations come from [this](https://www.kaggle.com/abhishek/maskrcnn-utils) utils package by Abishek, `VerticalFlip` is my adaption of HorizontalFlip, and `Normalize` is of my own.","metadata":{"papermill":{"duration":0.028028,"end_time":"2021-11-08T00:11:48.445872","exception":false,"start_time":"2021-11-08T00:11:48.417844","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# These are slight redefinitions of torch.transformation classes\n# The difference is that they handle the target and the mask\n# Copied from Abishek, added new ones\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\nclass VerticalFlip:\n    def __init__(self, prob):\n        self.prob = prob\n\n    def __call__(self, image, target):\n        if random.random() < self.prob:\n            height, width = image.shape[-2:]\n            image = image.flip(-2)\n            bbox = target[\"boxes\"]\n            bbox[:, [1, 3]] = height - bbox[:, [3, 1]]\n            target[\"boxes\"] = bbox\n            target[\"masks\"] = target[\"masks\"].flip(-2)\n        return image, target\n\nclass HorizontalFlip:\n    def __init__(self, prob):\n        self.prob = prob\n\n    def __call__(self, image, target):\n        if random.random() < self.prob:\n            height, width = image.shape[-2:]\n            image = image.flip(-1)\n            bbox = target[\"boxes\"]\n            bbox[:, [0, 2]] = width - bbox[:, [2, 0]]\n            target[\"boxes\"] = bbox\n            target[\"masks\"] = target[\"masks\"].flip(-1)\n        return image, target\n\nclass Normalize:\n    def __call__(self, image, target):\n        image = F.normalize(image, RESNET_MEAN, RESNET_STD)\n        return image, target\n\nclass ToTensor:\n    def __call__(self, image, target):\n        image = F.to_tensor(image)\n        return image, target\n    \n\ndef get_transform(train):\n    transforms = [ToTensor()]\n    if NORMALIZE:\n        transforms.append(Normalize())\n    \n    # Data augmentation for train\n    if train: \n        transforms.append(HorizontalFlip(0.5))\n        transforms.append(VerticalFlip(0.5))\n\n    return Compose(transforms)","metadata":{"papermill":{"duration":0.04587,"end_time":"2021-11-08T00:11:48.519757","exception":false,"start_time":"2021-11-08T00:11:48.473887","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.199980Z","iopub.execute_input":"2023-06-25T10:23:16.200724Z","iopub.status.idle":"2023-06-25T10:23:16.216064Z","shell.execute_reply.started":"2023-06-25T10:23:16.200684Z","shell.execute_reply":"2023-06-25T10:23:16.215114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Dataset and DataLoader","metadata":{"id":"hHT_aovU0zRd","papermill":{"duration":0.029828,"end_time":"2021-11-08T00:11:48.577508","exception":false,"start_time":"2021-11-08T00:11:48.547680","status":"completed"},"tags":[]}},{"cell_type":"code","source":"cell_type_dict = {'blood_vessel': 1, 'glomerulus': 2, 'unsure': 3}\nclass HuBMAPDataset(Dataset):\n    def __init__(self, image_dir, df, transforms=None, resize=False):\n        self.transforms = transforms\n        self.image_dir = image_dir\n        self.df = df\n        \n        self.should_resize = resize is not False\n        if self.should_resize:\n            self.height = int(HEIGHT * resize)\n            self.width = int(WIDTH * resize)\n            print(\"image size used:\", self.height, self.width)\n        else:\n            self.height = HEIGHT\n            self.width = WIDTH\n        \n        self.image_info = collections.defaultdict(dict)\n        temp_df = self.df.groupby([\"image_id\"])[['annotations','category_id']].agg(lambda x: list(x)).reset_index()\n        for index, row in temp_df.iterrows():\n            self.image_info[index] = {\n                    'image_id': row['image_id'],\n                    'image_path': os.path.join(self.image_dir, row['image_id'] + '.tif'),\n                    'annotations': list(row[\"annotations\"]),\n                    'category_name': list(row[\"category_id\"])\n                    }\n            \n    def get_box(self, a_mask):\n        ''' Get the bounding box of a given mask '''\n        pos = np.where(a_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        return [xmin, ymin, xmax, ymax]\n\n    def __getitem__(self, idx):\n        ''' Get the image and the target'''\n        \n        img_path = self.image_info[idx][\"image_path\"]\n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        info = self.image_info[idx]\n\n        n_objects = len(info['annotations'])\n        masks = np.zeros((len(info['annotations']), self.height, self.width), dtype=np.uint8)\n        boxes = []\n        labels = []\n        for i, annotation in enumerate(info['annotations']):\n            a_mask = rle_decode(annotation, (HEIGHT, WIDTH))        \n            a_mask = np.array(a_mask) > 0\n            masks[i, :, :] = a_mask\n            boxes.append(self.get_box(a_mask))\n        labels = info[\"category_name\"]        \n        boxes = torch.as_tensor(boxes, dtype=torch.float32)\n        labels = torch.as_tensor(labels, dtype=torch.int64)\n        masks = torch.as_tensor(masks, dtype=torch.uint8)\n        image_id = torch.tensor([idx])\n        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])\n        iscrowd = torch.zeros((n_objects,), dtype=torch.int64)\n        target = {\n            'boxes': boxes,\n            'labels': labels,\n            'masks': masks,\n            'image_id': image_id,\n            'area': area,\n            'iscrowd': iscrowd\n        }\n\n        if self.transforms is not None:\n            img, target = self.transforms(img, target)\n\n        return img, target#,img_path\n\n    def __len__(self):\n        return len(self.image_info)","metadata":{"id":"C9Y03YgA0zRd","papermill":{"duration":0.053348,"end_time":"2021-11-08T00:11:48.658827","exception":false,"start_time":"2021-11-08T00:11:48.605479","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.218720Z","iopub.execute_input":"2023-06-25T10:23:16.219613Z","iopub.status.idle":"2023-06-25T10:23:16.238354Z","shell.execute_reply.started":"2023-06-25T10:23:16.219564Z","shell.execute_reply":"2023-06-25T10:23:16.237317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Reading Dataset","metadata":{}},{"cell_type":"code","source":"df_ = pd.read_csv(TRAIN_CSV)\ndf_.head(3)","metadata":{"execution":{"iopub.status.busy":"2023-06-25T10:23:16.239700Z","iopub.execute_input":"2023-06-25T10:23:16.240102Z","iopub.status.idle":"2023-06-25T10:23:16.484671Z","shell.execute_reply.started":"2023-06-25T10:23:16.240069Z","shell.execute_reply":"2023-06-25T10:23:16.483534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_images = df_.groupby([\"image_id\", \"category_name\"]).agg({'annotations': 'count'}).sort_values(\"annotations\", ascending=False).reset_index()","metadata":{"execution":{"iopub.status.busy":"2023-06-25T10:23:16.489985Z","iopub.execute_input":"2023-06-25T10:23:16.490738Z","iopub.status.idle":"2023-06-25T10:23:16.517883Z","shell.execute_reply.started":"2023-06-25T10:23:16.490694Z","shell.execute_reply":"2023-06-25T10:23:16.516862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use the quantiles of amoount of annotations to stratify\ndf_images_train, df_images_val = train_test_split(df_images, stratify=df_images['category_name'], \n                                                  test_size=TEST_SIZE,\n                                                  random_state=1234)\ndf_train = df_[df_['image_id'].isin(df_images_train['image_id'])]\ndf_val = df_[df_['image_id'].isin(df_images_val['image_id'])]","metadata":{"id":"2v7VvtTp0zRf","outputId":"82690b02-dd4b-4c1d-ed5d-78e30bf084a2","papermill":{"duration":0.064614,"end_time":"2021-11-08T00:11:49.722066","exception":false,"start_time":"2021-11-08T00:11:49.657452","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.519692Z","iopub.execute_input":"2023-06-25T10:23:16.520125Z","iopub.status.idle":"2023-06-25T10:23:16.541092Z","shell.execute_reply.started":"2023-06-25T10:23:16.520086Z","shell.execute_reply":"2023-06-25T10:23:16.540184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = HuBMAPDataset(TRAIN_PATH, df_train, resize=resize_factor, transforms=get_transform(train=True))\ndl_train = DataLoader(ds_train, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True,\n                      num_workers=2, collate_fn=lambda x: tuple(zip(*x)))\n\nds_val = HuBMAPDataset(TRAIN_PATH, df_val, resize=resize_factor, transforms=get_transform(train=False))\ndl_val = DataLoader(ds_val, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True,\n                    num_workers=2, collate_fn=lambda x: tuple(zip(*x)))","metadata":{"id":"kUcpAbdO0zRg","papermill":{"duration":0.138046,"end_time":"2021-11-08T00:11:49.891312","exception":false,"start_time":"2021-11-08T00:11:49.753266","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.543529Z","iopub.execute_input":"2023-06-25T10:23:16.544253Z","iopub.status.idle":"2023-06-25T10:23:16.832696Z","shell.execute_reply.started":"2023-06-25T10:23:16.544220Z","shell.execute_reply":"2023-06-25T10:23:16.831775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train model","metadata":{"id":"y8JNMn770zRg","papermill":{"duration":0.031418,"end_time":"2021-11-08T00:11:49.953839","exception":false,"start_time":"2021-11-08T00:11:49.922421","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## setup model","metadata":{"papermill":{"duration":0.03149,"end_time":"2021-11-08T00:11:50.015790","exception":false,"start_time":"2021-11-08T00:11:49.984300","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp ../input/cocopre/maskrcnn_resnet50_fpn_coco-bf2d0c1e.pth /root/.cache/torch/hub/checkpoints/maskrcnn_resnet50_fpn_coco-bf2d0c1e.pth","metadata":{"id":"VMaqdcNa0zRg","outputId":"c5ca31a5-6d8d-4639-e547-f44e5772c725","papermill":{"duration":4.3581,"end_time":"2021-11-08T00:11:54.404736","exception":false,"start_time":"2021-11-08T00:11:50.046636","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:16.833975Z","iopub.execute_input":"2023-06-25T10:23:16.834324Z","iopub.status.idle":"2023-06-25T10:23:20.319947Z","shell.execute_reply.started":"2023-06-25T10:23:16.834292Z","shell.execute_reply":"2023-06-25T10:23:20.318677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(num_classes, model_chkpt=None):    \n    if NORMALIZE:\n        model = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=True,\n                                                                   box_detections_per_img=BOX_DETECTIONS_PER_IMG,\n                                                                   image_mean=RESNET_MEAN,\n                                                                   image_std=RESNET_STD)\n    else:\n        model = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=True,\n                                                                   box_detections_per_img=BOX_DETECTIONS_PER_IMG)\n\n    # get the number of input features for the classifier\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    # replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes+1)\n    # now get the number of input features for the mask classifier\n    in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n    hidden_layer = 256\n    # and replace the mask predictor with a new one\n    model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask, hidden_layer, num_classes+1)\n    if model_chkpt:\n        model.load_state_dict(torch.load(model_chkpt, map_location=DEVICE))\n    return model\n# Get the Mask R-CNN model\n# The model does classification, bounding boxes and MASKs for individuals, all at the same time\n# We only care about MASKS\nmodel = get_model(len(cell_type_dict))\nmodel.to(DEVICE)\n\n# TODO: try removing this for\nfor param in model.parameters():\n    param.requires_grad = True\n    \nmodel.train();","metadata":{"id":"3Ds5dHex0zRh","papermill":{"duration":4.429689,"end_time":"2021-11-08T00:11:58.867693","exception":false,"start_time":"2021-11-08T00:11:54.438004","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:20.322907Z","iopub.execute_input":"2023-06-25T10:23:20.323673Z","iopub.status.idle":"2023-06-25T10:23:24.349712Z","shell.execute_reply.started":"2023-06-25T10:23:20.323635Z","shell.execute_reply":"2023-06-25T10:23:24.348754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training loop!","metadata":{"id":"RvawgUM30zRh","papermill":{"duration":0.030847,"end_time":"2021-11-08T00:11:58.931530","exception":false,"start_time":"2021-11-08T00:11:58.900683","status":"completed"},"tags":[]}},{"cell_type":"code","source":"params = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.SGD(params, lr=LEARNING_RATE, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY)\n#optimizer = torch.optim.Adam(params, lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)\n\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\nn_batches, n_batches_val = len(dl_train), len(dl_val)\nvalidation_mask_losses = []\n\nfor epoch in range(1, NUM_EPOCHS + 1):\n    print(f\"Starting epoch {epoch} of {NUM_EPOCHS}\")\n\n    time_start = time.time()\n    loss_accum = 0.0\n    loss_mask_accum = 0.0\n    loss_classifier_accum = 0.0\n    for batch_idx, (images, targets) in enumerate(dl_train, 1):\n    \n        # Predict\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        loss = sum(loss for loss in loss_dict.values())\n        \n        # Backprop\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Logging\n        loss_mask = loss_dict['loss_mask'].item()\n        loss_accum += loss.item()\n        loss_mask_accum += loss_mask\n        loss_classifier_accum += loss_dict['loss_classifier'].item()\n        \n        if batch_idx % 500 == 0:\n            print(f\"    [Batch {batch_idx:3d} / {n_batches:3d}] Batch train loss: {loss.item():7.3f}. Mask-only loss: {loss_mask:7.3f}.\")\n                        \n    if USE_SCHEDULER:\n        lr_scheduler.step()\n\n    # Train losses\n    train_loss = loss_accum / n_batches\n    train_loss_mask = loss_mask_accum / n_batches\n    train_loss_classifier = loss_classifier_accum / n_batches\n\n    # Validation\n    val_loss_accum = 0\n    val_loss_mask_accum = 0\n    val_loss_classifier_accum = 0\n    \n    with torch.no_grad():\n        for batch_idx, (images, targets) in enumerate(dl_val, 1):\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            val_loss_dict = model(images, targets)\n            val_batch_loss = sum(loss for loss in val_loss_dict.values())\n            val_loss_accum += val_batch_loss.item()\n            val_loss_mask_accum += val_loss_dict['loss_mask'].item()\n            val_loss_classifier_accum += val_loss_dict['loss_classifier'].item()\n\n    # Validation losses\n    val_loss = val_loss_accum / n_batches_val\n    val_loss_mask = val_loss_mask_accum / n_batches_val\n    val_loss_classifier = val_loss_classifier_accum / n_batches_val\n    elapsed = time.time() - time_start\n\n    validation_mask_losses.append(val_loss_mask)\n\n    torch.save(model.state_dict(), f\"pytorch_model-e{epoch}.bin\")\n    prefix = f\"[Epoch {epoch:2d}/{NUM_EPOCHS:2d}]\"\n    print(f\"{prefix} -- train mask loss: {train_loss_mask:7.3f}, classes loss {train_loss_classifier:7.3f}\")\n    print(f\"{prefix} -- val mask loss  : {val_loss_mask:7.3f}, classes loss {val_loss_classifier:7.3f}\")\n    print(f\"{prefix} -- train loss: {train_loss:7.3f}. val loss: {val_loss:7.3f} [{elapsed:.0f} secs]\")\n    print(\"---------------------------------------------------------------------------------------------\")","metadata":{"id":"52B16JCW0zRh","outputId":"9b9c5ad9-58c1-4d50-dd7c-79b18b1e57b7","papermill":{"duration":8400.749925,"end_time":"2021-11-08T02:31:59.712988","exception":false,"start_time":"2021-11-08T00:11:58.963063","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:23:24.356570Z","iopub.execute_input":"2023-06-25T10:23:24.357007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Analyze prediction results for train set","metadata":{"id":"MspyyJlP0zRh","papermill":{"duration":0.04241,"end_time":"2021-11-08T02:31:59.796080","exception":false,"start_time":"2021-11-08T02:31:59.753670","status":"completed"},"tags":[]}},{"cell_type":"code","source":"colors = [ 'Set1', 'Set3','Set3_r'] \nlegend = {0: 'blood_vessel',1: 'glomerulus', 2: 'unsure'}\nimg, targets = ds_train[17]\n\nfig, axs = plt.subplots(1, 3, figsize=(15, 5))\naxs[0].imshow(img.numpy().transpose((1,2,0)))\naxs[0].set_title('Image')\naxs[1].imshow(img.numpy().transpose((1,2,0)))\n# Loop over each mask and its label in the targets dictionary\nfor mask, lbl in zip(targets['masks'], targets['labels']):\n    mask = np.ma.masked_where(mask == 0, mask)\n    color = colors[lbl.item()-1]\n    axs[1].imshow(mask, cmap=color, alpha=0.8)\n    axs[1].set_title('Ground truth')\n    handles = []\n    for cl in legend:\n        color = colors[cl]\n        handles.append(mpatches.Patch(color=plt.colormaps.get_cmap(color)(0)))\n    axs[1].legend(handles, legend.values(), bbox_to_anchor=(0.55, 1.3), loc='upper left')\nmodel.eval()\nwith torch.no_grad():\n    preds = model([img.to(DEVICE)])[0]\nprint(len(preds[\"labels\"]))\naxs[2].imshow(img.numpy().transpose((1,2,0)))\nmasks,labels=get_filtered_masks(preds)\nfor mask ,lbl in zip(masks,labels):\n    color = colors[lbl-1]\n    mask = np.ma.masked_where(mask == 0, mask)\n    axs[2].imshow(mask, cmap=color, alpha=0.8)\n    axs[2].set_title('Predicted Masks')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get the best model","metadata":{"id":"Jon4MSmk0zRj","papermill":{"duration":0.057254,"end_time":"2021-11-08T02:32:09.513282","exception":false,"start_time":"2021-11-08T02:32:09.456028","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Epochs with their losses and IOU scores\nval_scores = pd.DataFrame()\nfor e, val_loss in enumerate(validation_mask_losses):\n    model_chk = f\"pytorch_model-e{e+1}.bin\"\n    print(\"Loading:\", model_chk)\n    model = get_model(len(cell_type_dict), model_chk)\n    model.load_state_dict(torch.load(model_chk))\n    model = model.to(DEVICE)\n    val_scores.loc[e,\"mask_loss\"] = val_loss\n    val_scores.loc[e,\"score\"] = get_score(ds_val, model)\n     \ndisplay(val_scores.sort_values(\"score\", ascending=False))\nbest_epoch = np.argmax(val_scores[\"score\"])\nprint(best_epoch+1)","metadata":{"id":"O0ejRcer0zRj","outputId":"0806ad69-abcf-440a-c9ad-d5a2a454abe4","papermill":{"duration":4225.264552,"end_time":"2021-11-08T03:42:34.835241","exception":false,"start_time":"2021-11-08T02:32:09.570689","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{"id":"bTNGfMuQ0zRi","papermill":{"duration":0.089032,"end_time":"2021-11-08T03:42:35.013518","exception":false,"start_time":"2021-11-08T03:42:34.924486","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Test Dataset and DataLoader","metadata":{"id":"tRSo-FPt0zRi","papermill":{"duration":0.08476,"end_time":"2021-11-08T03:42:35.185476","exception":false,"start_time":"2021-11-08T03:42:35.100716","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class HuBMAPTestDataset(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    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        image_path = os.path.join(self.image_dir, image_id + '.tif')\n        image = cv2.imread(image_path, cv2.IMREAD_COLOR)\n\n        if self.transforms is not None:\n            image, _ = self.transforms(image=image, target=None)\n        return {'image': image, 'image_id': image_id,'image_path':image_path}\n\n    def __len__(self):\n        return len(self.image_ids)","metadata":{"id":"ijZzdcHB0zRj","papermill":{"duration":0.107528,"end_time":"2021-11-08T03:42:35.379122","exception":false,"start_time":"2021-11-08T03:42:35.271594","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_test = HuBMAPTestDataset(TEST_PATH, transforms=get_transform(train=False))","metadata":{"id":"WbciaVrJ0zRj","outputId":"3ec69aa1-4136-41a1-b853-67d38654ef9a","papermill":{"duration":0.102778,"end_time":"2021-11-08T03:42:35.569848","exception":false,"start_time":"2021-11-08T03:42:35.467070","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import typing as t\nfrom pycocotools import _mask as coco_mask\nimport zlib\nimport base64\ndef encode_binary_mask(mask: np.ndarray) -> t.Text:\n    # check input mask --\n    if mask.dtype != bool:\n        raise ValueError(\n            \"encode_binary_mask expects a binary mask, received dtype == %s\" %\n            mask.dtype)\n\n    mask = np.squeeze(mask)\n    if len(mask.shape) != 2:\n        raise ValueError(\n            \"encode_binary_mask expects a 2d mask, received shape == %s\" %\n            mask.shape)\n\n    # convert input mask to expected COCO API input --\n    mask_to_encode = mask.reshape(mask.shape[0], mask.shape[1], 1)\n    mask_to_encode = mask_to_encode.astype(np.uint8)\n    mask_to_encode = np.asfortranarray(mask_to_encode)\n\n    # RLE encode mask --\n    encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n\n    # compress and base64 encoding --\n    binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n    base64_str = base64.b64encode(binary_str)\n    return base64_str","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_chk = f\"pytorch_model-e{best_epoch+1}.bin\"\nprint(\"Loading:\", model_chk)\nmodel = get_model(len(cell_type_dict))\nmodel.load_state_dict(torch.load(model_chk))\nmodel = model.to(DEVICE)\n\nfor param in model.parameters():\n    param.requires_grad = False\n\nmodel.eval();\n\nsubmission = []\nfor sample in ds_test:\n    img = sample['image']\n    image_id = sample['image_id']\n    image_path = sample['image_path']\n    h, w, _ = cv2.imread(image_path).shape\n    with torch.no_grad():\n        result = model([img.to(DEVICE)])[0]\n    \n    previous_masks = []\n    masks_use = []\n    labels_use = []\n    pred_string=\"\"\n    for i, mask in enumerate(result[\"masks\"]):\n        # Filter-out low-scoring results.\n        score = result[\"scores\"][i].cpu().item()\n        label = result[\"labels\"][i].cpu().item()\n        if score > min_score_dict[label]:\n            mask = mask.cpu().numpy()\n            # Keep only highly likely pixels\n            binary_mask = mask > mask_threshold_dict[label]\n            binary_mask = remove_overlapping_pixels(binary_mask, previous_masks)\n            masks_use.append(binary_mask)\n            labels_use.append(label)\n            previous_masks.append(binary_mask)\n            encoded = encode_binary_mask(binary_mask)\n            label=label-1\n            if label != 0: continue\n            if i == 0:\n                pred_string += f\"{int(label)} {score} {encoded.decode('utf-8')}\"\n            else:\n                pred_string += f\" {int(label)} {score} {encoded.decode('utf-8')}\"\n\n    submission.append((image_id,w,h, pred_string))\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"colors = [ 'Set1', 'Set3','Set3_r'] \nlegend = {0: 'blood_vessel',1: 'glomerulus', 2: 'unsure'}\nfig, axs = plt.subplots(1, 2, figsize=(10, 5))\naxs[0].imshow(img.numpy().transpose((1,2,0)))\naxs[0].set_title('Image')\naxs[1].imshow(img.numpy().transpose((1,2,0)))\n# Loop over each mask and its label in the targets dictionary\nfor mask, lbl in zip(masks_use, labels_use):\n    mask = np.ma.masked_where(mask[0] == 0, mask[0])\n    color = colors[lbl-1]\n    axs[1].imshow(mask, cmap=color, alpha=0.8)\n    axs[1].set_title('Predicted Masks')\n    handles = []\n    for cl in legend:\n        color = colors[cl]\n        handles.append(mpatches.Patch(color=plt.colormaps.get_cmap(color)(0)))\n    axs[1].legend(handles, legend.values(), bbox_to_anchor=(0.55, 1.3), loc='upper left')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub = pd.DataFrame(submission, columns=['id','height','width','prediction_string'])\ndf_sub.to_csv(\"submission.csv\", index=False)\ndf_sub.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}