{"cells":[{"metadata":{},"cell_type":"markdown","source":"## Algorithm to compute \"optimal\" coordinates for patches/tiles\n\nSorry if the code is messy and/or unreadable! But I tried to document/comment here and there to make it a bit clearer. Also, the algorithm is not optimized in terms of run-time (it's rather slow actually), but aims to optimize the coordinates of the patches/tiles.\n\n\nMain sections:\n\n* [Computing patch coordinates with visualization](#Precompute-patch-coordinates-for-later-use)\n* [Computing patches and stitching them together with visualization](#TensorFlow:-Stitch-patches-together,-using-tf-operations-and-tf.data.Dataset)\n","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install imagecodecs","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2 \nfrom tqdm.notebook import tqdm\nimport skimage.io\nimport tensorflow as tf\nimport math\nimport glob","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"data = pd.read_csv('../input/prostate-cancer-grade-assessment/train.csv')\ninput_path = '../input/prostate-cancer-grade-assessment/train_images/'\ndata.head(3)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Reading image","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def read_image(image_path, resize_ratio=1):\n    \n    if not(isinstance(image_path, str)):\n        # if tensor with byte string\n        image_path = image_path.numpy().decode('utf-8')\n        \n    image_level_1 = skimage.io.MultiImage(image_path)[1]\n    \n    if resize_ratio != 1:\n        new_w = int(image_level_1.shape[1]*resize_ratio)\n        new_h = int(image_level_1.shape[0]*resize_ratio)\n        image_level_1 = cv2.resize(\n            image_level_1, (new_w, new_h), interpolation=cv2.INTER_AREA)\n    \n    return image_level_1\n\nimage = read_image(input_path + data.image_id[0] + '.tiff')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Masking image","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def _mask_tissue(image, kernel_size=(7, 7), gray_threshold=220):\n    \"\"\"Masks tissue in image. Uses gray-scaled image, as well as\n    dilation kernels and 'gap filling'\n    \"\"\"\n    # Define elliptic kernel\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, kernel_size)\n    # Convert rgb to gray scale for easier masking\n    gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n    # Now mask the gray-scaled image (capturing tissue in biopsy)\n    mask = np.where(gray < gray_threshold, 1, 0).astype(np.uint8)\n    # Use dilation and findContours to fill in gaps/holes in masked tissue\n    mask = cv2.dilate(mask, kernel, iterations=1)\n    contour, _ = cv2.findContours(mask, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_SIMPLE)\n    for cnt in contour:\n        cv2.drawContours(mask, [cnt], 0, 1, -1)\n    return mask\n\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 12))\n\nmask = _mask_tissue(image)\n\naxes[0].imshow(image)\naxes[1].imshow(mask)\naxes[0].axis('off')\naxes[1].axis('off');","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Compute patch coordinates\n\nhelper functions:\n```\n_pad_image(...)\n_transpose_image(...)\n_get_tissue_parts_indices(...)\n_get_tissue_subparts_coords(...)\n_eval_and_append_xy_coords(...)\n```\nmain function:\n```\ncompute_coords(image, patch_size, ...)\n```\n","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def _pad_image(image, pad_len, pad_val):\n    \"\"\"Pads inputted image, accepts both \n    2-d (mask) and 3-d (rgb image) arrays\n    \"\"\"\n    if image is None:\n        return None\n    elif image.ndim == 2:\n        return np.pad(\n            image, ((pad_len, pad_len), (pad_len, pad_len)), pad_val)\n    elif image.ndim == 3:\n        return np.pad(\n            image, ((pad_len, pad_len), (pad_len, pad_len), (0, 0)), pad_val)\n    return None\n\ndef _transpose_image(image):\n    \"\"\"Inputs an image and transposes it, accepts \n    both 2-d (mask) and 3-d (rgb image) arrays\n    \"\"\"\n    if image is None:\n        return None\n    elif image.ndim == 2:\n        return np.transpose(image, (1, 0)).copy()\n    elif image.ndim == 3:\n        return np.transpose(image, (1, 0, 2)).copy()\n    return None\n\ndef _get_tissue_parts_indices(tissue, min_consec_info):\n    \"\"\"If there are multiple tissue parts in 'tissue', 'tissue' will be \n    split. Each tissue part will be taken care of separately (later on), \n    and if the tissue part is less than min_consec_info, it's considered \n    to small and won't be returned.\n    \"\"\"\n    split_points = np.where(np.diff(tissue) != 1)[0]+1\n    tissue_parts = np.split(tissue, split_points)\n    return [\n        tp for tp in tissue_parts if len(tp) >= min_consec_info\n    ]\n\ndef _get_tissue_subparts_coords(subtissue, patch_size, min_decimal_keep):\n    \"\"\"Inputs a tissue part resulting from '_get_tissue_parts_indices'.\n    This tissue part is divided into N subparts and returned.\n    Argument min_decimal_keep basically decides if we should reduce the\n    N subparts to N-1 subparts, due to overflow.\n    \"\"\"\n    start, end = subtissue[0], subtissue[-1]\n    num_subparts = (end-start)/patch_size\n    if num_subparts % 1 < min_decimal_keep and num_subparts >= 1:\n        num_subparts = math.floor(num_subparts)\n    else:\n        num_subparts = math.ceil(num_subparts)\n\n    excess = (num_subparts*patch_size) - (end-start)\n    shift = excess // 2\n\n    return [\n        i * patch_size + start - shift \n        for i in range(num_subparts)\n    ]\n\ndef _eval_and_append_xy_coords(coords,\n                               image, \n                               mask, \n                               patch_size, \n                               x, y, \n                               min_patch_info,\n                               transposed,\n                               precompute):\n    \"\"\"Based on computed x and y coordinates of patch: \n    slices out patch from original image, flattens it,\n    preprocesses it, and finally evaluates its mask.\n    If patch contains more info than min_patch_info,\n    the patch coordinates are kept, along with a value \n    'val1' that estimates how much information there \n    is in the patch. Smaller 'val1' assumes more info.\n    \"\"\"\n    patch_1d = (\n        image[y: y+patch_size, x:x+patch_size, :]\n        .mean(axis=2)\n        .reshape(-1)\n    )\n    idx_tissue = np.where(patch_1d <= 210)[0]\n    idx_black = np.where(patch_1d < 5)[0]\n    idx_background = np.where(patch_1d > 210)[0]\n\n    if len(idx_tissue) > 0:\n        patch_1d[idx_black] = 210\n        patch_1d[idx_background] = 210\n        val1 = int(patch_1d.mean())\n        val2 = mask[y:y+patch_size, x:x+patch_size].mean()\n        if val2 > min_patch_info:\n            if precompute:\n                if transposed:\n                    coords = np.concatenate([\n                        coords, [[val1, x-patch_size, y-patch_size]]\n                    ])\n                else:\n                    coords = np.concatenate([\n                        coords, [[val1, y-patch_size, x-patch_size]]\n                    ])\n            else:\n                coords = np.concatenate([\n                    coords, [[val1, y, x]]\n                ])\n               \n    return coords\n\ndef compute_coords(image,\n                   patch_size=256,\n                   precompute=False,\n                   min_patch_info=0.35,\n                   min_axis_info=0.35,\n                   min_consec_axis_info=0.35,\n                   min_decimal_keep=0.7):\n\n    \"\"\"\n    Input:\n        image : 3-d np.ndarray\n        patch_size : size of patches/tiles, will be of \n            size (patch_size x patch_size x 3)\n        precompute : If True, only coordinates will be returned,\n            these coordinates match the inputted 'original' image.\n            If False, both an image and coordinates will be returned,\n            the coordinates does not match the inputted image but the\n            image that it is returned with.\n        min_patch_info : Minimum required information in patch\n            (see '_eval_and_append_xy_coords')\n        min_axis_info : Minimum fraction of on-bits in x/y dimension to be \n            considered enough information. For x, this would be fraction of \n            on-bits in x-dimension of a y:y+patch_size slice. For y, this would \n            be the fraction of on-bits for the whole image in y-dimension\n        min_consec_axis_info : Minimum consecutive x/y on-bits\n            (see '_get_tissue_parts_indices')\n        min_decimal_keep : Threshold for decimal point for removing \"excessive\" patch\n            (see '_get_tissue_subparts_coords')\n    \n    Output:\n        image [only if precompute is False] : similar to input image, but fits \n            to the computed coordinates\n        coords : the coordinates that will be used to compute the patches later on\n    \"\"\"\n    \n    \n    if type(image) != np.ndarray:\n        # if image is a Tensor\n        image = image.numpy()\n    \n    # masked tissue will be used to compute the coordinates\n    mask = _mask_tissue(image)\n\n    # initialize coordinate accumulator\n    coords = np.zeros([0, 3], dtype=int)\n\n    # pad image and mask to make sure no tissue is potentially missed out\n    image = _pad_image(image, patch_size, 'maximum')\n    mask = _pad_image(mask, patch_size, 'minimum')\n    \n    y_sum = mask.sum(axis=1)\n    x_sum = mask.sum(axis=0)\n    # if on bits in x_sum is greater than in y_sum, the tissue is\n    # likely aligned horizontally. The algorithm works better if\n    # the image is aligned vertically, thus the image will be transposed\n    if len(np.where(x_sum > 0)[0]) > len(np.where(y_sum > 0)[0]):\n        image = _transpose_image(image)\n        mask = _transpose_image(mask)\n        y_sum, _ = x_sum, y_sum\n        transposed = True\n    else:\n        transposed = False\n    \n    # where y_sum is more than the minimum number of on-bits\n    y_tissue = np.where(y_sum >= (patch_size*min_axis_info))[0]\n    \n    if len(y_tissue) < 1:\n        warnings.warn(\"Not enough tissue in image (y-dim)\", RuntimeWarning)\n        if precompute: return [(0, 0, 0)]\n        else: return image, [(0, 0, 0)]\n    \n    y_tissue_parts_indices = _get_tissue_parts_indices(\n        y_tissue, patch_size*min_consec_axis_info)\n    \n    if len(y_tissue_parts_indices) < 1: \n        warnings.warn(\"Not enough tissue in image (y-dim)\", RuntimeWarning)\n        if precompute: return [(0, 0, 0)]\n        else: return image, [(0, 0, 0)]\n    \n    # loop over the tissues in y-dimension\n    for yidx in y_tissue_parts_indices:\n        y_tissue_subparts_coords = _get_tissue_subparts_coords(\n            yidx, patch_size, min_decimal_keep)\n        \n        for y in y_tissue_subparts_coords:\n            # in y_slice, where x_slice_sum is more than the minimum number of on-bits\n            x_slice_sum = mask[y:y+patch_size, :].sum(axis=0)\n            x_tissue = np.where(x_slice_sum >= (patch_size*min_axis_info))[0]\n            \n            x_tissue_parts_indices = _get_tissue_parts_indices(\n                x_tissue, patch_size*min_consec_axis_info)\n            \n            # loop over tissues in x-dimension (inside y_slice 'y:y+patch_size')\n            for xidx in x_tissue_parts_indices:\n                x_tissue_subparts_coords = _get_tissue_subparts_coords(\n                    xidx, patch_size, min_decimal_keep)\n                \n                for x in x_tissue_subparts_coords:\n                    coords = _eval_and_append_xy_coords(\n                        coords, image, mask, patch_size, x, y, \n                        min_patch_info, transposed, precompute\n                    )     \n    \n    if len(coords) < 1:\n        warnings.warn(\"Not enough tissue in image (x-dim)\", RuntimeWarning)\n        if precompute: return [(0, 0, 0)]\n        else: return image, [(0, 0, 0)]\n    \n    if precompute: return coords\n    else: return image, coords\n\n\ncoords = compute_coords(image, precompute=True)\nprint(\"    val  y   x\\n\", coords[:10])\n\nimage, coords = compute_coords(image, precompute=False)\nprint(\"    val  y   x\\n\", coords[:10])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Precompute patch coordinates for later use\n\n`compute_coords(..., precompute=True)`\n\nHyperparameters to vary (between 0 and 1):\n\n```\npatch_size\n\nmin_patch_info\nmin_axis_info\nmin_consec_axis_info\nmin_decimal_keep\n```\n","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"fig, axes = plt.subplots(10, 1, figsize=(20, 140))\n\npatch_size = 256\n\nfor i, ax in enumerate(axes.reshape(-1)):\n    image_path = input_path + data.image_id[i+500] + '.tiff'\n    image = read_image(image_path, 1)\n    \n    coords = compute_coords(image,\n                            patch_size=patch_size,\n                            precompute=True,\n                            min_patch_info=0.35,\n                            min_axis_info=0.35,\n                            min_consec_axis_info=0.35,\n                            min_decimal_keep=0.7)\n    \n    # sort coords (high info -> low info)\n    coords = sorted(coords, key= lambda x: x[0], reverse=False)\n    for (v, y, x) in coords:\n        end_point = (x, y)\n        start_point = (x+patch_size, y+patch_size)\n        image = cv2.rectangle(image, start_point, end_point, 2, 14)\n    \n    ax.imshow(image)\n    ax.axis('off')\n    ax.set_title(\n        \"num patches = \"+str(len(coords))+\", isup grade = \"+str(data.isup_grade[i+500]),\n        fontsize=20)\n\nplt.subplots_adjust(hspace=0.05, wspace=0.05)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### TensorFlow: Stitch patches together, using tf-operations and tf.data.Dataset\n\n```\ncompute_coords(..., precompute=False) -> patch_image(..., coords)\n```","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"def _patch_augment(patch, p=0.5):\n    \"\"\"Performs random rotation, random flip (u/d, l/r),\n    and random transpose, based on probability p\"\"\"\n    r1 = tf.random.uniform(\n        shape=(4,), minval=0, maxval=1)\n    r2 = tf.random.uniform(\n        shape=(), minval=0, maxval=4, dtype=tf.int32)\n    if r1[0] > (1-p):\n        patch = tf.image.rot90(patch, k=r2)\n    if r1[1] > (1-p):\n        patch = tf.image.random_flip_left_right(patch)\n    if r1[2] > (1-p):\n        patch = tf.image.random_flip_up_down(patch)\n    if r1[3] > (1-p):\n        patch = tf.transpose(patch, (1, 0, 2))\n    return patch\n\ndef _excess_coords_filtering(coords, sample_size, proportion=0.25):\n    \"\"\"filters out a portion of excessive coordinates.\n    coordinates with higher values are filtered out.\n    \"\"\"\n    if len(coords) > sample_size:\n        c = tf.transpose(coords)\n        v = tf.gather(c, 0)\n        num = tf.cast(len(v), tf.float32)\n        sample_size = tf.cast(sample_size, tf.float32)\n        indices_reduced = int(tf.math.ceil(\n            num * (1 - ((num - sample_size) / num) * proportion)\n        ))\n        v_argsort = tf.argsort(v)\n        indices = tf.gather(v_argsort, tf.range(indices_reduced))\n        indices = tf.sort(indices)\n        coords = tf.gather(coords, indices)\n    return coords\n\n@tf.function\ndef patch_image(image, coords, sample_size=36, patch_size=256):\n\n    l = tf.cast(tf.math.sqrt(tf.cast(sample_size, tf.float32)), tf.int32)\n    \n    coords = _excess_coords_filtering(coords, sample_size)\n    # coords = tf.random.shuffle(coords)\n    if len(coords) < sample_size:\n        indices = tf.tile(\n            tf.range(len(coords)), [tf.math.ceil(sample_size/len(coords))])\n        indices = indices[:sample_size]\n    else:\n        indices = tf.range(sample_size)\n\n    coords = tf.gather(coords, indices)\n    \n    coords = tf.random.shuffle(coords) # Update: shuffle here instead\n    \n    patched_image = tf.zeros(\n        [0, patch_size, patch_size, 3], dtype=tf.dtypes.uint8)\n\n    for i in range(sample_size):\n        y = tf.gather_nd(coords, [i, 1])\n        x = tf.gather_nd(coords, [i, 2])\n        shape = tf.shape(image)\n        h = tf.gather(shape, 0)\n        w = tf.gather(shape, 1)\n        if y < 0: y = 0\n        if x < 0: x = 0\n        if y > h-patch_size: y = h-patch_size\n        if x > w-patch_size: x = w-patch_size\n            \n        patch = tf.slice(\n            image, \n            tf.stack([y, x, 0]), \n            tf.stack([patch_size, patch_size, -1]))\n\n        patch = _patch_augment(patch)\n        patched_image = tf.concat([\n            patched_image, tf.expand_dims(patch, 0)], axis=0)\n    \n    patched_image = tf.reshape(patched_image, (-1, patch_size*l, patch_size, 3))\n    patched_image = tf.transpose(patched_image, (0, 2, 1, 3))\n    patched_image = tf.reshape(patched_image, (patch_size*l, patch_size*l, 3))\n        \n    return patched_image\n\n\nnp.random.seed(42)\nimage_paths = glob.glob(input_path + '*')\nnp.random.shuffle(image_paths)\n\ndataset = tf.data.Dataset.from_tensor_slices(image_paths)\n\ndataset = dataset.map(\n    lambda x: tf.py_function(\n        func=read_image,\n        inp=[x],\n        Tout=[tf.uint8]), \n    num_parallel_calls=tf.data.experimental.AUTOTUNE)\n\ndataset = dataset.map(\n    lambda x: tf.py_function(\n        func=compute_coords,\n        inp=[x],\n        Tout=[tf.uint8, tf.int32]\n    ),\n    num_parallel_calls= tf.data.experimental.AUTOTUNE)\n\ndataset = dataset.map(\n    patch_image, \n    num_parallel_calls=tf.data.experimental.AUTOTUNE)\n\ndataset = dataset.batch(4)\ndataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)\n\nfor x in dataset.take(1):\n    pass","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fig, axes = plt.subplots(2, 2, figsize=(15, 15))\n\nfor i, ax in enumerate(axes.reshape(-1)):\n    ax.imshow(x.numpy()[i])\n    ax.axis('off')\n    \nplt.subplots_adjust(hspace=0.035, wspace=0.035)","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}