{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# MaskRCNN Training | HuBMAP - Hacking the Human Vasculature Competition","metadata":{}},{"cell_type":"markdown","source":"[EDA](https://www.kaggle.com/code/khalilrejiba/hubmap-vasculature-eda-interactive)","metadata":{}},{"cell_type":"markdown","source":"## Utility Scripts","metadata":{}},{"cell_type":"code","source":"!pip -q install cython\n!pip install -qU pycocotools ","metadata":{"id":"DBIoe_tHTQgV","outputId":"a8aa5ac8-fe40-46a2-d0af-106d7da6e149","papermill":{"duration":152.133302,"end_time":"2023-07-28T21:51:31.623313","exception":false,"start_time":"2023-07-28T21:48:59.490011","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git clone https://github.com/pytorch/vision.git\n%cd /kaggle/working/vision\n!git checkout 59ec1df\n%cd /kaggle/working","metadata":{"papermill":{"duration":0.203577,"end_time":"2023-07-28T21:51:31.859444","exception":false,"start_time":"2023-07-28T21:51:31.655867","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%bash\ncp vision/references/detection/utils.py /kaggle/working\ncp vision/references/detection/engine.py /kaggle/working\ncp vision/references/detection/coco_eval.py /kaggle/working\ncp vision/references/detection/coco_utils.py /kaggle/working\ncp vision/references/detection/transforms.py /kaggle/working","metadata":{"papermill":{"duration":5.178028,"end_time":"2023-07-28T21:51:37.069741","exception":false,"start_time":"2023-07-28T21:51:31.891713","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Import Statements","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport json\nfrom pathlib import Path\nimport gc\nimport matplotlib.pyplot as plt\nimport PIL\nimport skimage\nfrom shapely.geometry import LinearRing as ShapelyContour\nfrom shapely.geometry import Polygon as ShapelyPolygon\nimport random\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch import nn\nimport torchvision\ntorchvision.disable_beta_transforms_warning()\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nfrom engine import train_one_epoch, evaluate\nimport utils\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom pycocotools import _mask as coco_mask\nimport base64\nimport zlib","metadata":{"id":"LDjuVFgexFfh","outputId":"f2e0c5c5-a8d0-4999-c94d-1660f153cd00","papermill":{"duration":5.759132,"end_time":"2023-07-28T21:51:42.861928","exception":false,"start_time":"2023-07-28T21:51:37.102796","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setting Global Variables","metadata":{}},{"cell_type":"code","source":"# For reproducibility\nseed = 123\nrandom.seed(seed) # Used in Albumentations\nnp.random.seed(seed)\ntorch.manual_seed(seed);","metadata":{"papermill":{"duration":0.045883,"end_time":"2023-07-28T21:51:42.941059","exception":false,"start_time":"2023-07-28T21:51:42.895176","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_DIR = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train\"\nTEST_DIR = \"/kaggle/input/hubmap-hacking-the-human-vasculature/test\"\nANNOT_PATH = \"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\"\nWSI_TILE_CSV = \"/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv\"\nIMG_SIZE = 512","metadata":{"papermill":{"duration":0.041576,"end_time":"2023-07-28T21:51:43.015002","exception":false,"start_time":"2023-07-28T21:51:42.973426","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hidden_test_ids = [path.stem for path in Path(TEST_DIR).glob(\"*.tif\")]","metadata":{"papermill":{"duration":0.045584,"end_time":"2023-07-28T21:51:43.093619","exception":false,"start_time":"2023-07-28T21:51:43.048035","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_names = [\"background\", \"blood_vessel\", \"glomerulus\", \"unsure\"]\n\nwsi_df = pd.read_csv(WSI_TILE_CSV)\n# Tiles with Annotations\ndataset12_ids = wsi_df[wsi_df[\"dataset\"] != 3].id.values.tolist() \n# Tiles with Expert Reviewed Annotations\ndataset1_ids = wsi_df[wsi_df[\"dataset\"] == 1].id.values.tolist()\n# Tiles with Sparse Annotations\ndataset2_ids = wsi_df[wsi_df[\"dataset\"] == 2].id.values.tolist()\n# Tiles without Annotations\ndataset3_ids = wsi_df[wsi_df[\"dataset\"] == 3].id.values.tolist() ","metadata":{"papermill":{"duration":0.0722,"end_time":"2023-07-28T21:51:43.198697","exception":false,"start_time":"2023-07-28T21:51:43.126497","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing","metadata":{"papermill":{"duration":0.031884,"end_time":"2023-07-28T21:51:43.263289","exception":false,"start_time":"2023-07-28T21:51:43.231405","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def load_annotations(annotations_path=ANNOT_PATH):\n    \n    annotations = {}\n    duplicates_ids = []\n\n    with open(annotations_path, \"r\") as f:\n        for line in f:\n            entry = json.loads(line)\n            identifier = entry[\"id\"]\n            annotations_list = entry[\"annotations\"]\n\n            revised_annotations = [] # [[label_str, polygon_list], ...]\n\n            for annotation in annotations_list:\n                label = annotation[\"type\"]\n                polygon = annotation[\"coordinates\"][0]\n\n                # Remove Duplicates\n                if not any(ShapelyContour(p).equals(ShapelyContour(polygon)) and l==label for l, p in revised_annotations):\n                    revised_annotations.append([label, polygon])\n\n            if len(annotations_list) != len(revised_annotations):\n                duplicates_ids.append(identifier)\n\n            annotations[identifier] = revised_annotations\n    return annotations, duplicates_ids\n\ndef clean_annotations(annotations):\n    \n    new_annotations = {}\n    anomalies_ids = []\n\n    for identifier in annotations.keys():\n\n        annotations_list = annotations[identifier]\n        revised_annotations = []\n        skip = False\n\n        for (label, polygon) in annotations_list:\n\n            for i, (l, p) in enumerate(revised_annotations):\n                # New Polygon Inside Old Polygon -> Skip New\n                if ShapelyPolygon(p).buffer(0).contains(ShapelyPolygon(polygon).buffer(0)):\n                    skip = True\n                    break\n                # Old Polygon Inside New Polygon -> Delete Old\n                if ShapelyPolygon(p).buffer(0).within(ShapelyPolygon(polygon).buffer(0)):\n                    revised_annotations.pop(i)\n            if skip:\n                continue\n            revised_annotations.append([label, polygon])\n\n        if len(annotations_list) != len(revised_annotations):\n            anomalies_ids.append(identifier)\n\n        new_annotations[identifier] = revised_annotations\n\n    return new_annotations, anomalies_ids ","metadata":{"papermill":{"duration":0.049625,"end_time":"2023-07-28T21:51:43.345997","exception":false,"start_time":"2023-07-28T21:51:43.296372","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"annotations, _ = load_annotations()\nannotations, _ = clean_annotations(annotations)\nassert len(annotations) == len(dataset12_ids)","metadata":{"papermill":{"duration":66.773746,"end_time":"2023-07-28T21:52:50.151983","exception":false,"start_time":"2023-07-28T21:51:43.378237","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_gt_instance_seg(identifiers, annotations, class_names, keep_only=None, erode=0):\n    \"\"\"\n    Transform polygon annotations into instance masks\n    \"\"\"   \n    num_classes = len(class_names) - 1\n    if keep_only is None: \n        mask_shape = (IMG_SIZE, IMG_SIZE, num_classes)\n    else:\n        mask_shape = (IMG_SIZE, IMG_SIZE)\n    new_ids, masks, polygons = [], [], []\n        \n    for identifier in identifiers:\n        \n        labeled_mask = np.zeros(mask_shape, dtype=\"uint8\")\n        overlap = np.zeros(mask_shape, dtype=\"uint8\")\n        polygons_list = [] # [[label_int, polygon_list], ...]\n        \n        num_objects = 0\n        \n        for (label, polygon) in annotations[identifier]:\n            if label == keep_only or keep_only is None:\n                num_objects += 1\n                label_int = class_names.index(label)\n                polygons_list.append([label_int, polygon])\n                mask = skimage.draw.polygon2mask(mask_shape[:2], polygon).T\n                if erode:\n                    mask = skimage.morphology.erosion(mask, skimage.morphology.square(7))\n                    \n                if keep_only is None: \n                    intersection = (labeled_mask * np.stack([mask]*num_classes, axis=-1)) > 0\n                else:\n                    intersection = (labeled_mask * mask) > 0\n                    \n                if intersection.sum():\n                    overlap += intersection\n                    \n                if keep_only is None: \n                    labeled_mask[:,:,label_int-1] += (num_objects * mask).astype(\"uint8\")\n                else:\n                    labeled_mask += (num_objects * mask).astype(\"uint8\")\n        \n        # Overlap is treated as background\n        labeled_mask[overlap > 0] = 0\n        \n        if labeled_mask.sum() == 0:\n            continue\n            \n        new_ids.append(identifier)\n        masks.append(labeled_mask)\n        polygons.append(polygons_list)\n    \n    masks = np.array(masks) # (N, H, W) or (N, H, W, C)\n    \n    return new_ids, masks, polygons","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def shuffle_sequence(seq):\n    array = np.array(seq)\n    np.random.shuffle(array)\n    return array.tolist()\n\ndef retrieve_identifiers(wsi_df, source, shuffle=True):\n    array = wsi_df[wsi_df[\"source\"]==source].id.values.tolist()\n    if shuffle:    \n        return shuffle_sequence(array)\n    return array\n\nwsi_df = pd.read_csv(WSI_TILE_CSV)\nwsi_df[\"source\"] = wsi_df[[\"dataset\", \"source_wsi\"]].apply(tuple, axis=1)\n\nsource11 = retrieve_identifiers(wsi_df, (1, 1))\nsource12 = retrieve_identifiers(wsi_df, (1, 2))\nsource21 = retrieve_identifiers(wsi_df, (2, 1))\nsource22 = retrieve_identifiers(wsi_df, (2, 2))\nsource23 = retrieve_identifiers(wsi_df, (2, 3))\nsource24 = retrieve_identifiers(wsi_df, (2, 4))\n\nvalid_ids = source11[:50] + source12[:50]\ntrain_ids = source11[50:] + source12[50:] + source21 + source22 + source23 + source24","metadata":{"papermill":{"duration":0.133628,"end_time":"2023-07-28T21:52:50.395282","exception":false,"start_time":"2023-07-28T21:52:50.261654","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HuBMAPVasculatureDataset(Dataset):\n    \n    def __init__(self, ids, masks=None, root=TRAIN_DIR, transforms=None):\n        self.root = root\n        self.ext = \"tif\"\n        self.transforms = transforms\n        self.ids = ids\n        self.masks = masks\n        self.multiclass = masks.ndim == 4\n\n    def __getitem__(self, idx):\n        img_path = Path(self.root) / f\"{self.ids[idx]}.{self.ext}\"\n        img = PIL.Image.open(img_path).convert(\"RGB\")\n        img = np.array(img)\n        mask = self.masks[idx]\n        obj_ids = np.unique(mask)\n        obj_ids = obj_ids[1:]\n        \n        if self.multiclass:\n            masks = (mask == obj_ids[:, None, None, None])\n            labels = (np.argmax(np.any(masks, axis=(1, 2)), axis=-1) + 1).astype(np.int64)\n            masks = np.any(masks, axis=-1).astype(\"uint8\")\n        else:\n            masks = (mask == obj_ids[:, None, None]).astype(\"uint8\")\n            labels = np.ones((num_objs,), dtype=np.int64)\n\n        num_objs = len(obj_ids)\n        boxes = []\n        for i in range(num_objs):\n            pos = np.where(masks[i])\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            boxes.append([xmin, ymin, xmax, ymax])\n        boxes = np.array(boxes, dtype=np.float32)\n        \n        image_id = torch.tensor(idx)\n        area = torch.as_tensor((boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0]))\n        iscrowd = torch.zeros((num_objs,), dtype=torch.int64)\n\n        target = {}\n        target[\"boxes\"] = boxes\n        target[\"labels\"] = labels\n        target[\"masks\"] = masks\n        target[\"image_id\"] = image_id\n        target[\"area\"] = area # Can't rely on area after augmentation\n        target[\"iscrowd\"] = iscrowd\n\n        if self.transforms is not None:\n            target[\"masks\"] = np.transpose(target[\"masks\"], (1, 2, 0))\n            \n            data = self.transforms(image=img, bboxes=target[\"boxes\"], mask=target[\"masks\"], class_labels=target[\"labels\"])\n\n            img = data[\"image\"] / 255.\n            target[\"masks\"] = data[\"mask\"]\n            target[\"boxes\"] = torch.as_tensor(data[\"bboxes\"], dtype=torch.float32)\n            target[\"labels\"] = torch.as_tensor(data[\"class_labels\"], dtype=torch.int64)\n\n        return img, target\n\n    def __len__(self):\n        return len(self.ids)","metadata":{"id":"mTgWtixZTs3X","papermill":{"duration":0.049861,"end_time":"2023-07-28T21:52:50.477532","exception":false,"start_time":"2023-07-28T21:52:50.427671","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HuBMAPVasculatureTestset(Dataset):\n    \n    def __init__(self, ids, root=TEST_DIR, transforms=None):\n        self.root = root\n        self.ext = \"tif\"\n        self.ids = ids\n        self.transforms = transforms\n\n    def __getitem__(self, idx):\n        img_path = Path(self.root) / f\"{self.ids[idx]}.{self.ext}\"\n        img = PIL.Image.open(img_path).convert(\"RGB\")\n        img = np.array(img)\n        if self.transforms is not None:\n            img = self.transforms(image=img)[\"image\"]\n        return img / 255\n\n    def __len__(self):\n        return len(self.ids)","metadata":{"papermill":{"duration":0.041982,"end_time":"2023-07-28T21:52:50.551434","exception":false,"start_time":"2023-07-28T21:52:50.509452","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_instance_segmentation_model(num_classes=2, hidden_layer=256):\n    model = torchvision.models.detection.maskrcnn_resnet50_fpn(weights=\"DEFAULT\")\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n    in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels\n    model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask,\n                                                       hidden_layer,\n                                                       num_classes)\n    return model","metadata":{"id":"YjNHjVMOyYlH","papermill":{"duration":0.042421,"end_time":"2023-07-28T21:52:50.625693","exception":false,"start_time":"2023-07-28T21:52:50.583272","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(train=False, s=IMG_SIZE):\n    transforms = [\n        A.Resize(height=s, width=s),\n    ]\n    if train:\n        transforms.extend([\n#             A.ShiftScaleRotate(p=0.5, shift_limit=0.05, scale_limit=(-0.01, 0.1), rotate_limit=5, border_mode=4),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.ColorJitter (p=0.5, brightness=0.25, contrast=0.25, saturation=0.25, hue=0.1),\n        ])  \n    transforms.append(ToTensorV2(transpose_mask=True)) \n    \n    transforms = A.Compose(\n        transforms,\n        bbox_params=A.BboxParams(format='pascal_voc', label_fields=['class_labels'], \n                                 min_area=0, min_visibility=0)\n    )\n    \n    return transforms\n\n# With ShiftScaleRotate, you risk having images with no objects","metadata":{"id":"l79ivkwKy357","papermill":{"duration":0.042223,"end_time":"2023-07-28T21:52:50.699940","exception":false,"start_time":"2023-07-28T21:52:50.657717","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms_pred(s=IMG_SIZE):\n    transforms = [\n        A.Resize(height=s, width=s),\n        ToTensorV2()\n    ]    \n    transforms = A.Compose(transforms)\n    return transforms","metadata":{"papermill":{"duration":0.039653,"end_time":"2023-07-28T21:52:50.771552","exception":false,"start_time":"2023-07-28T21:52:50.731899","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" Under the hood, the MaskRCNN model uses the following preprocessing steps:   \n  (transform): GeneralizedRCNNTransform(  \n      Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  \n      Resize(min_size=(800,), max_size=1333, mode='bilinear')  \n  )","metadata":{"papermill":{"duration":0.031411,"end_time":"2023-07-28T21:52:50.835047","exception":false,"start_time":"2023-07-28T21:52:50.803636","status":"completed"},"tags":[]}},{"cell_type":"code","source":"ids, masks_gt, polygons_gt = prepare_gt_instance_seg(dataset12_ids,\n                                                     annotations,\n                                                     class_names)","metadata":{"papermill":{"duration":84.401475,"end_time":"2023-07-28T21:54:15.268360","exception":false,"start_time":"2023-07-28T21:52:50.866885","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK_DIR = Path(\"/kaggle/working\") / \"masks\"\nMASK_DIR.mkdir(parents=True, exist_ok=True)\nfor identifier, mask in zip(ids, masks_gt):\n    path = MASK_DIR / f\"{identifier}.npy\"\n    np.save(path, mask)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_indices = [ids.index(i) for i in train_ids if i in ids]\nvalid_indices = [ids.index(i) for i in valid_ids if i in ids]\n\ntrain_dataset = HuBMAPVasculatureDataset(ids, masks_gt, transforms=get_transforms(train=True))\nvalid_dataset = HuBMAPVasculatureDataset(ids, masks_gt, transforms=get_transforms(train=False))\ntrain_dataset = torch.utils.data.Subset(train_dataset, train_indices)\nvalid_dataset = torch.utils.data.Subset(valid_dataset, valid_indices)","metadata":{"papermill":{"duration":0.128963,"end_time":"2023-07-28T21:54:15.431408","exception":false,"start_time":"2023-07-28T21:54:15.302445","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"papermill":{"duration":0.031281,"end_time":"2023-07-28T21:54:15.495057","exception":false,"start_time":"2023-07-28T21:54:15.463776","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_loader = torch.utils.data.DataLoader(\n    train_dataset, batch_size=8, shuffle=True, num_workers=0,\n    collate_fn=utils.collate_fn\n)\n\nvalid_loader = torch.utils.data.DataLoader(\n    valid_dataset, batch_size=16, shuffle=False, num_workers=0,\n    collate_fn=utils.collate_fn\n)","metadata":{"papermill":{"duration":0.040744,"end_time":"2023-07-28T21:54:15.567457","exception":false,"start_time":"2023-07-28T21:54:15.526713","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nnum_epochs = 10\nSAVE_PATH = f\"maskrcnn_resnet50_fpn_finetune_{num_epochs}epochs.pth\"\n\ntrain_success = False\nnum_try = 0\nwhile not train_success and num_try < 10:\n    try:\n        model = get_instance_segmentation_model(num_classes=len(class_names))\n        model.to(device)\n\n        params = [p for p in model.parameters() if p.requires_grad]\n        optimizer = torch.optim.SGD(params,\n                                    lr=0.005,\n                                    momentum=0.9,\n                                    weight_decay=0.0005)\n\n        lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer,\n                                                       step_size=3,\n                                                       gamma=0.1)\n    \n        mean_ap_best = 0\n\n        for epoch in range(num_epochs):\n            torch.cuda.empty_cache()\n            train_one_epoch(model, optimizer, train_loader, device, epoch, print_freq=500)\n            if epoch >= (num_epochs//2):\n                lr_scheduler.step()\n            torch.cuda.empty_cache()\n            coco_evaluator = evaluate(model, valid_loader, device=device)\n            mean_ap = coco_evaluator.coco_eval[\"segm\"].stats[0]\n            if mean_ap > mean_ap_best:\n                mean_ap_best = mean_ap\n                print(\"Saving checkpoint: mAP = \", mean_ap_best)\n                torch.save(model.state_dict(), SAVE_PATH)\n        train_success = True\n        \n    except (RuntimeError, ConnectionRefusedError, EOFError) as e:\n        print(e)\n        num_try += 1\n        torch.cuda.empty_cache()\n        gc.collect()\n        pass\n    \nif train_success:\n    print(\"Finished training succesfully!\")","metadata":{"id":"zoenkCj18C4h","papermill":{"duration":4870.700029,"end_time":"2023-07-28T23:15:26.299371","exception":false,"start_time":"2023-07-28T21:54:15.599342","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Post-processing","metadata":{"papermill":{"duration":0.042209,"end_time":"2023-07-28T23:15:26.844725","exception":false,"start_time":"2023-07-28T23:15:26.802516","status":"completed"},"tags":[]}},{"cell_type":"code","source":"model.load_state_dict(torch.load(SAVE_PATH))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def apply_nms_box(boxes, scores, iou_thres=0.5):\n    \"\"\"\n    Non Maximum Suppression using bounding boxes\n    \"\"\"\n    x0 = boxes[:, 0]\n    y0 = boxes[:, 1]\n    x1 = boxes[:, 2]\n    y1 = boxes[:, 3]\n\n    areas = (x1 - x0) * (y1 - y0)\n\n    indices = scores.argsort(descending=True)\n \n    revised_indices = []\n    \n    while len(indices) > 0:\n        \n        idx = indices[0]\n        indices = indices[1:]\n        revised_indices.append(idx.item())\n        \n        xx0 = torch.max(x0[indices], x0[idx])\n        yy0 = torch.max(y0[indices], y0[idx])\n        xx1 = torch.min(x1[indices], x1[idx])\n        yy1 = torch.min(y1[indices], y1[idx])\n         \n        w = torch.clamp(xx1 - xx0, min=0)\n        h = torch.clamp(yy1 - yy0, min=0)\n        intersection = w * h\n        union = areas[indices] - intersection\n        IoU = intersection / union\n\n        indices = indices[IoU < iou_thres]\n\n    return revised_indices","metadata":{"papermill":{"duration":0.055673,"end_time":"2023-07-28T23:15:26.943229","exception":false,"start_time":"2023-07-28T23:15:26.887556","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def apply_nms_mask(masks, scores, iou_thres=0.5):\n    \"\"\"\n    Non Maximum Suppression using masks\n    \"\"\"\n    indices = scores.argsort(descending=True)\n \n    revised_indices = []\n    \n    while len(indices) > 0:\n        \n        idx = indices[0]\n        indices = indices[1:]\n        revised_indices.append(idx.item())\n         \n        intersection = masks[indices] * masks[idx]\n        union = masks[indices] + masks[idx]\n        union = union.clip(0, 1)\n\n        IoU = intersection.sum(dim=(1, 2, 3)) / union.sum(dim=(1, 2, 3))\n\n        indices = indices[IoU < iou_thres]\n\n    return revised_indices","metadata":{"papermill":{"duration":0.054745,"end_time":"2023-07-28T23:15:27.040892","exception":false,"start_time":"2023-07-28T23:15:26.986147","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_predictions(model, img, label=1, conf_thres=0.5, iou_thres=0.6, nms_mask=True):\n    model.eval()\n    with torch.no_grad():\n        prediction = model([img.to(device)])\n    prediction = prediction[0]\n    \n    conf_filter = prediction['scores'] > conf_thres\n    label_filter = prediction['labels'] == label\n    for k, v in prediction.items():    \n        prediction[k] = v[conf_filter * label_filter]\n    \n    masks = prediction['masks']\n    boxes = prediction['boxes']\n    scores = prediction['scores']\n    \n    # Soft Masks to Binary Masks\n    masks = (masks >  0.5).byte()\n    \n    # Non Maximum Suppression\n    if nms_mask:\n        indices = apply_nms_mask(masks, scores, iou_thres)\n    else:\n        indices = apply_nms_box(boxes, scores, iou_thres)\n\n    return masks[indices].cpu().numpy(), boxes[indices].cpu().numpy(), scores[indices].cpu().numpy()","metadata":{"papermill":{"duration":0.054528,"end_time":"2023-07-28T23:15:27.138381","exception":false,"start_time":"2023-07-28T23:15:27.083853","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = HuBMAPVasculatureTestset(hidden_test_ids, transforms=get_transforms_pred())\n\nimg = test_dataset[np.random.randint(len(test_dataset))]\n\nplt.imshow(img.permute(1, 2, 0))\n\nmasks, _, _ = get_predictions(model, img)\n\nfor mask in masks:\n    mask = mask.squeeze()\n    polygon = skimage.measure.find_contours(mask)[0]\n    yy, xx = zip(*polygon)\n    plt.fill(xx, yy, facecolor='none', edgecolor=np.random.random(size=3), linewidth=2)","metadata":{"papermill":{"duration":0.739024,"end_time":"2023-07-28T23:15:27.922802","exception":false,"start_time":"2023-07-28T23:15:27.183778","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = HuBMAPVasculatureTestset(dataset3_ids, root=TRAIN_DIR, transforms=get_transforms_pred())\n\nimg = test_dataset[np.random.randint(len(test_dataset))]\n\nplt.imshow(img.permute(1, 2, 0))\n\nmasks, _, _ = get_predictions(model, img)\n\nfor mask in masks:\n    mask = mask.squeeze()\n    polygon = skimage.measure.find_contours(mask)[0]\n    yy, xx = zip(*polygon)\n    plt.fill(xx, yy, facecolor='none', edgecolor=np.random.random(size=3), linewidth=2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"glomerus_ids = [\"d40ccc91c77d\", \"26d850f54795\", \"4b19cae08534\"] # Dataset 3 (unseen)\ntest_dataset = HuBMAPVasculatureTestset(glomerus_ids, root=TRAIN_DIR, transforms=get_transforms_pred())\n\ncolors = [\"red\", \"green\", \"blue\"]\n\nfig, axs = plt.subplots(1, 3)\n\nfor idx, image_id in enumerate(glomerus_ids):\n    ax = axs[idx]\n    ax.set_title(image_id)\n    \n    img = test_dataset[idx]\n    ax.imshow(img.permute(1, 2, 0))\n    \n    for label_id in range(1, 4):\n        masks, _, _ = get_predictions(model, img, label=label_id)\n        for mask in masks:\n            mask = mask.squeeze()\n            polygon = skimage.measure.find_contours(mask)[0]\n            yy, xx = zip(*polygon)\n            ax.fill(xx, yy, facecolor='none', edgecolor=colors[label_id-1], linewidth=2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_grid(dataset, tile_ids, masks, polygons, class_names, model=None, nrows=2, ncols=2, divide_by=2, dpi=100):\n    figsize = (((IMG_SIZE/divide_by)*ncols)/dpi, ((IMG_SIZE/divide_by)*nrows)/dpi)\n    fig, axs = plt.subplots(nrows, ncols, figsize=figsize, dpi=dpi)\n    fig.subplots_adjust(left=0, bottom=0, right=1, top=1, wspace=0, hspace=0)\n    axs = axs.flatten()\n    for i, ax in enumerate(axs):\n        identifier = tile_ids[i]\n        img = dataset[i]\n        mask = masks[i]\n        polygons_list = polygons[i]\n        ax.imshow(img.cpu().permute(1, 2, 0))\n        ax.imshow(mask, alpha=0.25, cmap=\"turbo\", interpolation=\"nearest\")\n        if model is not None:\n            for mask in get_predictions(model, img)[0]:\n                mask = mask.squeeze()\n                polygon = skimage.measure.find_contours(mask)[0]\n                yy, xx = zip(*polygon)\n                ax.fill(xx, yy, facecolor='none', edgecolor=\"blue\", linewidth=1)\n        for label, polygon in polygons_list:\n            xx, yy = zip(*polygon)\n            if label == 1:\n                ax.fill(xx, yy, facecolor='none', edgecolor=\"lightgreen\", linewidth=1)\n            elif label == 2:\n                ax.fill(xx, yy, facecolor='none', edgecolor=\"red\", linewidth=1)\n            else:\n                ax.fill(xx, yy, facecolor='none', edgecolor=\"lightgreen\", linestyle=\"--\", linewidth=1)\n        ax.set_xticks([])\n        ax.set_yticks([])\n        ax.set_aspect(\"equal\")\n    fig.show()","metadata":{"papermill":{"duration":0.062592,"end_time":"2023-07-28T23:15:28.032944","exception":false,"start_time":"2023-07-28T23:15:27.970352","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = HuBMAPVasculatureTestset(ids, transforms=get_transforms_pred(), root=TRAIN_DIR)","metadata":{"papermill":{"duration":0.056841,"end_time":"2023-07-28T23:15:28.137195","exception":false,"start_time":"2023-07-28T23:15:28.080354","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display Ground Truth\nshow_grid(dataset, ids, masks_gt, polygons_gt, class_names, divide_by=1.5)","metadata":{"papermill":{"duration":1.059081,"end_time":"2023-07-28T23:15:29.246742","exception":false,"start_time":"2023-07-28T23:15:28.187661","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Display Ground Truth and Model Predictions\nshow_grid(dataset, ids, masks_gt, polygons_gt, class_names, model, divide_by=1.5)","metadata":{"papermill":{"duration":1.869486,"end_time":"2023-07-28T23:15:31.180661","exception":false,"start_time":"2023-07-28T23:15:29.311175","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  Submission","metadata":{"papermill":{"duration":0.078658,"end_time":"2023-07-28T23:15:31.346927","exception":false,"start_time":"2023-07-28T23:15:31.268269","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def encode_mask(mask):\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    encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n    binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n    base64_str = base64.b64encode(binary_str)\n    return base64_str","metadata":{"papermill":{"duration":0.088661,"end_time":"2023-07-28T23:15:31.513693","exception":false,"start_time":"2023-07-28T23:15:31.425032","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_list = []\nfor i, test_id in enumerate(hidden_test_ids):\n    prediction_string = \"\"\n    \n    img = test_dataset[i]\n    _, h, w = img.shape\n    masks, _, scores = get_predictions(model, img)\n    masks_glomerulus, _, _ = get_predictions(model, img, label=2)\n    for mask, score in zip(masks, scores):\n        for glm_msk in masks_glomerulus:\n            # Test if vessel is mainly outside glomerulus\n            intersection = mask * glm_msk\n            if intersection.sum() < 0.5 * mask.sum():\n                mask = mask.squeeze()\n                mask_str = encode_mask(mask).decode('UTF-8')\n                prediction_string += f\"0 {score} {mask_str} \"\n                entry = {\n                    \"id\": test_id,\n                    \"height\": h,\n                    \"width\": w,\n                    \"prediction_string\": prediction_string,\n                }\n                submission_list.append(entry)","metadata":{"papermill":{"duration":0.226144,"end_time":"2023-07-28T23:15:31.817608","exception":false,"start_time":"2023-07-28T23:15:31.591464","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.DataFrame(submission_list)\nsubmission_df = submission_df.set_index('id')\nsubmission_df.to_csv(\"submission.csv\")","metadata":{"papermill":{"duration":0.092792,"end_time":"2023-07-28T23:15:31.991159","exception":false,"start_time":"2023-07-28T23:15:31.898367","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.head()","metadata":{"papermill":{"duration":0.096516,"end_time":"2023-07-28T23:15:32.165190","exception":false,"start_time":"2023-07-28T23:15:32.068674","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Modified from: https://pytorch.org/tutorials/intermediate/torchvision_tutorial.html","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"}}