{
  "cells": [
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "baa9a008-482d-dedf-ec30-86d571680925"
      },
      "outputs": [],
      "source": [
        "%matplotlib inline\n",
        "import pandas as pd\n",
        "import numpy as np\n",
        "\n",
        "from shapely.wkt import loads as wkt_loads\n",
        "from matplotlib.patches import Polygon, Patch\n",
        "\n",
        "# decartes package makes plotting with holes much easier\n",
        "from descartes.patch import PolygonPatch\n",
        "\n",
        "import matplotlib.pyplot as plt\n",
        "import tifffile as tiff\n",
        "\n",
        "import pylab\n",
        "# turn interactive mode on so that plots immediately\n",
        "# See: http://stackoverflow.com/questions/2130913/no-plot-window-in-matplotlib\n",
        "# pylab.ion()\n",
        "\n",
        "inDir = '../input'\n",
        "\n",
        "# Give short names, sensible colors and zorders to object types\n",
        "CLASSES = {\n",
        "        1 : 'Bldg',\n",
        "        2 : 'Struct',\n",
        "        3 : 'Road',\n",
        "        4 : 'Track',\n",
        "        5 : 'Trees',\n",
        "        6 : 'Crops',\n",
        "        7 : 'Fast H20',\n",
        "        8 : 'Slow H20',\n",
        "        9 : 'Truck',\n",
        "        10 : 'Car',\n",
        "        }\n",
        "COLORS = {\n",
        "        1 : '0.7',\n",
        "        2 : '0.4',\n",
        "        3 : '#b35806',\n",
        "        4 : '#dfc27d',\n",
        "        5 : '#1b7837',\n",
        "        6 : '#a6dba0',\n",
        "        7 : '#74add1',\n",
        "        8 : '#4575b4',\n",
        "        9 : '#f46d43',\n",
        "        10: '#d73027',\n",
        "        }\n",
        "ZORDER = {\n",
        "        1 : 5,\n",
        "        2 : 5,\n",
        "        3 : 4,\n",
        "        4 : 1,\n",
        "        5 : 3,\n",
        "        6 : 2,\n",
        "        7 : 7,\n",
        "        8 : 8,\n",
        "        9 : 9,\n",
        "        10: 10,\n",
        "        }\n",
        "\n",
        "# read the training data from train_wkt_v4.csv\n",
        "df = pd.read_csv(inDir + '/train_wkt_v4.csv')\n",
        "print(df.head())\n",
        "\n",
        "# grid size will also be needed later..\n",
        "gs = pd.read_csv(inDir + '/grid_sizes.csv', names=['ImageId', 'Xmax', 'Ymin'], skiprows=1)\n",
        "print(gs.head())\n",
        "\n",
        "# imageIds in a DataFrame\n",
        "allImageIds = gs.ImageId.unique()\n",
        "trainImageIds = df.ImageId.unique()\n",
        "\n",
        "\n",
        "def get_image_names(imageId):\n",
        "    '''\n",
        "    Get the names of the tiff files\n",
        "    '''\n",
        "    d = {'3': '{}/three_band/{}.tif'.format(inDir, imageId),\n",
        "         'A': '{}/sixteen_band/{}_A.tif'.format(inDir, imageId),\n",
        "         'M': '{}/sixteen_band/{}_M.tif'.format(inDir, imageId),\n",
        "         'P': '{}/sixteen_band/{}_P.tif'.format(inDir, imageId),\n",
        "         }\n",
        "    return d\n",
        "\n",
        "\n",
        "def get_images(imageId, img_key = None):\n",
        "    '''\n",
        "    Load images correspoding to imageId\n",
        "\n",
        "    Parameters\n",
        "    ----------\n",
        "    imageId : str\n",
        "        imageId as used in grid_size.csv\n",
        "    img_key : {None, '3', 'A', 'M', 'P'}, optional\n",
        "        Specify this to load single image\n",
        "        None loads all images and returns in a dict\n",
        "        '3' loads image from three_band/\n",
        "        'A' loads '_A' image from sixteen_band/\n",
        "        'M' loads '_M' image from sixteen_band/\n",
        "        'P' loads '_P' image from sixteen_band/\n",
        "\n",
        "    Returns\n",
        "    -------\n",
        "    images : dict\n",
        "        A dict of image data from TIFF files as numpy array\n",
        "    '''\n",
        "    img_names = get_image_names(imageId)\n",
        "    images = dict()\n",
        "    if img_key is None:\n",
        "        for k in img_names.keys():\n",
        "            images[k] = tiff.imread(img_names[k])\n",
        "    else:\n",
        "        images[img_key] = tiff.imread(img_names[img_key])\n",
        "    return images\n",
        "\n",
        "\n",
        "def get_size(imageId):\n",
        "    \"\"\"\n",
        "    Get the grid size of the image\n",
        "\n",
        "    Parameters\n",
        "    ----------\n",
        "    imageId : str\n",
        "        imageId as used in grid_size.csv\n",
        "    \"\"\"\n",
        "    xmax, ymin = gs[gs.ImageId == imageId].iloc[0,1:].astype(float)\n",
        "    W, H = get_images(imageId, '3')['3'].shape[1:]\n",
        "    return (xmax, ymin, W, H)\n",
        "\n",
        "\n",
        "def is_training_image(imageId):\n",
        "    '''\n",
        "    Returns\n",
        "    -------\n",
        "    is_training_image : bool\n",
        "        True if imageId belongs to training data\n",
        "    '''\n",
        "    return any(trainImageIds == imageId)\n",
        "\n",
        "\n",
        "def plot_polygons(fig, ax, polygonsList):\n",
        "    '''\n",
        "    Plot descrates.PolygonPatch from list of polygons objs for each CLASS\n",
        "    '''\n",
        "    legend_patches = []\n",
        "    for cType in polygonsList:\n",
        "        print('{} : {} \\tcount = {}'.format(cType, CLASSES[cType], len(polygonsList[cType])))\n",
        "        legend_patches.append(Patch(color=COLORS[cType],\n",
        "                                    label='{} ({})'.format(CLASSES[cType], len(polygonsList[cType]))))\n",
        "        for polygon in polygonsList[cType]:\n",
        "            mpl_poly = PolygonPatch(polygon,\n",
        "                                    color=COLORS[cType],\n",
        "                                    lw=0,\n",
        "                                    alpha=0.7,\n",
        "                                    zorder=ZORDER[cType])\n",
        "            ax.add_patch(mpl_poly)\n",
        "    # ax.relim()\n",
        "    ax.autoscale_view()\n",
        "    ax.set_title('Objects')\n",
        "    ax.set_xticks([])\n",
        "    ax.set_yticks([])\n",
        "    return legend_patches\n",
        "\n",
        "def plot_image(fig, ax, imageId, img_key, selected_channels=None):\n",
        "    '''\n",
        "    Plot get_images(imageId)[image_key] on axis/fig\n",
        "    Optional: select which channels of the image are used (used for sixteen_band/ images)\n",
        "    Parameters\n",
        "    ----------\n",
        "    img_key : str, {'3', 'P', 'N', 'A'}\n",
        "        See get_images for description.\n",
        "    '''\n",
        "    images = get_images(imageId, img_key)\n",
        "    img = images[img_key]\n",
        "    title_suffix = ''\n",
        "    if selected_channels is not None:\n",
        "        img = img[selected_channels]\n",
        "        title_suffix = ' (' + ','.join([ repr(i) for i in selected_channels ]) + ')'\n",
        "    if len(img.shape) == 2:\n",
        "        new_img = np.zeros((3, img.shape[0], img.shape[1]))\n",
        "        new_img[0] = img\n",
        "        new_img[1] = img\n",
        "        new_img[2] = img\n",
        "        img = new_img\n",
        "    \n",
        "    tiff.imshow(img, figure=fig, subplot=ax)\n",
        "    ax.set_title(imageId + ' - ' + img_key + title_suffix)\n",
        "    ax.set_xlabel(img.shape[-2])\n",
        "    ax.set_ylabel(img.shape[-1])\n",
        "    ax.set_xticks([])\n",
        "    ax.set_yticks([])\n",
        "\n",
        "\n",
        "\n",
        "def visualize_image(imageId, plot_all=True):\n",
        "    '''         \n",
        "    Plot all images and object-polygons\n",
        "    \n",
        "    Parameters\n",
        "    ----------\n",
        "    imageId : str\n",
        "        imageId as used in grid_size.csv\n",
        "    plot_all : bool, True by default\n",
        "        If True, plots all images (from three_band/ and sixteen_band/) as subplots.\n",
        "        Otherwise, only plots Polygons.\n",
        "    '''         \n",
        "    df_image = df[df.ImageId == imageId]\n",
        "    xmax, ymin, W, H = get_size(imageId)\n",
        "    \n",
        "    if plot_all:\n",
        "        fig, axArr = plt.subplots(figsize=(10, 10), nrows=3, ncols=3)\n",
        "        ax = axArr[0][0]\n",
        "    else:\n",
        "        fig, axArr = plt.subplots(figsize=(10, 10))\n",
        "        ax = axArr\n",
        "    if is_training_image(imageId):\n",
        "        print('ImageId : {}'.format(imageId))\n",
        "        polygonsList = {}\n",
        "        for cType in CLASSES.keys():\n",
        "            polygonsList[cType] = wkt_loads(df_image[df_image.ClassType == cType].MultipolygonWKT.values[0])\n",
        "        legend_patches = plot_polygons(fig, ax, polygonsList)\n",
        "        ax.set_xlim(0, xmax)\n",
        "        ax.set_ylim(ymin, 0)\n",
        "        ax.set_xlabel(xmax)\n",
        "        ax.set_ylabel(ymin)\n",
        "    if plot_all:\n",
        "        plot_image(fig, axArr[0][1], imageId, '3')\n",
        "        plot_image(fig, axArr[0][2], imageId, 'P')\n",
        "        plot_image(fig, axArr[1][0], imageId, 'A', [0, 3, 6])\n",
        "        plot_image(fig, axArr[1][1], imageId, 'A', [1, 4, 7])\n",
        "        plot_image(fig, axArr[1][2], imageId, 'A', [2, 5, 0])\n",
        "        plot_image(fig, axArr[2][0], imageId, 'M', [0, 3, 6])\n",
        "        plot_image(fig, axArr[2][1], imageId, 'M', [1, 4, 7])\n",
        "        plot_image(fig, axArr[2][2], imageId, 'M', [2, 5, 0])\n",
        "\n",
        "    if is_training_image(imageId):\n",
        "        ax.legend(handles=legend_patches,\n",
        "                   # loc='upper center',\n",
        "                   bbox_to_anchor=(0.9, 1),\n",
        "                   bbox_transform=plt.gcf().transFigure,\n",
        "                   ncol=5,\n",
        "                   fontsize='x-small',\n",
        "                   title='Objects-' + imageId,\n",
        "                   # mode=\"expand\",\n",
        "                   framealpha=0.3)\n",
        "    return (fig, axArr, ax)\n",
        "\n",
        "# Loop over few training images and save to files\n",
        "i = 1\n",
        "for imageId in trainImageIds:\n",
        "    fig, axArr, ax = visualize_image(imageId, plot_all=False)\n",
        "    #plt.savefig('Objects--' + imageId + '.png')\n",
        "    #plt.clf()\n",
        "    plt.show()\n",
        "    i = i+1\n",
        "    if(i>2):\n",
        "        break\n"
      ]
    }
  ],
  "metadata": {
    "_change_revision": 0,
    "_is_fork": false,
    "kernelspec": {
      "display_name": "Python 3",
      "language": "python",
      "name": "python3"
    },
    "language_info": {
      "codemirror_mode": {
        "name": "ipython",
        "version": 3
      },
      "file_extension": ".py",
      "mimetype": "text/x-python",
      "name": "python",
      "nbconvert_exporter": "python",
      "pygments_lexer": "ipython3",
      "version": "3.5.2"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 0
}