From 09ec44fa80fdd3604d2261370ab53319d8545b9e Mon Sep 17 00:00:00 2001 From: Somshubra Majumdar Date: Sat, 18 Nov 2017 13:38:12 -0600 Subject: [PATCH] Improve auxilary head --- keras_contrib/applications/nasnet.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/keras_contrib/applications/nasnet.py b/keras_contrib/applications/nasnet.py index 15e7247..cddc682 100644 --- a/keras_contrib/applications/nasnet.py +++ b/keras_contrib/applications/nasnet.py @@ -195,7 +195,8 @@ def NASNet(input_shape=None, auxilary_x = None if use_auxilary_branch: - img_dim = 2 if K.image_data_format() == 'channels_first' else -2 + img_height = 1 if K.image_data_format() == 'channels_first' else 2 + img_width = 2 if K.image_data_format() == 'channels_first' else 3 with K.name_scope('auxilary_branch'): auxilary_x = Activation('relu')(x) @@ -206,8 +207,9 @@ def NASNet(input_shape=None, name='aux_bn_projection')(auxilary_x) auxilary_x = Activation('relu')(auxilary_x) - auxilary_x = Conv2D(768, auxilary_x._keras_shape[img_dim], padding='valid', use_bias=False, - kernel_initializer='he_normal', name='aux_conv_reduction')(auxilary_x) + auxilary_x = Conv2D(768, (auxilary_x._keras_shape[img_height], auxilary_x._keras_shape[img_width]), + padding='valid', use_bias=False, kernel_initializer='he_normal', + name='aux_conv_reduction')(auxilary_x) auxilary_x = BatchNormalization(axis=channel_dim, momentum=_BN_DECAY, epsilon=_BN_EPSILON, name='aux_bn_reduction')(auxilary_x) auxilary_x = Activation('relu')(auxilary_x)