{"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":"> **[Mask RCNN Source Code](https://github.com/pytorch/vision/blob/main/torchvision/models/detection/mask_rcnn.py)**","metadata":{}},{"cell_type":"code","source":"from torchvision.models.detection import maskrcnn_resnet50_fpn, MaskRCNN_ResNet50_FPN_Weights\nfrom glob import glob\nimport pandas as pd\nimport json\nimport numpy as np\nimport matplotlib as mpl\nimport matplotlib.pyplot as plt\nimport torch\nimport cv2\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.transforms import transforms\nfrom PIL import Image\nimport os\nfrom skimage.draw import polygon2mask\nfrom matplotlib.patches import Polygon\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\nimport torch.optim as optim\nimport time\nfrom tqdm import tqdm","metadata":{"papermill":{"duration":0.028281,"end_time":"2023-05-27T17:49:59.290563","exception":false,"start_time":"2023-05-27T17:49:59.262282","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-08-10T10:06:44.850288Z","iopub.execute_input":"2023-08-10T10:06:44.850884Z","iopub.status.idle":"2023-08-10T10:06:50.045194Z","shell.execute_reply.started":"2023-08-10T10:06:44.850844Z","shell.execute_reply":"2023-08-10T10:06:50.044225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Path Configs","metadata":{}},{"cell_type":"code","source":"base_path = \"/kaggle/input/hubmap-hacking-the-human-vasculature/\"\nimages_folder = base_path + \"/train/\"\nlabels_path = base_path + \"/polygons.jsonl\"\nmetadata_path = base_path + \"/tile_meta.csv\"","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:06:50.047320Z","iopub.execute_input":"2023-08-10T10:06:50.048249Z","iopub.status.idle":"2023-08-10T10:06:50.053553Z","shell.execute_reply.started":"2023-08-10T10:06:50.048213Z","shell.execute_reply":"2023-08-10T10:06:50.052781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = []\nwith open(labels_path, 'r') as file:\n    for line in file:\n        item = json.loads(line)\n        data.append(item)\n\ndf = pd.DataFrame(data)\ndf.head(1)","metadata":{"papermill":{"duration":3.735182,"end_time":"2023-05-27T17:50:03.408175","exception":false,"start_time":"2023-05-27T17:49:59.672993","status":"completed"},"scrolled":true,"tags":[],"execution":{"iopub.status.busy":"2023-08-10T10:06:50.055078Z","iopub.execute_input":"2023-08-10T10:06:50.055853Z","iopub.status.idle":"2023-08-10T10:06:54.988862Z","shell.execute_reply.started":"2023-08-10T10:06:50.055706Z","shell.execute_reply":"2023-08-10T10:06:54.987802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Creation","metadata":{"papermill":{"duration":0.007215,"end_time":"2023-05-27T17:50:04.103073","exception":false,"start_time":"2023-05-27T17:50:04.095858","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class HubMap_Dataset(Dataset):\n    \n    def is_dataset_1_file(self, metadata_file, file_name):\n        df_tile = pd.read_csv(metadata_file)\n        dataset_1_files = df_tile[df_tile[\"dataset\"]==1]\n        files = dataset_1_files[\"id\"].tolist()\n        if file_name in files:\n            return True\n        else:\n            return False\n    \n    def __init__(self, img_path, labels_file, metadata_file):\n        self.json_labels = []    \n        self.metadata_file = metadata_file\n        self.image_dir = img_path\n        with open(labels_file, 'r') as json_file:\n            for line in json_file:\n                temp = json.loads(line)\n                flag = self.is_dataset_1_file(self.metadata_file, temp[\"id\"])\n                if flag:\n                    self.json_labels.append(json.loads(line))\n                \n    def __len__(self):\n        return len(self.json_labels)\n\n    def __getitem__(self, idx):\n        image_path = self.image_dir + \"/\" + self.json_labels[idx][\"id\"] + \".tif\"\n        img_id = self.json_labels[idx][\"id\"]\n        image = Image.open(image_path)\n        \n        \n        boxes = []\n        labels = []\n        masks = []\n        all_masked = np.zeros((512, 512), dtype=np.float32)\n        all_filled_masks = np.zeros((512, 512), dtype=np.float32)\n        filled_masks = []\n        for annot in self.json_labels[idx][\"annotations\"]:\n            if annot['type'] == \"blood_vessel\":\n                coordinates = np.array(annot[\"coordinates\"])\n                temp_filled_mask = np.zeros((512, 512), dtype=np.float32)\n                all_filled_masks = cv2.fillPoly(all_filled_masks, [coordinates], 1)\n                filled_masks.append(np.array(cv2.fillPoly(temp_filled_mask, [coordinates], 1), dtype=np.float32))\n                for cord in coordinates:\n                    mask = np.zeros((512, 512), dtype=np.float32)\n                    x, y = np.array([i[1] for i in cord]), np.asarray([i[0] for i in cord])\n                    min_x, min_y, max_x, max_y = min(x), min(y), max(x), max(y)\n                    boxes.append([min_y, min_x, max_y, max_x])\n                    labels.append(1)\n                    mask[x,y] = 1\n                    all_masked[x,y] = 1\n                    masks.append(np.array(mask, dtype=np.float32))\n        \n        image = np.array(image)\n        image = torch.tensor(np.array(image/255.0), dtype=torch.float32)\n        masks = torch.tensor(np.array(masks), dtype=torch.float32)\n        filled_masks = torch.tensor(np.array(filled_masks), dtype=torch.float32)\n        boxes = torch.tensor(np.array(boxes), dtype=torch.float32)\n        labels = torch.tensor(np.array(labels), dtype=torch.int64)\n        \n        return image, filled_masks, boxes, labels, all_filled_masks, masks, all_masked, img_id\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:06:54.991423Z","iopub.execute_input":"2023-08-10T10:06:54.991804Z","iopub.status.idle":"2023-08-10T10:06:55.011115Z","shell.execute_reply.started":"2023-08-10T10:06:54.991769Z","shell.execute_reply":"2023-08-10T10:06:55.009817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = HubMap_Dataset(img_path=images_folder, labels_file=labels_path, metadata_file=metadata_path)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:06:55.012368Z","iopub.execute_input":"2023-08-10T10:06:55.013258Z","iopub.status.idle":"2023-08-10T10:07:13.598331Z","shell.execute_reply.started":"2023-08-10T10:06:55.013232Z","shell.execute_reply":"2023-08-10T10:07:13.597180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_set, val_set = torch.utils.data.random_split(dataset, [360, 62])","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:07:54.402414Z","iopub.execute_input":"2023-08-10T10:07:54.402921Z","iopub.status.idle":"2023-08-10T10:07:54.411346Z","shell.execute_reply.started":"2023-08-10T10:07:54.402879Z","shell.execute_reply":"2023-08-10T10:07:54.410319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Random Plots","metadata":{}},{"cell_type":"code","source":"sample = dataset[0]\nall_masks = sample[4]\nimage = sample[0]\nboxes = sample[2]","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:07:55.594465Z","iopub.execute_input":"2023-08-10T10:07:55.594933Z","iopub.status.idle":"2023-08-10T10:07:55.751302Z","shell.execute_reply.started":"2023-08-10T10:07:55.594893Z","shell.execute_reply":"2023-08-10T10:07:55.750238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig,ax = plt.subplots(1)\nax.imshow(image.numpy())\nfor box in boxes.numpy():\n    width = abs(box[0]-box[2])\n    height = abs(box[1]-box[3])\n    rect = mpl.patches.Rectangle((box[0],box[1]),width, height,fill=False)\n    ax.add_patch(rect)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:07:55.884189Z","iopub.execute_input":"2023-08-10T10:07:55.884611Z","iopub.status.idle":"2023-08-10T10:07:56.379468Z","shell.execute_reply.started":"2023-08-10T10:07:55.884576Z","shell.execute_reply":"2023-08-10T10:07:56.378579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(all_masks)","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:07:56.380876Z","iopub.execute_input":"2023-08-10T10:07:56.381188Z","iopub.status.idle":"2023-08-10T10:07:56.767796Z","shell.execute_reply.started":"2023-08-10T10:07:56.381159Z","shell.execute_reply":"2023-08-10T10:07:56.766811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Phase","metadata":{"papermill":{"duration":0.013345,"end_time":"2023-05-27T17:50:12.356285","exception":false,"start_time":"2023-05-27T17:50:12.342940","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"- ### Model","metadata":{}},{"cell_type":"code","source":"def get_model(num_classes):    \n    model = maskrcnn_resnet50_fpn(pretrained=True,box_detections_per_img=20)\n\n    # get the number of input features for the classifier\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    \n    # replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes+1)\n    \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    \n    # and replace the mask predictor with a new one\n    model.roi_heads.mask_predictor = MaskRCNNPredictor(in_features_mask,\n                                                       hidden_layer, num_classes+1)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:08:05.142165Z","iopub.execute_input":"2023-08-10T10:08:05.142542Z","iopub.status.idle":"2023-08-10T10:08:05.149811Z","shell.execute_reply.started":"2023-08-10T10:08:05.142508Z","shell.execute_reply":"2023-08-10T10:08:05.148434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- ### Model Configs","metadata":{}},{"cell_type":"code","source":"def my_collate_fn(data):\n    return data\n\ntrain_dl = DataLoader(train_set, batch_size=8, shuffle=True,collate_fn=lambda x: tuple(x))\nval_dl = DataLoader(val_set, batch_size=8, shuffle=True,collate_fn=lambda x: tuple(x))\n\nmodel = get_model(num_classes=1)\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = optim.Adam(params, lr=0.001)\nnum_epochs = 20\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)\nmodel.to(device);","metadata":{"papermill":{"duration":470.246425,"end_time":"2023-05-27T17:58:02.653327","exception":false,"start_time":"2023-05-27T17:50:12.406902","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-08-10T10:08:14.240718Z","iopub.execute_input":"2023-08-10T10:08:14.241167Z","iopub.status.idle":"2023-08-10T10:08:15.051779Z","shell.execute_reply.started":"2023-08-10T10:08:14.241125Z","shell.execute_reply":"2023-08-10T10:08:15.050718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- ### Training Loop","metadata":{}},{"cell_type":"code","source":"model.train()\nfaulty = 0\nnum_epochs = 20\n\nfor epoch in range(num_epochs):\n    running_loss = 0.0\n    start_time = time.time()\n    \n    for data in tqdm(train_dl):\n        images = []\n        targets = []\n        for i in data:\n            if i[2].shape[-1] == 4:\n                images.append(i[0].permute(2,0,1).to(device))\n                targets.append({\n                    \"boxes\" : torch.tensor(np.array(i[2]),dtype=torch.float32).to(device),\n                    \"labels\" : torch.tensor(np.array(i[3]),dtype=torch.int64).to(device),\n                    \"masks\" : torch.tensor(np.array(i[1]),dtype=torch.float32).to(device)\n                })\n            else:\n                faulty += 1\n                continue\n        optimizer.zero_grad()\n        \n        outputs = model(images,targets)\n        total_loss = 0\n        for loss in outputs:\n            if loss == \"loss_mask\" or loss == \"loss_box_reg\":\n                total_loss += outputs[loss]*1.5\n            else:\n                total_loss += outputs[loss]\n        \n        total_loss.backward()\n        optimizer.step()\n\n        running_loss += total_loss.item()\n        \n    with torch.no_grad():\n        running_val_loss = 0.0\n        for data in tqdm(val_dl):\n            images = []\n            targets = []\n            for i in data:\n                if i[2].shape[-1] == 4:\n                    images.append(i[0].permute(2,0,1).to(device))\n                    targets.append({\n                        \"boxes\" : torch.tensor(np.array(i[2]),dtype=torch.float32).to(device),\n                        \"labels\" : torch.tensor(np.array(i[3]),dtype=torch.int64).to(device),\n                        \"masks\" : torch.tensor(np.array(i[1]),dtype=torch.float32).to(device)\n                    })\n                else:\n                    continue\n\n            outputs = model(images,targets)\n            total_loss = 0\n            for loss in outputs:\n                if loss == \"loss_mask\" or loss == \"loss_box_reg\":\n                    total_loss += outputs[loss]*1.5\n                else:\n                    total_loss += outputs[loss]\n\n            running_val_loss += total_loss.item()\n\n    epoch_loss = running_loss / len(train_dl)\n    epoch_time = time.time() - start_time\n    print(f\"Epoch {epoch+1}/{num_epochs}, Loss: {epoch_loss:.4f}, Time: {epoch_time:.2f} seconds\")\n    print(\"Validation Loss : \", round(running_val_loss,3)/ len(val_dl))\n","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-08-10T10:08:21.854301Z","iopub.execute_input":"2023-08-10T10:08:21.854951Z","iopub.status.idle":"2023-08-10T10:28:29.213313Z","shell.execute_reply.started":"2023-08-10T10:08:21.854907Z","shell.execute_reply":"2023-08-10T10:28:29.212336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outputs","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:30:13.640249Z","iopub.execute_input":"2023-08-10T10:30:13.640723Z","iopub.status.idle":"2023-08-10T10:30:13.655761Z","shell.execute_reply.started":"2023-08-10T10:30:13.640671Z","shell.execute_reply":"2023-08-10T10:30:13.654398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Testing","metadata":{"papermill":{"duration":0.013376,"end_time":"2023-05-27T17:58:02.680228","exception":false,"start_time":"2023-05-27T17:58:02.666852","status":"completed"},"tags":[]}},{"cell_type":"code","source":"model.eval();\nindex = 89\npreds = model([dataset[index][0].permute(2,0,1).to(device)])\npred_masks = preds[0][\"masks\"]","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:30:16.900942Z","iopub.execute_input":"2023-08-10T10:30:16.901326Z","iopub.status.idle":"2023-08-10T10:30:17.111322Z","shell.execute_reply.started":"2023-08-10T10:30:16.901297Z","shell.execute_reply":"2023-08-10T10:30:17.110282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count = 0\nfiltered_mask = np.zeros((512,512))\nfor mask in tqdm(pred_masks):\n    mask = mask.squeeze().cpu().detach().numpy()\n    indexes= np.where(mask>=0)\n    filtered_mask[indexes[0],indexes[1]] = 1    ","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:30:18.936609Z","iopub.execute_input":"2023-08-10T10:30:18.937859Z","iopub.status.idle":"2023-08-10T10:30:19.057855Z","shell.execute_reply.started":"2023-08-10T10:30:18.937808Z","shell.execute_reply":"2023-08-10T10:30:19.056777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visual Results","metadata":{}},{"cell_type":"code","source":"for temp in val_set:\n    fig, axs = plt.subplots(nrows=1, ncols=2,)\n    preds = model([temp[0].permute(2,0,1).to(device)])\n    pred_masks = preds[0][\"masks\"]\n    \n    count = 0\n    filtered_mask = np.zeros((512,512))\n    for mask in pred_masks:\n        mask = mask.squeeze().cpu().detach().numpy()\n        indexes= np.where(mask>0.6)\n        filtered_mask[indexes[0],indexes[1]] = 1\n\n    print(temp[-1])\n    \n    axs[0].imshow(temp[4])\n    axs[0].set_title(\"Original\")\n    axs[1].imshow(filtered_mask)\n    axs[1].set_title(\"Predicted\")\n    plt.show()\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-10T10:30:23.664405Z","iopub.execute_input":"2023-08-10T10:30:23.664833Z","iopub.status.idle":"2023-08-10T10:30:57.216115Z","shell.execute_reply.started":"2023-08-10T10:30:23.664795Z","shell.execute_reply":"2023-08-10T10:30:57.215162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}