{"cells":[{"metadata":{"_uuid":"34b6997bb115a11f47a7f54ce9d9052791b2e707","_cell_guid":"c67e806c-e0e6-415b-8cf0-ed92eba5ed37"},"cell_type":"markdown","source":"# Overview\nThe goal is to make a nice retinopathy model by using a pretrained inception v3 as a base and retraining some modified final layers with attention\n\nThis can be massively improved with \n* high-resolution images\n* better data sampling\n* ensuring there is no leaking between training and validation sets, ```sample(replace = True)``` is real dangerous\n* better target variable (age) normalization\n* pretrained models\n* attention/related techniques to focus on areas"},{"metadata":{"_uuid":"c163e45042a69905855f7c04a65676e5aca4837b","_cell_guid":"e94de3e7-de1f-4c18-ad1f-c8b686127340","trusted":true,"collapsed":true},"cell_type":"code","source":"# copy the weights and configurations for the pre-trained models\n!mkdir ~/.keras\n!mkdir ~/.keras/models\n!cp ../input/keras-pretrained-models/*notop* ~/.keras/models/\n!cp ../input/keras-pretrained-models/imagenet_class_index.json ~/.keras/models/","execution_count":9,"outputs":[]},{"metadata":{"_uuid":"725d378daf5f836d4885d67240fc7955f113309d","collapsed":true,"_cell_guid":"c3cc4285-bfa4-4612-ac5f-13d10678c09a","trusted":true},"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt # showing and rendering figures\n# io related\nfrom skimage.io import imread\nimport os\nfrom glob import glob\n# not needed in Kaggle, but required in Jupyter\n# matplotlib inline ","execution_count":1,"outputs":[]},{"metadata":{"_uuid":"346da81db6ee7a34af8da8af245b42e681f2ba48","_cell_guid":"c4b38df6-ffa1-4847-b605-511e72b68231","trusted":true,"collapsed":true},"cell_type":"code","source":"base_image_dir = os.path.join('..', 'input', 'diabetic-retinopathy-detection')\nretina_df = pd.read_csv(os.path.join(base_image_dir, 'trainLabels.csv'))\nretina_df['PatientId'] = retina_df['image'].map(lambda x: x.split('_')[0])\nretina_df['path'] = retina_df['image'].map(lambda x: os.path.join(base_image_dir,\n                                                         '{}.jpeg'.format(x)))\nretina_df['exists'] = retina_df['path'].map(os.path.exists)\nprint(retina_df['exists'].sum(), 'images found of', retina_df.shape[0], 'total')\nretina_df['eye'] = retina_df['image'].map(lambda x: 1 if x.split('_')[-1]=='left' else 0)\nfrom keras.utils.np_utils import to_categorical\nretina_df['level_cat'] = retina_df['level'].map(lambda x: to_categorical(x, 1+retina_df['level'].max()))\n\nretina_df.dropna(inplace = True)\nretina_df = retina_df[retina_df['exists']]\nretina_df.sample(3)","execution_count":2,"outputs":[]},{"metadata":{"_uuid":"688e4340238e013b8459b6f6470993c7de492d83","_cell_guid":"818da6ca-bbff-4ca0-ad57-ef3a145ae863"},"cell_type":"markdown","source":"# Examine the distribution of eye and severity"},{"metadata":{"_uuid":"60a8111c4093ca6f69d27a4499442ba7dd750839","_cell_guid":"5c8bd288-8261-4cbe-a954-e62ac795cc3e","trusted":true,"collapsed":true},"cell_type":"code","source":"retina_df[['level', 'eye']].hist(figsize = (10, 5))","execution_count":4,"outputs":[]},{"metadata":{"_uuid":"4df45776bae0b8a1bf9d3eb4eaaebce6e24d726d","_cell_guid":"0ba697ed-85bb-4e9a-9765-4c367db078d1"},"cell_type":"markdown","source":"# Split Data into Training and Validation"},{"metadata":{"_uuid":"a48b300ca4d37a6e8b39f82e3c172739635e4baa","_cell_guid":"1192c6b3-a940-4fa0-a498-d7e0d400a796","trusted":true,"collapsed":true},"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nrr_df = retina_df[['PatientId', 'level']].drop_duplicates()\ntrain_ids, valid_ids = train_test_split(rr_df['PatientId'], \n                                   test_size = 0.25, \n                                   random_state = 2018,\n                                   stratify = rr_df['level'])\nraw_train_df = retina_df[retina_df['PatientId'].isin(train_ids)]\nvalid_df = retina_df[retina_df['PatientId'].isin(valid_ids)]\nprint('train', raw_train_df.shape[0], 'validation', valid_df.shape[0])\n# dropna handle missing data\n# drop_duplicates 去除重复项","execution_count":5,"outputs":[]},{"metadata":{"_uuid":"26e566d6cec5bd41f9afe392f456ddf7ceb306ea","_cell_guid":"f8060459-da1e-4293-8f61-c7f99de1de9f"},"cell_type":"markdown","source":"# Balance the distribution in the training set"},{"metadata":{"_uuid":"ba7befa238b8c11f9672e3539ac58f3da6955bd9","_cell_guid":"7a130199-fbf6-4c60-95f5-0797b2f3eaf1","trusted":true,"collapsed":true},"cell_type":"code","source":"train_df = raw_train_df.groupby(['level', 'eye']).apply(lambda x: x.sample(75, replace = True)\n                                                      ).reset_index(drop = True)\nprint('New Data Size:', train_df.shape[0], 'Old Size:', raw_train_df.shape[0])\ntrain_df[['level', 'eye']].hist(figsize = (10, 5))","execution_count":6,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"b1e2be537329bc2958def2034f575a2ed983df25"},"cell_type":"code","source":"import tensorflow as tf\nfrom keras import backend as K\nfrom keras.applications.inception_v3 import preprocess_input\nimport numpy as np\nIMG_SIZE = (512, 512) # slightly smaller than vgg16 normally expects\n# This part is used to handle batch process\n# tf.map_fn given the function and batch data, it will process every element in batch data\ndef tf_image_loader(out_size, \n                      horizontal_flip = True, \n                      vertical_flip = False, \n                     random_brightness = True,\n                     random_contrast = True,\n                    random_saturation = True,\n                    random_hue = True,\n                      color_mode = 'rgb',\n                       preproc_func = preprocess_input,\n                       on_batch = False):\n    def _func(X):\n        with tf.name_scope('image_augmentation'):\n            with tf.name_scope('input'):\n                X = tf.image.decode_png(tf.read_file(X), channels = 3 if color_mode == 'rgb' else 0)\n                X = tf.image.resize_images(X, out_size)\n            with tf.name_scope('augmentation'):\n                if horizontal_flip:\n                    X = tf.image.random_flip_left_right(X)\n                if vertical_flip:\n                    X = tf.image.random_flip_up_down(X)\n                if random_brightness:\n                    X = tf.image.random_brightness(X, max_delta = 0.1)\n                if random_saturation:\n                    X = tf.image.random_saturation(X, lower = 0.75, upper = 1.5)\n                if random_hue:\n                    X = tf.image.random_hue(X, max_delta = 0.15)\n                if random_contrast:\n                    X = tf.image.random_contrast(X, lower = 0.75, upper = 1.5)\n                return preproc_func(X)\n    if on_batch: \n        # we are meant to use it on a batch\n        def _batch_func(X, y):\n            return tf.map_fn(_func, X), y\n        return _batch_func\n    else:\n        # we apply it to everything\n        def _all_func(X, y):\n            return _func(X), y         \n        return _all_func","execution_count":8,"outputs":[]},{"metadata":{"_uuid":"9529ab766763a9f122786464c24ab1ebe22c6006","_cell_guid":"9954bfda-29bd-4c4d-b526-0a972b3e43e2","trusted":true,"collapsed":true},"cell_type":"code","source":"def tf_augmentor(out_size,\n                intermediate_size = (640, 640),\n                 intermediate_trans = 'crop',\n                 batch_size = 16,\n                   horizontal_flip = True, \n                  vertical_flip = False, \n                 random_brightness = True,\n                 random_contrast = True,\n                 random_saturation = True,\n                    random_hue = True,\n                  color_mode = 'rgb',\n                   preproc_func = preprocess_input,\n                   min_crop_percent = 0.001,\n                   max_crop_percent = 0.005,\n                   crop_probability = 0.5,\n                   rotation_range = 10):\n    \n    load_ops = tf_image_loader(out_size = intermediate_size, \n                               horizontal_flip=horizontal_flip, \n                               vertical_flip=vertical_flip, \n                               random_brightness = random_brightness,\n                               random_contrast = random_contrast,\n                               random_saturation = random_saturation,\n                               random_hue = random_hue,\n                               color_mode = color_mode,\n                               preproc_func = preproc_func,\n                               on_batch=False)\n    def batch_ops(X, y):\n        batch_size = tf.shape(X)[0]\n        with tf.name_scope('transformation'):\n            # code borrowed from https://becominghuman.ai/data-augmentation-on-gpu-in-tensorflow-13d14ecf2b19\n            # The list of affine transformations that our image will go under.\n            # Every element is Nx8 tensor, where N is a batch size.\n            transforms = []\n            identity = tf.constant([1, 0, 0, 0, 1, 0, 0, 0], dtype=tf.float32)\n            if rotation_range > 0:\n                angle_rad = rotation_range / 180 * np.pi\n                angles = tf.random_uniform([batch_size], -angle_rad, angle_rad)\n                transforms += [tf.contrib.image.angles_to_projective_transforms(angles, intermediate_size[0], intermediate_size[1])]\n                #angles, image_height, image_width,\n                #Returns: A tensor of shape (num_images, 8). Projective transforms which can be given to tf.contrib.image.transform.\n\n            if crop_probability > 0:\n                crop_pct = tf.random_uniform([batch_size], min_crop_percent, max_crop_percent)\n                left = tf.random_uniform([batch_size], 0, intermediate_size[0] * (1.0 - crop_pct))\n                top = tf.random_uniform([batch_size], 0, intermediate_size[1] * (1.0 - crop_pct))\n                # tf.stack 在新的维度上合并张量\n                # (num_images, 8) 的张量\n                crop_transform = tf.stack([\n                      crop_pct,\n                      tf.zeros([batch_size]), top,\n                      tf.zeros([batch_size]), crop_pct, left,\n                      tf.zeros([batch_size]),\n                      tf.zeros([batch_size])\n                  ], 1)\n                coin = tf.less(tf.random_uniform([batch_size], 0, 1.0), crop_probability)  # less 若第一个元素较小则返回true，相当于掷硬币，概率为crop_probability\n                transforms += [tf.where(coin, crop_transform, tf.tile(tf.expand_dims(identity, 0), [batch_size, 1]))]\n                # tf.where(condition, x=None, y=None, name=None) x (if true) or y (if false).\n                # tf.tile The output tensor's i'th dimension has input.dims(i) * multiples[i] elements, and the values of input are replicated multiples[i] times along the 'i'th dimension. For example, tiling [a b c d] by [2] produces [a b c d a b c d].\n                # tf.ex\n            if len(transforms)>0:\n                X = tf.contrib.image.transform(X,\n                      tf.contrib.image.compose_transforms(*transforms),\n                      interpolation='BILINEAR') # or 'NEAREST'\n            if intermediate_trans=='scale':\n                X = tf.image.resize_images(X, out_size)\n            elif intermediate_trans=='crop':\n                X = tf.image.resize_image_with_crop_or_pad(X, out_size[0], out_size[1])\n            else:\n                raise ValueError('Invalid Operation {}'.format(intermediate_trans))\n            return X, y\n    def _create_pipeline(in_ds):\n        batch_ds = in_ds.map(load_ops, num_parallel_calls=4).batch(batch_size)\n        return batch_ds.map(batch_ops)\n    # map: same operator to each element\n    return _create_pipeline","execution_count":9,"outputs":[]},{"metadata":{"_uuid":"07851e798db3d89ba13db7d4b56ab2b759221464","collapsed":true,"_cell_guid":"b5767f42-da63-4737-8f50-749c1a25aa84","trusted":true},"cell_type":"code","source":"def flow_from_dataframe(idg, \n                        in_df, \n                        path_col,\n                        y_col, \n                        shuffle = True, \n                        color_mode = 'rgb'):\n    files_ds = tf.data.Dataset.from_tensor_slices((in_df[path_col].values,          #利用tf.data.Dataset.from_tensor_slices创建每个元素是一个tuple的dataset也是可以的：\n                                                   np.stack(in_df[y_col].values,0)))\n    in_len = in_df[path_col].values.shape[0]\n    while True:\n        if shuffle:\n            files_ds = files_ds.shuffle(in_len) # shuffle the whole dataset\n        \n        next_batch = idg(files_ds).repeat().make_one_shot_iterator().get_next()\n        for i in range(max(in_len//32,1)):\n            # NOTE: if we loop here it is 'thread-safe-ish' if we loop on the outside it is completely unsafe\n            yield K.get_session().run(next_batch)\n            \n#这里使用了较难理解的回调函数编程\n#1.定义函数：tf_image_loader。返回的是函数 _batch_func 或 _all_func。返回的函数，进一步返回的是主体操作 _func,如果是批处理，那么使用 map_fn处理输入的X中的每一个。最终输出： 处理过后的X 和 没有处理的 y。\n#2.定义函数：tf_augmentor。返回的是函数 _create_pipeline。_create_pipeline 第一步对提供给_create_pipeline的数据的每一个元素进行tf_image_loader的操作，也就是1中的结果。第二步，对 每个元素做 batch_ops 操作,也就是透视变换。\n#3.定义函数 flow_from_dataframe。idg是tf_augmentor返回的分别是一个_create_pipeline， in_df是pandas list。path_col, y_col, 是数据集dataframe中两列。","execution_count":10,"outputs":[]},{"metadata":{"_uuid":"1848f5048a9e00668c3778a85deea97f980e4f1c","_cell_guid":"810bd229-fec9-43c4-b3bd-afd62e3e9552","trusted":true,"collapsed":true},"cell_type":"code","source":"batch_size = 48\ncore_idg = tf_augmentor(out_size = IMG_SIZE, \n                        color_mode = 'rgb', \n                        vertical_flip = True,\n                        crop_probability=0.0, # crop doesn't work yet\n                        batch_size = batch_size) \nvalid_idg = tf_augmentor(out_size = IMG_SIZE, color_mode = 'rgb', \n                         crop_probability=0.0, \n                         horizontal_flip = False, \n                         vertical_flip = False, \n                         random_brightness = False,\n                         random_contrast = False,\n                         random_saturation = False,\n                         random_hue = False,\n                         rotation_range = 0,\n                        batch_size = batch_size)\n\ntrain_gen = flow_from_dataframe(core_idg, train_df, \n                             path_col = 'path',\n                            y_col = 'level_cat')\n\nvalid_gen = flow_from_dataframe(valid_idg, valid_df, \n                             path_col = 'path',\n                            y_col = 'level_cat') # we can use much larger batches for evaluation\n#core_idg 和 Valid_idg 分别是一个_create_pipeline 函数。\n#train_df 和 valid_df 是上文处理好的训练集和验证集。","execution_count":11,"outputs":[]},{"metadata":{"_uuid":"2ad184b936d5cebae91a265a247d8e0e25920566"},"cell_type":"markdown","source":"# Validation Set\nWe do not perform augmentation at all on these images"},{"metadata":{"trusted":true,"_uuid":"6810407e25b887dd8b352f1e46fb3faceaa58ab7","collapsed":true},"cell_type":"code","source":"t_x, t_y = next(valid_gen)\nfig, m_axs = plt.subplots(2, 4, figsize = (16, 8))\nfor (c_x, c_y, c_ax) in zip(t_x, t_y, m_axs.flatten()):\n    c_ax.imshow(np.clip(c_x*127+127, 0, 255).astype(np.uint8))\n    c_ax.set_title('Severity {}'.format(np.argmax(c_y, -1)))\n    c_ax.axis('off')\n#可迭代对象：Iterable\n# generator 不但可以作用于for循环，还可以被next()函数不断调用并返回下一个值，直到最后抛出StopIteration错误表示无法继续返回下一个值了","execution_count":12,"outputs":[]},{"metadata":{"_uuid":"34ce892a19c9734511e2da1d0f2552b361dc826d"},"cell_type":"markdown","source":"# Training Set\nThese are augmented and a real mess"},{"metadata":{"_uuid":"8190b4ad60d49fa65af074dd138a19cb8787e983","scrolled":true,"_cell_guid":"2d62234f-aeb0-4eba-8a38-d713d819abf6","trusted":true,"collapsed":true},"cell_type":"code","source":"","execution_count":13,"outputs":[]},{"metadata":{"_uuid":"55d665e1e8a8d83b9db005a66a965f8a90c62da1","_cell_guid":"da22790a-672c-474e-b118-9eef15b53160"},"cell_type":"markdown","source":"# Attention Model\nThe basic idea is that a Global Average Pooling is too simplistic since some of the regions are more relevant than others. So we build an attention mechanism to turn pixels in the GAP on an off before the pooling and then rescale (Lambda layer) the results based on the number of pixels. The model could be seen as a sort of 'global weighted average' pooling. There is probably something published about it and it is very similar to the kind of attention models used in NLP.\nIt is largely based on the insight that the winning solution annotated and trained a UNET model to segmenting the hand and transforming it. This seems very tedious if we could just learn attention."},{"metadata":{"_uuid":"1f0dfaccda346d7bc4758e7329d61028d254a8d6","_cell_guid":"eeb36110-0cde-4450-a43c-b8f707adb235","trusted":true,"collapsed":true},"cell_type":"code","source":"from keras.applications.vgg16 import VGG16 as PTModel\nfrom keras.applications.inception_resnet_v2 import InceptionResNetV2 as PTModel\nfrom keras.applications.inception_v3 import InceptionV3 as PTModel\nfrom keras.layers import GlobalAveragePooling2D, Dense, Dropout, Flatten, Input, Conv2D, multiply, LocallyConnected2D, Lambda\nfrom keras.models import Model\nin_lay = Input(t_x.shape[1:])\nbase_pretrained_model = PTModel(input_shape =  t_x.shape[1:], include_top = False, weights = 'imagenet')\nbase_pretrained_model.trainable = False\npt_depth = base_pretrained_model.get_output_shape_at(0)[-1]\npt_features = base_pretrained_model(in_lay)\nfrom keras.layers import BatchNormalization\nbn_features = BatchNormalization()(pt_features)\n\n# here we do an attention mechanism to turn pixels in the GAP on an off\n\nattn_layer = Conv2D(64, kernel_size = (1,1), padding = 'same', activation = 'relu')(Dropout(0.5)(bn_features))\nattn_layer = Conv2D(16, kernel_size = (1,1), padding = 'same', activation = 'relu')(attn_layer)\nattn_layer = Conv2D(8, kernel_size = (1,1), padding = 'same', activation = 'relu')(attn_layer)\nattn_layer = Conv2D(1, \n                    kernel_size = (1,1), \n                    padding = 'valid', \n                    activation = 'sigmoid')(attn_layer)\n# fan it out to all of the channels\nup_c2_w = np.ones((1, 1, 1, pt_depth))\nup_c2 = Conv2D(pt_depth, kernel_size = (1,1), padding = 'same', \n               activation = 'linear', use_bias = False, weights = [up_c2_w])\nup_c2.trainable = False\nattn_layer = up_c2(attn_layer)\n\nmask_features = multiply([attn_layer, bn_features])\ngap_features = GlobalAveragePooling2D()(mask_features)\ngap_mask = GlobalAveragePooling2D()(attn_layer)\n# to account for missing values from the attention model\ngap = Lambda(lambda x: x[0]/x[1], name = 'RescaleGAP')([gap_features, gap_mask])\ngap_dr = Dropout(0.25)(gap)\ndr_steps = Dropout(0.25)(Dense(128, activation = 'relu')(gap_dr))\nout_layer = Dense(t_y.shape[-1], activation = 'softmax')(dr_steps)\nretina_model = Model(inputs = [in_lay], outputs = [out_layer])\nfrom keras.metrics import top_k_categorical_accuracy\ndef top_2_accuracy(in_gt, in_pred):\n    return top_k_categorical_accuracy(in_gt, in_pred, k=2)\n\nretina_model.compile(optimizer = 'adam', loss = 'categorical_crossentropy',\n                           metrics = ['categorical_accuracy', top_2_accuracy])\nretina_model.summary()\n# Keras 提供的applications中的网络具有预训练权值。","execution_count":17,"outputs":[]},{"metadata":{"_uuid":"48b9764e16fb5af52aed35c82bae6299e67d5bc7","_cell_guid":"17803ae1-bed8-41a4-9a2c-e66287a24830","trusted":true,"collapsed":true},"cell_type":"code","source":"from keras.callbacks import ModelCheckpoint, LearningRateScheduler, EarlyStopping, ReduceLROnPlateau\nweight_path=\"{}_weights.best.hdf5\".format('retina')\n\ncheckpoint = ModelCheckpoint(weight_path, monitor='val_loss', verbose=1, \n                             save_best_only=True, mode='min', save_weights_only = True)\n\nreduceLROnPlat = ReduceLROnPlateau(monitor='val_loss', factor=0.8, patience=3, verbose=1, mode='auto', epsilon=0.0001, cooldown=5, min_lr=0.0001)\nearly = EarlyStopping(monitor=\"val_loss\", \n                      mode=\"min\", \n                      patience=6) # probably needs to be more patient, but kaggle time is limited\ncallbacks_list = [checkpoint, early, reduceLROnPlat]\n# 回调函数是一组在训练的特定阶段被调用的函数集，你可以使用回调函数来观察训练过程中网络内部的状态和统计信息。通过传递回调函数列表到模型的.fit()中，即可在给定的训练阶段调用该函数集中的函数。\n# 当评价指标不在提升时，减少学习率 当学习停滞时，减少2倍或10倍的学习率常常能获得较好的效果","execution_count":13,"outputs":[]},{"metadata":{"_uuid":"78dfa383c51777377c1f81e42017cbcca5f5736f","collapsed":true,"_cell_guid":"84f7cdec-ca00-460c-9991-55b1f7f02f20","trusted":true},"cell_type":"code","source":"!rm -rf ~/.keras # clean up before starting training","execution_count":14,"outputs":[]},{"metadata":{"_uuid":"b2148479bfe41c5d9fd0faece4c75adea509dabe","_cell_guid":"58a75586-442b-4804-84a6-d63d5a42ea14","trusted":true,"collapsed":true},"cell_type":"code","source":"retina_model.fit_generator(train_gen, \n                           steps_per_epoch = train_df.shape[0]//batch_size,\n                           validation_data = valid_gen, \n                           validation_steps = valid_df.shape[0]//batch_size,\n                              epochs = 25, \n                              callbacks = callbacks_list,\n                             workers = 0, # tf-generators are not thread-safe\n                             use_multiprocessing=False, \n                             max_queue_size = 0\n                            )","execution_count":39,"outputs":[]},{"metadata":{"_uuid":"3a90f05dd206cd76c72d8c6278ebb93da41ee45f","collapsed":true,"_cell_guid":"4d0c45b0-bb23-48d2-83eb-bc3990043e26","trusted":true},"cell_type":"code","source":"# load the best version of the model\nretina_model.load_weights(weight_path)\nretina_model.save('full_retina_model.h5')","execution_count":40,"outputs":[]},{"metadata":{"_uuid":"2b74f4ab850c6e82549d732b6f0524724b95b53c","_cell_guid":"f37dd4d8-ecd6-487a-90d8-74fe14a9a318","trusted":true,"collapsed":true},"cell_type":"code","source":"##### create one fixed dataset for evaluating\nfrom tqdm import tqdm_notebook\n# fresh valid gen\nvalid_gen = flow_from_dataframe(valid_idg, valid_df, \n                             path_col = 'path',\n                            y_col = 'level_cat') \nvbatch_count = (valid_df.shape[0]//batch_size-1)\nout_size = vbatch_count*batch_size\ntest_X = np.zeros((out_size,)+t_x.shape[1:], dtype = np.float32)\ntest_Y = np.zeros((out_size,)+t_y.shape[1:], dtype = np.float32)\nfor i, (c_x, c_y) in zip(tqdm_notebook(range(vbatch_count)), \n                         valid_gen):\n    j = i*batch_size\n    test_X[j:(j+c_x.shape[0])] = c_x\n    test_Y[j:(j+c_x.shape[0])] = c_y","execution_count":41,"outputs":[]},{"metadata":{"_uuid":"cca170eb40bc591f89748ede8aa35de4308faaaf","_cell_guid":"11f33f0a-61eb-488a-b7ea-4bc9d15ba8f9"},"cell_type":"markdown","source":"# Show Attention\nDid our attention model learn anything useful?"},{"metadata":{"_uuid":"ad5b085d351e79b950bf0c2ddc476799d5b0692f","_cell_guid":"e41a063f-35c9-410f-be63-f66b63ff9683","trusted":true,"collapsed":true},"cell_type":"code","source":"# get the attention layer since it is the only one with a single output dim\nfor attn_layer in retina_model.layers:\n    c_shape = attn_layer.get_output_shape_at(0)\n    if len(c_shape)==4:\n        if c_shape[-1]==1:\n            print(attn_layer)\n            break","execution_count":42,"outputs":[]},{"metadata":{"_uuid":"00850972ae4298f49ed1838b3fc49c2d8fb07547","_cell_guid":"340eef36-f5b2-4b15-a59f-440061a427eb","trusted":true,"collapsed":true},"cell_type":"code","source":"import keras.backend as K\nrand_idx = np.random.choice(range(len(test_X)), size = 6)\nattn_func = K.function(inputs = [retina_model.get_input_at(0), K.learning_phase()],\n           outputs = [attn_layer.get_output_at(0)]\n          )\nfig, m_axs = plt.subplots(len(rand_idx), 2, figsize = (8, 4*len(rand_idx)))\n[c_ax.axis('off') for c_ax in m_axs.flatten()]\nfor c_idx, (img_ax, attn_ax) in zip(rand_idx, m_axs):\n    cur_img = test_X[c_idx:(c_idx+1)]\n    attn_img = attn_func([cur_img, 0])[0]\n    img_ax.imshow(np.clip(cur_img[0,:,:,:]*127+127, 0, 255).astype(np.uint8))\n    attn_ax.imshow(attn_img[0, :, :, 0]/attn_img[0, :, :, 0].max(), cmap = 'viridis', \n                   vmin = 0, vmax = 1, \n                   interpolation = 'lanczos')\n    real_cat = np.argmax(test_Y[c_idx, :])\n    img_ax.set_title('Eye Image\\nCat:%2d' % (real_cat))\n    pred_cat = retina_model.predict(cur_img)\n    attn_ax.set_title('Attention Map\\nPred:%2.2f%%' % (100*pred_cat[0,real_cat]))\nfig.savefig('attention_map.png', dpi = 300)","execution_count":43,"outputs":[]},{"metadata":{"_uuid":"244bac80d1ea2074e47932e367996e32cbab6a3d","_cell_guid":"24796de7-b1e9-4b3b-bcc6-d997aa3e6d16"},"cell_type":"markdown","source":"# Evaluate the results\nHere we evaluate the results by loading the best version of the model and seeing how the predictions look on the results. We then visualize spec"},{"metadata":{"_uuid":"b421b6183b1919a7414482f0b1ac611079e45174","_cell_guid":"d0edaf00-4b7c-4f65-af0b-e5a03b9b8428","trusted":true,"collapsed":true},"cell_type":"code","source":"from sklearn.metrics import accuracy_score, classification_report\npred_Y = retina_model.predict(test_X, batch_size = 32, verbose = True)\npred_Y_cat = np.argmax(pred_Y, -1)\ntest_Y_cat = np.argmax(test_Y, -1)\nprint('Accuracy on Test Data: %2.2f%%' % (accuracy_score(test_Y_cat, pred_Y_cat)))\nprint(classification_report(test_Y_cat, pred_Y_cat))","execution_count":44,"outputs":[]},{"metadata":{"_uuid":"10162e055ca7cd52878a289bab377231787ab732","_cell_guid":"15189df2-3fed-495e-9661-97bb2b712dfd","trusted":true,"collapsed":true},"cell_type":"code","source":"import seaborn as sns\nfrom sklearn.metrics import confusion_matrix\nsns.heatmap(confusion_matrix(test_Y_cat, pred_Y_cat), \n            annot=True, fmt=\"d\", cbar = False, cmap = plt.cm.Blues, vmax = test_X.shape[0]//16)","execution_count":45,"outputs":[]},{"metadata":{"_uuid":"12dfe39ea80194062068589699953c6645e285d6","_cell_guid":"70827da6-bf91-4b65-80e9-bf1e6b885db3"},"cell_type":"markdown","source":"# ROC Curve for healthy vs sick\nHere we make an ROC curve for healthy (```severity == 0```) and sick (```severity>0```) to see how well the model works at just identifying the disease"},{"metadata":{"_uuid":"2b2aaee6c83043f721b0c9ed5bc4229eb7165200","_cell_guid":"829475ab-7db2-4421-b9ad-1a51971fd459","trusted":true,"collapsed":true},"cell_type":"code","source":"\nfrom sklearn.metrics import roc_curve, roc_auc_score\nsick_vec = test_Y_cat>0\nsick_score = np.sum(pred_Y[:,1:],1)\nfpr, tpr, _ = roc_curve(sick_vec, sick_score)\nfig, ax1 = plt.subplots(1,1, figsize = (6, 6), dpi = 150)\nax1.plot(fpr, tpr, 'b.-', label = 'Model Prediction (AUC: %2.2f)' % roc_auc_score(sick_vec, sick_score))\nax1.plot(fpr, fpr, 'g-', label = 'Random Guessing')\nax1.legend()\nax1.set_xlabel('False Positive Rate')\nax1.set_ylabel('True Positive Rate');","execution_count":46,"outputs":[]},{"metadata":{"_uuid":"ba87d0e7c3a77181487b99ca64d13de2aa8a21ee","scrolled":false,"_cell_guid":"c34f049f-b032-45bf-9d5e-a756ecc46a82","trusted":true,"collapsed":true},"cell_type":"code","source":"fig, m_axs = plt.subplots(2, 4, figsize = (32, 20))\nfor (idx, c_ax) in enumerate(m_axs.flatten()):\n    c_ax.imshow(np.clip(test_X[idx]*127+127,0 , 255).astype(np.uint8), cmap = 'bone')\n    c_ax.set_title('Actual Severity: {}\\n{}'.format(test_Y_cat[idx], \n                                                           '\\n'.join(['Predicted %02d (%04.1f%%): %s' % (k, 100*v, '*'*int(10*v)) for k, v in sorted(enumerate(pred_Y[idx]), key = lambda x: -1*x[1])])), loc='left')\n    c_ax.axis('off')\nfig.savefig('trained_img_predictions.png', dpi = 300)","execution_count":47,"outputs":[]},{"metadata":{"_uuid":"eb6752295030ba512263433f8383711e4ca1c14c","collapsed":true,"_cell_guid":"f2e189dc-f80a-4b16-bb1d-5c05a155a80b","trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.5","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}