{"cells":[{"metadata":{"_uuid":"f09542bd1bf5bc3ee2bc3bd7b931fbb7cffe480f"},"cell_type":"markdown","source":"# Overview\nThis is a very simple (very bad) model for estimating the systolic volume directly from the images. As the dataset doesn't have any segmentations all we can do here is try and predict the volume based on a time series of images. As a very simple model, we use 3D convolutions to combine the different frames (first axis) and spatial (second and thirds axis) together and extract meaningful information from the image\n\n## Note\nThe train/test split here is very poor so please come up with a better validation strategy before experimenting too much"},{"metadata":{"trusted":true,"_uuid":"ec860566ccc754297c9c676c2e7adc9ea42d9c68","collapsed":true},"cell_type":"code","source":"import os\nimport h5py\nimport matplotlib.pyplot as plt\nfrom skimage.util.montage import montage2d\nfrom sklearn.preprocessing import LabelEncoder\nfrom keras.utils.np_utils import to_categorical\nimport numpy as np\nimport gc\ngc.enable() # we come close to the memory limits and this seems to minimize kernel resets\nmontage3d = lambda x, **k: montage2d(np.stack([montage2d(y, **k) for y in x],0))\ndata_dir = '../input/mri-heart-processing/'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"73b93122fd63f98eba8ad7e05c5b75e3ca325265"},"cell_type":"code","source":"with h5py.File(os.path.join(data_dir, 'train_mri_128_128.h5'), 'r') as w:\n    full_data = w['image'].value\n    n_group = w['id'].value\n    n_scalar = w['area_multiplier'].value\n    y_target = w['systole'].value / n_scalar # remove the area scalar since we dont have this in the images","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"e046b323cbcbf9452c59f6e79b212137f72ad75d","collapsed":true},"cell_type":"code","source":"y_target.min(), y_target.max(), y_target.mean()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"c12507086a3eb0d8a1a6b8798ce0721a02937794"},"cell_type":"code","source":"offset_value = 50\nscale_factor = 100","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"f511db258ffe441cf6a06cb4a3fea54a6b4dbdaf","collapsed":true},"cell_type":"code","source":"y_target_class = ((y_target-offset_value)/scale_factor).clip(-1.5,1.5).reshape((-1,1))\n_ = plt.hist(y_target_class)\ny_target_class.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"1231bfa4919adaab139ee66359b5477bce80e5d6","collapsed":true},"cell_type":"code","source":"# instance normalization\nsafe_norm_func = lambda x: np.clip((x-x.mean())/(0.1+x.std()), -2, 2)\nnorm_ch_x_data = np.expand_dims(np.apply_along_axis(safe_norm_func, 0, full_data), -1) # add channels\ndel full_data\nnorm_ch_x_data.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"51c578a1eb535261bd996d9e12330f588b96190c","collapsed":true},"cell_type":"code","source":"fig, ax1 = plt.subplots(1,1, figsize = (8,8))\nax1.imshow(montage3d(norm_ch_x_data[np.random.choice(norm_ch_x_data.shape[0], size = 4), :, :, :, 0]))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"beba8596232d67fbac3c6a2fadbbd09fb9f103b9","collapsed":true},"cell_type":"code","source":"%matplotlib inline\nplt.hist(norm_ch_x_data[:5].ravel())","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"63428fc05f1b07b4bd048dd100318dc8eae90481"},"cell_type":"markdown","source":"# Build the Model\nHere we make a simple sequential model for processing the MRI frames and estimating the systolic volume"},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"6b6058c5aae5f3be946515a3d76cd9f6eb8b375d"},"cell_type":"code","source":"from keras.models import Sequential\nfrom keras.layers import SpatialDropout3D, Dropout, Activation\nfrom keras.layers import Conv3D, BatchNormalization, Dense, Flatten, Reshape, GlobalAveragePooling3D, MaxPooling3D","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"998797558a7fea009c6385b1b752e4837ba7d53d"},"cell_type":"code","source":"in_shape = norm_ch_x_data.shape[1:]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"f915d69a813d5f3758f8ef6c79f5d3f8f9b5d1ea","collapsed":true},"cell_type":"code","source":"simple_model = Sequential()\nsimple_model.add(Conv3D(filters = 8, \n                        kernel_size = (3,1,1), \n                        input_shape = in_shape, \n                        activation = 'linear',\n                        strides = (2, 1, 1),\n                       use_bias = False))\nsimple_model.add(BatchNormalization())\nsimple_model.add(Activation('relu'))\nsimple_model.add(Conv3D(filters = 16, \n                        kernel_size = (3,1,1), \n                        input_shape = in_shape, \n                        activation = 'linear',\n                        strides = (2, 1, 1),\n                       use_bias = False))\nsimple_model.add(BatchNormalization())\nsimple_model.add(Activation('relu'))\nsimple_model.add(Conv3D(filters = 64, kernel_size = (3,3,3)))\nsimple_model.add(Conv3D(filters = 64, kernel_size = (1,3,3)))\nsimple_model.add(MaxPooling3D((1,2,2)))\nsimple_model.add(Conv3D(filters = 128, kernel_size = (1,3,3)))\nsimple_model.add(Conv3D(filters = 128, kernel_size = (1,3,3)))\nsimple_model.add(MaxPooling3D((1,2,2)))\nsimple_model.add(Conv3D(filters = 256, kernel_size = (1,3,3)))\nsimple_model.add(MaxPooling3D((1,2,2)))\nsimple_model.add(Conv3D(filters = 512, kernel_size = (1,3,3)))\nsimple_model.add(Conv3D(filters = 1024, kernel_size = (3,1,1)))\nsimple_model.add(SpatialDropout3D(0.5))\nsimple_model.add(GlobalAveragePooling3D())\nsimple_model.add(Dense(256))\nsimple_model.add(Dropout(0.5))\nsimple_model.add(Dense(y_target_class.shape[1], activation = 'tanh'))\nsimple_model.summary()","execution_count":null,"outputs":[]},{"metadata":{"collapsed":true,"trusted":true,"_uuid":"90de080f7d9240a2b75d91b1e0c301c41f9b1f80"},"cell_type":"code","source":"from keras.optimizers import Adam\nsimple_model.compile(loss = 'mse', \n                     optimizer = Adam(1e-4, decay = 1e-6), \n                     metrics = ['mae'])\nloss_history = []","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"932a3abb66f6aa9d25796580a1adf0550c844d2b","collapsed":true},"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n# a simpler one is better here since the same patients are spread over multiple slices and we want to minimize leak without making too much hassle\ndef train_test_split(x, y, train_size, random_state):\n    last_train_idx = int(train_size*x.shape[0])\n    return x[:last_train_idx], x[last_train_idx+1:], y[:last_train_idx], y[last_train_idx+1:]\nfrom keras.utils.np_utils import to_categorical\nX_train, X_test, y_train, y_test = train_test_split(norm_ch_x_data, y_target_class, \n                                                   train_size = 0.7,\n                                                   random_state = 2017)\ndel norm_ch_x_data","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"61404276814a5ab844f968aa7d0ad24c000a076b","collapsed":true},"cell_type":"code","source":"from keras.callbacks import ModelCheckpoint, LearningRateScheduler, EarlyStopping, ReduceLROnPlateau\nweight_path=\"{}_weights.best.hdf5\".format('systole_model')\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.5, \n                                   patience=3, \n                                   verbose=1, mode='min', epsilon=0.0001, cooldown=2, min_lr=1e-6)\nearly = EarlyStopping(monitor=\"val_loss\", \n                      mode=\"min\", \n                      patience=15) # probably needs to be more patient, but kaggle time is limited\ncallbacks_list = [checkpoint, early, reduceLROnPlat]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"8ee0c6ba8d7794cf3497bd99ae0d1b854a80289a","collapsed":true},"cell_type":"code","source":"loss_history += [simple_model.fit(X_train, y_train, \n          validation_data=(X_test, y_test),\n                           shuffle = True,\n                           batch_size = 32,\n                           epochs = 30,\n                                 callbacks = callbacks_list)]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"ff135321c9df565fbf41260335df0081125cddba"},"cell_type":"code","source":"simple_model.load_weights(weight_path)\nsimple_model.save('full_systolic_model.h5')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"e482eccba784afe0c0a32c8725bffb47747a81f2","collapsed":true},"cell_type":"code","source":"pred_test = simple_model.predict(X_test, verbose = 1)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"2d13de03649a93c6707954b34052d0b6eb87c088"},"cell_type":"markdown","source":"# Comparing Predictions to Real Values\nHere we compare the predictions to the real values on the test data. We ideally see a line indicating perfect correlation between the two datasets."},{"metadata":{"trusted":true,"_uuid":"534d5c8a82835f517ad9c9da6e2690703165b423","collapsed":true},"cell_type":"code","source":"fig, (ax1) = plt.subplots(1,1, figsize = (8, 8))\nax1.scatter(y_test, pred_test)\nax1.plot(y_test, y_test, 'r-')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"5fbe62589e1085b4b07f7b79b660f98f83bf9eed","collapsed":true},"cell_type":"code","source":"for v, f in zip(simple_model.evaluate(X_test, y_test, verbose = 1), \n                simple_model.metrics_names):\n    print('{}: difference - {:2.2f} ml'.format(f, scale_factor*v))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"40f23dbc5f765566e217fd6afea8e070cb3c8aab"},"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.4","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}