{"cells":[{"metadata":{},"cell_type":"markdown","source":"Mask -RCNN model to segment and classify two different cell types. For the code to run, tensorflow version 1.x must be utilitised. This can be achieved by [copy + edit] an existing older notebook, here https://www.kaggle.com/hmendonca/mask-rcnn-and-coco-transfer-learning-lb-0-155. \n\nThis notebook can also train on grayscale images if only pictures of the cytoplasma are given. This increases the computation speed. However, for images with nucleus channel this channel should be used as well. The changes necessary for grayscale images are indicated by '#grayscale'. \n\nThe dataset datacells consists of 1200 images for each cell type, N2a and HEK. The data was segmented using cellpose.ipynb and the annotation .json files were created using cells_to_coco.ipynb. \n\nIt is expected that size of the dataset is sufficient to avoid overfitting, especially since in some images there are a lot of cells (especially HEK). Therefore, at this point no data augmentation is used. But, since it is build-in in MaskRCNN, the inclusion should not be hard. \n\nInspiration was taken from https://github.com/waspinator/deep-learning-explorer/tree/master/mask-rcnn and the shapes sample in the Mask R-CNN github repository."},{"metadata":{"trusted":true},"cell_type":"code","source":"#import relevant packages\nimport os\nimport sys\nimport random\nimport math\nimport re\nimport time\nimport numpy as np\nimport cv2\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n\n#data directory\nDATA_DIR = '/kaggle/input/datacells2/data-cells2'\n\n# Directory to save logs and trained model\nROOT_DIR = '/kaggle/working'\n\n#install pycocotools to treat data in COCO format\n!pip install pycocotools\n\n#clone mask rcnn\n!git clone https://www.github.com/matterport/Mask_RCNN.git\nos.chdir('Mask_RCNN')\n\n!python3 setup.py install\n\n#note that mask rcnn need tensorflow version 1.x to work, hence in kaggle an older notebook needed to be copied, in google colab this can be easy implemented \n#and on a local machine a special anaconda environment has to be utilised\nimport tensorflow as tf\nprint(tf.__version__)\n\n# Import Mask RCNN\nsys.path.append(os.path.join(ROOT_DIR, 'Mask_RCNN'))  # To find local version of the library\nfrom mrcnn.config import Config\nfrom mrcnn import utils\nimport mrcnn.model as modellib\nfrom mrcnn import visualize\nfrom mrcnn.model import log","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import samples.coco.coco as coco\n\n%matplotlib inline \n\n# Directory to save logs and trained model\nMODEL_DIR = os.path.join(ROOT_DIR, \"logs\")\n#DATA_DIR = os.path.join(ROOT_DIR, \"data\")\nWEIGHTS_DIR = os.path.join(ROOT_DIR, \"data\")\n\n# Local path to trained weights file\nCOCO_MODEL_PATH = os.path.join(ROOT_DIR, \"mask_rcnn_coco.h5\")\n# Download COCO trained weights from Releases if needed\nif not os.path.exists(COCO_MODEL_PATH):\n    utils.download_trained_weights(COCO_MODEL_PATH)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import skimage\n\n#modify dataset to input images in grayscale\nclass CellsDataset(coco.CocoDataset):\n\n    def load_image(self, image_id):\n        \"\"\"Load the specified image and return a [H,W] Numpy array. Before it was a [H,W,3] array.\n        \"\"\"\n        # Load image\n        image = skimage.io.imread(self.image_info[image_id]['path'], as_gray=True)\n        image = image[..., np.newaxis]\n        # If grayscale. Convert to RGB for consistency.\n        #if image.ndim != 3:\n        #    image = skimage.color.gray2rgb(image)\n        # If has an alpha channel, remove it for consistency\n        if image.shape[-1] == 4:\n            image = image[..., :3]\n        return image","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Directory containing the input data has to be of the form: \n\n    DATA_DIR\n    |\n    |\n    |___ annotations \n    |\n    |--- cells_train2020\n    |\n    |--- cells_validate2020\n    |\n    |--- cells_test2020"},{"metadata":{"trusted":true},"cell_type":"code","source":"#load datasets\ndataset_train = coco.CocoDataset() #CellsDataset() #grayscale\ndataset_train.load_coco(DATA_DIR, subset=\"cells_train\", year=\"2020\")\ndataset_train.prepare()\n\ndataset_validate = coco.CocoDataset() #CellsDataset()\ndataset_validate.load_coco(DATA_DIR, subset=\"cells_validate\", year=\"2020\")\ndataset_validate.prepare()\n\ndataset_test = coco.CocoDataset() #CellsDataset()\ndataset_test.load_coco(DATA_DIR, subset=\"cells_test\", year=\"2020\")\ndataset_test.prepare()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Load and display random samples\nimage_ids = np.random.choice(dataset_train.image_ids, 4)\nfor image_id in image_ids:\n    image = dataset_train.load_image(image_id)\n    mask, class_ids = dataset_train.load_mask(image_id)\n    #visualize.display_top_masks(image[:,:,0], mask, class_ids, dataset_train.class_names) #grayscale\n    visualize.display_top_masks(image, mask, class_ids, dataset_train.class_names)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Adjust configuration file for the images\n# E.g. the images are reduced in size when input etc.\n\nimage_size = 512\nrpn_anchor_template = (1, 2, 4, 8, 16) # anchor sizes in pixels\nrpn_anchor_scales = tuple(i * (image_size // 16) for i in rpn_anchor_template)\n\nclass CellsConfig(Config):\n    \"\"\"Configuration for training on the cells dataset.\n    \"\"\"\n    NAME = \"cells\"\n\n    # Train on 1 GPU and 2 images per GPU. Put multiple images on each\n    # GPU if the images are small. Batch size is 2 (GPUs * images/GPU).\n    GPU_COUNT = 1\n    IMAGES_PER_GPU = 4\n\n    # Number of classes (including background)\n    NUM_CLASSES = 1 + 2  # background + 2 cell types (MEF, HEK)\n\n    # Use smaller images for faster training. \n    IMAGE_MAX_DIM = image_size\n    IMAGE_MIN_DIM = image_size\n    \n    # Use smaller anchors because our image and objects are small\n    RPN_ANCHOR_SCALES = rpn_anchor_scales\n\n    # Aim to allow ROI sampling to pick 33% positive ROIs.\n    TRAIN_ROIS_PER_IMAGE = 200#100 #32\n\n    # Input grayscale images\n    #IMAGE_CHANNEL_COUNT = 1\n    #MEAN_PIXEL = 1\n\n    STEPS_PER_EPOCH = 100\n\n    VALIDATION_STEPS = STEPS_PER_EPOCH / 20\n    \nconfig = CellsConfig()\nconfig.display()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_ax(rows=1, cols=1, size=8):\n    \"\"\"Return a Matplotlib Axes array to be used in\n    all visualizations in the notebook. Provide a\n    central point to control graph sizes.\n    \n    Change the default size attribute to control the size\n    of rendered images\n    \"\"\"\n    _, ax = plt.subplots(rows, cols, figsize=(size*cols, size*rows))\n    return ax","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Create model in training mode\nmodel = modellib.MaskRCNN(mode=\"training\", config=config,\n                          model_dir=MODEL_DIR)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Load weights, either start with imagenet or coco weights or start with last weights "},{"metadata":{"trusted":true},"cell_type":"code","source":"# Which weights to start with?\ninit_with = \"last\"  # imagenet, coco, or last\n\nif init_with == \"imagenet\":\n    model.load_weights(model.get_imagenet_weights(), by_name=True)\nelif init_with == \"coco\":\n    # Load weights trained on MS COCO, but skip layers that\n    # are different due to the different number of classes\n    # See README for instructions to download the COCO weights\n    model.load_weights(COCO_MODEL_PATH, by_name=True,\n                       exclude=[\"mrcnn_class_logits\", \"mrcnn_bbox_fc\", \n                                \"mrcnn_bbox\", \"mrcnn_mask\"])#, \"conv1\"]) #grayscale\nelif init_with == \"last\":\n    # Load the last model you trained and continue training\n    model.load_weights(model.find_last(), by_name=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#can also load manually created weights via \nmodel.load_weights('/kaggle/input/mrcnn-weights/mask_rcnn_cells_0064.h5', by_name=True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Train in two step process: First only the head branches and then the full network. "},{"metadata":{"trusted":true},"cell_type":"code","source":"# Train the head branches\n# Passing layers=\"heads\" freezes all layers except the head\n# layers. You can also pass a regular expression to select\n# which layers to train by name pattern.\nmodel.train(dataset_train, dataset_validate, \n            learning_rate=config.LEARNING_RATE, \n            epochs=1, \n            #layers=r\"(conv1)|(mrcnn\\_.*)|(rpn\\_.*)|(fpn\\_.*)\") #grayscale\n            layers='heads')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# necessary to download weights from training directory, otherwise irrelevant\nimport os\nos.rename('/kaggle/working/logs/cells20210310T1243/mask_rcnn_cells_0020.h5', '/kaggle/working/mask_rcnn_cells_0030.h5')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"At the current setting with 100 steps per epoch, 8 images per GPU and an image size of (256, 256, 3) the code takes ~22 min for one epoch training on the heads."},{"metadata":{"trusted":true},"cell_type":"code","source":"# Fine tune all layers\n# Passing layers=\"all\" trains all layers. You can also \n# pass a regular expression to select which layers to\n# train by name pattern.\nmodel.train(dataset_train, dataset_validate, \n            learning_rate=config.LEARNING_RATE / 10,\n            epochs=2, \n            layers=\"all\")","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Inference"},{"metadata":{"trusted":true},"cell_type":"code","source":"class InferenceConfig(CellsConfig):\n    GPU_COUNT = 1\n    IMAGES_PER_GPU = 1\n\ninference_config = InferenceConfig()\n\n# Recreate the model in inference mode\nmodel = modellib.MaskRCNN(mode=\"inference\", \n                          config=inference_config,\n                          model_dir=MODEL_DIR)\n\n# Get path to saved weights\n# Either set a specific path or find last trained weights\nmodel.load_weights('/kaggle/input/mrcnn-weights/mask_rcnn_cells_0064.h5', by_name=True)\n#print(model.find_last()[1])\n#model_path = model.find_last()\n\n# Load trained weights\n#assert model_path != \"\", \"Provide path to trained weights\"\n#print(\"Loading weights from \", model_path)\n#model.load_weights(model_path, by_name=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"original_image = skimage.io.imread('/kaggle/input/mixed-images/Mixed_0011_merged.png')\n\nresults = model.detect([original_image], verbose=1)\n\nr = results[0]\nvisualize.display_instances(original_image, r['rois'], r['masks'], r['class_ids'], \n                            dataset_validate.class_names, r['scores'], ax=get_ax(), show_bbox=False, show_mask=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#determine the ratio of HEK vs N2a cells detected in the mixed images and compare with predicted 50% of each celltype\nimport time\n\nHEKnum = 0\nN2anum = 0\ntotnum = 0\ncount = 0\nstart = time.time()\nfor path, directories, files in os.walk('/kaggle/input/mixed-images/'):\n    files.sort()\n    for file in files:\n        original_image = skimage.io.imread(path + file)\n        results = model.detect([original_image], verbose=0)\n        r = results[0]\n        count += 1\n        for c in r['class_ids']:\n            if c == 1:\n                HEKnum += 1\n            else:\n                N2anum += 1\n            totnum += 1\nend = time.time()\nprint('Finished', count, 'images. Took', end-start, 'seconds.')\nprint('Percentage of HEKs', HEKnum/totnum)\nprint('Percentage of N2as', N2anum/totnum)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.imshow(original_image); plt.axis('off')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#compare prediction with actual output for a random image from the test dataset\n\nimage_id = random.choice(dataset_test.image_ids)\noriginal_image, image_meta, gt_class_id, gt_bbox, gt_mask =\\\n    modellib.load_image_gt(dataset_test, inference_config, \n                           image_id) #removed 'use_mini_mask=False'\n\nlog(\"original_image\", original_image)\nlog(\"image_meta\", image_meta)\nlog(\"gt_class_id\", gt_class_id)\nlog(\"gt_bbox\", gt_bbox)\nlog(\"gt_mask\", gt_mask)\n\n#rgb_image = skimage.color.gray2rgb(original_image[:, :, 0]) #grayscale\n\nrgb_image = original_image\n\nvisualize.display_instances(rgb_image, gt_bbox, gt_mask, gt_class_id, \n                            dataset_train.class_names, figsize=(8, 8), show_bbox=False, show_mask=False)\n\nresults = model.detect([original_image], verbose=1)\n\nr = results[0]\nvisualize.display_instances(rgb_image, r['rois'], r['masks'], r['class_ids'], \n                            dataset_validate.class_names, None, ax=get_ax(), show_bbox=False, show_mask=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Compute VOC-Style mAP @ IoU=0.5\n\nimage_ids = np.random.choice(dataset_test.image_ids, 60)\nAPs = []\n\nfor image_id in image_ids:\n    # Load image and ground truth data\n    image, image_meta, gt_class_id, gt_bbox, gt_mask =\\\n        modellib.load_image_gt(dataset_test, inference_config,\n                               image_id, use_mini_mask=False)\n    molded_images = np.expand_dims(modellib.mold_image(image, inference_config), 0)\n    # Run object detection\n    results = model.detect([image], verbose=0)\n    r = results[0]\n    # Compute AP\n    AP, precisions, recalls, overlaps =\\\n        utils.compute_ap(gt_bbox, gt_class_id, gt_mask,\n                         r[\"rois\"], r[\"class_ids\"], r[\"scores\"], r['masks'])\n    APs.append(AP)\n    \nprint(\"mAP: \", np.mean(APs))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#define functions to estimate the precision per class\n\ndef compute_per_class_precision(gt_boxes, gt_class_ids, gt_masks,\n              pred_boxes, pred_class_ids, pred_scores, pred_masks,\n              class_infos, iou_threshold=0.5):\n    \"\"\"\n        Compute per class precision\n    \"\"\"\n    \n    class_precisions = {}\n    \n    for class_info in class_infos:\n        if class_info[\"name\"] == \"BG\":\n            continue\n        \n        class_gt_indexes = np.where(gt_class_ids == class_info[\"id\"])\n        class_gt_boxes = gt_boxes[class_gt_indexes]\n        class_gt_masks = gt_masks[:, :, class_gt_indexes[0]]        \n        class_gt_ids = np.full(np.size(class_gt_indexes), class_info[\"id\"])\n        \n        class_pred_indexes = np.where(pred_class_ids == class_info[\"id\"])\n        class_pred_boxes = pred_boxes[class_pred_indexes]\n        class_pred_masks = pred_masks[:, :, class_pred_indexes[0]]\n        class_pred_scores = pred_scores[class_pred_indexes]\n        class_pred_ids = np.full(np.size(class_pred_indexes), class_info[\"id\"])\n        \n        if np.shape(class_gt_masks)[2] == 0 and np.shape(class_pred_masks)[2] == 0:\n            continue   \n\n        if np.shape(class_gt_masks)[2] == 0:\n            class_gt_indexes = (np.array([0]),)\n            class_gt_boxes = np.array([[1, 1, 1, 1]])\n            class_gt_masks =  np.zeros([np.shape(class_gt_masks)[0], np.shape(class_gt_masks)[1], 1])\n            class_gt_ids = np.full(np.size(class_gt_indexes), class_info[\"id\"])\n\n        if np.shape(class_pred_masks)[2] == 0:\n            class_pred_indexes = (np.array([0]),)\n            class_pred_masks =  np.zeros([np.shape(class_gt_masks)[0], np.shape(class_gt_masks)[1], 1])\n            class_pred_boxes = np.array([[1, 1, 1, 1]])\n            class_pred_scores = np.array([0])\n            class_pred_ids = np.full(np.size(class_pred_indexes), class_info[\"id\"])\n \n        AP, precisions, recalls, overlaps =\\\n            utils.compute_ap(class_gt_boxes, class_gt_ids, class_gt_masks,\n                            class_pred_boxes, class_pred_ids, class_pred_scores, class_pred_masks,\n                            iou_threshold)\n        \n        class_precisions[class_info[\"name\"]] = {\n            \"average_precision\": AP,\n            \"precisions\": precisions,\n            \"recalls\": recalls,\n            \"overlaps\": overlaps\n        }\n    \n    return class_precisions\n\ndef compute_multiple_per_class_precision(model, inference_config, dataset, \n                                        number_of_images=10, iou_threshold=0.5):\n    \"\"\"\n        Compute per class precision on multiple images\n    \"\"\"\n\n    image_ids = np.random.choice(dataset.image_ids, number_of_images, replace=False)\n\n    class_precisions = {}\n\n    for image_id in image_ids:\n        image, _, gt_class_id, gt_bbox, gt_mask =\\\n            modellib.load_image_gt(dataset, inference_config,\n                                image_id, use_mini_mask=False)\n\n        results = model.detect([image], verbose=0)\n        r = results[0]\n\n        class_precision_info =\\\n        compute_per_class_precision(gt_bbox, gt_class_id, gt_mask,\n                r[\"rois\"], r[\"class_ids\"], r[\"scores\"], r[\"masks\"],\n                dataset.class_info, iou_threshold)\n        \n        for class_name in class_precision_info:\n            if class_precisions.get(class_name):\n                class_precisions[class_name].append(class_precision_info[class_name]['average_precision'])\n            else:\n                class_precisions[class_name] = [class_precision_info[class_name]['average_precision']]\n                \n    return class_precisions","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Comparing the precision on the different datasets can hint towards overfitting."},{"metadata":{"trusted":true},"cell_type":"code","source":"#output the precision on the test datset\npredictions = compute_multiple_per_class_precision(model, inference_config, dataset_test,\n                                                 number_of_images=59, iou_threshold=0.5)\ncomplete_predictions = []\n\nfor shape in predictions:\n    complete_predictions += predictions[shape]\n    print(\"{} ({}): {}\".format(shape, len(predictions[shape]), np.mean(predictions[shape])))\n\nprint(\"--------\")\nprint(\"average: {}\".format(np.mean(complete_predictions)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#output the precision on the validation datset\npredictions = compute_multiple_per_class_precision(model, inference_config, dataset_validate,\n                                                 number_of_images=59, iou_threshold=0.5)\ncomplete_predictions = []\n\nfor shape in predictions:\n    complete_predictions += predictions[shape]\n    print(\"{} ({}): {}\".format(shape, len(predictions[shape]), np.mean(predictions[shape])))\n\nprint(\"--------\")\nprint(\"average: {}\".format(np.mean(complete_predictions)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#output the precision on the training datset\npredictions = compute_multiple_per_class_precision(model, inference_config, dataset_train,\n                                                 number_of_images=59, iou_threshold=0.5)\ncomplete_predictions = []\n\nfor shape in predictions:\n    complete_predictions += predictions[shape]\n    print(\"{} ({}): {}\".format(shape, len(predictions[shape]), np.mean(predictions[shape])))\n\nprint(\"--------\")\nprint(\"average: {}\".format(np.mean(complete_predictions)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}