From bdc64bf2f3b70398d1d7d7a0c3ca7eaa858089c9 Mon Sep 17 00:00:00 2001 From: Somshubra Majumdar Date: Wed, 22 Feb 2017 11:07:13 -0600 Subject: [PATCH] Added Batch Renormalization Layer (#28) * Began work on batch renormalization * Correct the implementation to use broadcasting properly, so it supports th dim ordering * Correct the implementation to use broadcasting properly, so it supports th dim ordering * Corrected normalizing term for d * Moved location of updates * Fixed the training problems and now auto adjusts the r_max and d_max values * Removed unused code comments * Finished implementation of batch renormalization * Corrected the get_config method * Added tests * Added user changeable parameter t_delta to change the rate at which time steps increase * Updated tests to state renorm everywhere * Attempt a fix for mode=1 test fail (giving nans) * Applied fix for nans to all calculations involving running mean and std * Corrected time delta default value * Reverts a mistake in pytest settings * Corrections to calculations of r and d due to wrong variable naming (running_std perserves running_variance) * Added t_delta to get_config * Fix commit change mixup * Corrected the check for uses_learning_flag * Made a few corrections, added moments to backend, added test to backend * Probable fix for test * Corrected backend test * Added support for mode 2 --- keras_contrib/backend/tensorflow_backend.py | 6 + keras_contrib/backend/theano_backend.py | 9 + keras_contrib/layers/normalization.py | 230 ++++++++++++++++++ tests/keras_contrib/backend/test_backends.py | 27 ++ .../layers/test_normalization.py | 111 +++++++++ 5 files changed, 383 insertions(+) diff --git a/keras_contrib/backend/tensorflow_backend.py b/keras_contrib/backend/tensorflow_backend.py index a7eaa99..ec02c69 100644 --- a/keras_contrib/backend/tensorflow_backend.py +++ b/keras_contrib/backend/tensorflow_backend.py @@ -99,3 +99,9 @@ def extract_image_patches(X, ksizes, ssizes, border_mode="same", dim_ordering="t if dim_ordering == "tf": patches = KTF.permute_dimensions(patches, [0, 1, 2, 4, 5, 3]) return patches + + +def moments(x, axes, shift=None, keep_dims=False): + ''' Wrapper over tensorflow backend call ''' + + return tf.nn.moments(x, axes, shift=shift, keep_dims=keep_dims) diff --git a/keras_contrib/backend/theano_backend.py b/keras_contrib/backend/theano_backend.py index e94f552..396b589 100644 --- a/keras_contrib/backend/theano_backend.py +++ b/keras_contrib/backend/theano_backend.py @@ -115,3 +115,12 @@ def extract_image_patches(X, ksizes, strides, border_mode="valid", dim_ordering= if dim_ordering == "tf": patches = KTH.permute_dimensions(patches, [0, 1, 2, 4, 5, 3]) return patches + + +def moments(x, axes, shift=None, keep_dims=False): + ''' Calculates and returns the mean and variance of the input ''' + + mean_batch = KTH.mean(x, axis=axes, keepdims=keep_dims) + var_batch = KTH.var(x, axis=axes, keepdims=keep_dims) + + return mean_batch, var_batch diff --git a/keras_contrib/layers/normalization.py b/keras_contrib/layers/normalization.py index a3e9d98..179f2e4 100644 --- a/keras_contrib/layers/normalization.py +++ b/keras_contrib/layers/normalization.py @@ -1,3 +1,233 @@ from keras.engine import Layer, InputSpec from .. import initializations, regularizers from .. import backend as K + +import numpy as np + + +class BatchRenormalization(Layer): + """Batch renormalization layer (Sergey Ioffe, 2017). + + Normalize the activations of the previous layer at each batch, + i.e. applies a transformation that maintains the mean activation + close to 0 and the activation standard deviation close to 1. + + # Arguments + epsilon: small float > 0. Fuzz parameter. + Theano expects epsilon >= 1e-5. + mode: integer, 0, 1 or 2. + - 0: feature-wise normalization. + Each feature map in the input will + be normalized separately. The axis on which + to normalize is specified by the `axis` argument. + Note that if the input is a 4D image tensor + using Theano conventions (samples, channels, rows, cols) + then you should set `axis` to `1` to normalize along + the channels axis. + During training and testing we use running averages + computed during the training phase to normalize the data + - 1: sample-wise normalization. This mode assumes a 2D input. + - 2: feature-wise normalization, like mode 0, but + using per-batch statistics to normalize the data during both + testing and training. + axis: integer, axis along which to normalize in mode 0. For instance, + if your input tensor has shape (samples, channels, rows, cols), + set axis to 1 to normalize per feature map (channels axis). + momentum: momentum in the computation of the + exponential average of the mean and standard deviation + of the data, for feature-wise normalization. + r_max_value: Upper limit of the value of r_max. + d_max_value: Upper limit of the value of d_max. + t_delta: At each iteration, increment the value of t by t_delta. + weights: Initialization weights. + List of 2 Numpy arrays, with shapes: + `[(input_shape,), (input_shape,)]` + Note that the order of this list is [gamma, beta, mean, std] + beta_init: name of initialization function for shift parameter + (see [initializations](../initializations.md)), or alternatively, + Theano/TensorFlow function to use for weights initialization. + This parameter is only relevant if you don't pass a `weights` argument. + gamma_init: name of initialization function for scale parameter (see + [initializations](../initializations.md)), or alternatively, + Theano/TensorFlow function to use for weights initialization. + This parameter is only relevant if you don't pass a `weights` argument. + gamma_regularizer: instance of [WeightRegularizer](../regularizers.md) + (eg. L1 or L2 regularization), applied to the gamma vector. + beta_regularizer: instance of [WeightRegularizer](../regularizers.md), + applied to the beta vector. + + # Input shape + Arbitrary. Use the keyword argument `input_shape` + (tuple of integers, does not include the samples axis) + when using this layer as the first layer in a model. + + # Output shape + Same shape as input. + + # References + - [Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift](https://arxiv.org/abs/1502.03167) + """ + + def __init__(self, epsilon=1e-3, mode=0, axis=-1, momentum=0.99, + r_max_val=3., d_max_val=5., t_delta=1., weights=None, beta_init='zero', + gamma_init='one', gamma_regularizer=None, beta_regularizer=None, + **kwargs): + self.supports_masking = True + self.beta_init = initializations.get(beta_init) + self.gamma_init = initializations.get(gamma_init) + self.epsilon = epsilon + self.mode = mode + self.axis = axis + self.momentum = momentum + self.gamma_regularizer = regularizers.get(gamma_regularizer) + self.beta_regularizer = regularizers.get(beta_regularizer) + self.initial_weights = weights + self.r_max_value = r_max_val + self.d_max_value = d_max_val + self.t_delta = t_delta + if self.mode == 0: + self.uses_learning_phase = True + super(BatchRenormalization, self).__init__(**kwargs) + + def build(self, input_shape): + self.input_spec = [InputSpec(shape=input_shape)] + shape = (input_shape[self.axis],) + + self.gamma = self.add_weight(shape, + initializer=self.gamma_init, + regularizer=self.gamma_regularizer, + name='{}_gamma'.format(self.name)) + self.beta = self.add_weight(shape, + initializer=self.beta_init, + regularizer=self.beta_regularizer, + name='{}_beta'.format(self.name)) + self.running_mean = self.add_weight(shape, initializer='zero', + name='{}_running_mean'.format(self.name), + trainable=False) + # Note: running_std actually holds the running variance, not the running std. + self.running_std = self.add_weight(shape, initializer='one', + name='{}_running_std'.format(self.name), + trainable=False) + + self.r_max = K.variable(np.ones((1,)), name='{}_r_max'.format(self.name)) + + self.d_max = K.variable(np.zeros((1,)), name='{}_d_max'.format(self.name)) + + self.t = K.variable(np.zeros((1,)), name='{}_t'.format(self.name)) + + if self.initial_weights is not None: + self.set_weights(self.initial_weights) + del self.initial_weights + self.built = True + + def call(self, x, mask=None): + if self.mode == 0 or self.mode == 2: + assert self.built, 'Layer must be built before being called' + input_shape = K.int_shape(x) + + reduction_axes = list(range(len(input_shape))) + del reduction_axes[self.axis] + broadcast_shape = [1] * len(input_shape) + broadcast_shape[self.axis] = input_shape[self.axis] + + mean_batch, var_batch = K.moments(x, reduction_axes, shift=None, keep_dims=False) + std_batch = (K.sqrt(var_batch + self.epsilon)) + + r_max_val = K.get_value(self.r_max) + r = std_batch / (K.sqrt(self.running_std + self.epsilon)) + r = K.stop_gradient(K.clip(r, 1 / r_max_val, r_max_val)) + + d_max_val = K.get_value(self.d_max) + d = (mean_batch - self.running_mean) / K.sqrt(self.running_std + self.epsilon) + d = K.stop_gradient(K.clip(d, -d_max_val, d_max_val)) + + if sorted(reduction_axes) == range(K.ndim(x))[:-1]: + x_normed_batch = (x - mean_batch) / std_batch + x_normed = (x_normed_batch * r + d) * self.gamma + self.beta + else: + # need broadcasting + broadcast_mean = K.reshape(mean_batch, broadcast_shape) + broadcast_std = K.reshape(std_batch, broadcast_shape) + broadcast_r = K.reshape(r, broadcast_shape) + broadcast_d = K.reshape(d, broadcast_shape) + broadcast_beta = K.reshape(self.beta, broadcast_shape) + broadcast_gamma = K.reshape(self.gamma, broadcast_shape) + + x_normed_batch = (x - broadcast_mean) / broadcast_std + x_normed = (x_normed_batch * broadcast_r + broadcast_d) * broadcast_gamma + broadcast_beta + + # explicit update to moving mean and standard deviation + self.add_update([K.moving_average_update(self.running_mean, mean_batch, self.momentum), + K.moving_average_update(self.running_std, std_batch ** 2, self.momentum)], x) + + # update r_max and d_max + t_val = K.get_value(self.t) + r_val = self.r_max_value / (1 + (self.r_max_value - 1) * np.exp(-t_val)) + d_val = self.d_max_value / (1 + ((self.d_max_value / 1e-3) - 1) * np.exp(-(2 * t_val))) + t_val += float(self.t_delta) + + self.add_update([K.update(self.r_max, r_val), + K.update(self.d_max, d_val), + K.update(self.t, t_val)], x) + + if self.mode == 0: + if sorted(reduction_axes) == range(K.ndim(x))[:-1]: + x_normed_running = K.batch_normalization( + x, self.running_mean, self.running_std, + self.beta, self.gamma, + epsilon=self.epsilon) + else: + # need broadcasting + broadcast_running_mean = K.reshape(self.running_mean, broadcast_shape) + broadcast_running_std = K.reshape(self.running_std, broadcast_shape) + broadcast_beta = K.reshape(self.beta, broadcast_shape) + broadcast_gamma = K.reshape(self.gamma, broadcast_shape) + x_normed_running = K.batch_normalization( + x, broadcast_running_mean, broadcast_running_std, + broadcast_beta, broadcast_gamma, + epsilon=self.epsilon) + + # pick the normalized form of x corresponding to the training phase + # for batch renormalization, inference time remains same as batchnorm + x_normed = K.in_train_phase(x_normed, x_normed_running) + + elif self.mode == 1: + # sample-wise normalization + m = K.mean(x, axis=self.axis, keepdims=True) + std = K.sqrt(K.var(x, axis=self.axis, keepdims=True) + self.epsilon) + x_normed_batch = (x - m) / (std + self.epsilon) + + r_max_val = K.get_value(self.r_max) + r = std / (self.running_std + self.epsilon) + r = K.stop_gradient(K.clip(r, 1 / r_max_val, r_max_val)) + + d_max_val = K.get_value(self.d_max) + d = (m - self.running_mean) / (self.running_std + self.epsilon) + d = K.stop_gradient(K.clip(d, -d_max_val, d_max_val)) + + x_normed = ((x_normed_batch * r) + d) * self.gamma + self.beta + + # update r_max and d_max + t_val = K.get_value(self.t) + r_val = self.r_max_value / (1 + (self.r_max_value - 1) * np.exp(-t_val)) + d_val = self.d_max_value / (1 + ((self.d_max_value / 1e-3) - 1) * np.exp(-(2 * t_val))) + t_val += float(self.t_delta) + + self.add_update([K.update(self.r_max, r_val), + K.update(self.d_max, d_val), + K.update(self.t, t_val)], x) + + return x_normed + + def get_config(self): + config = {'epsilon': self.epsilon, + 'mode': self.mode, + 'axis': self.axis, + 'gamma_regularizer': self.gamma_regularizer.get_config() if self.gamma_regularizer else None, + 'beta_regularizer': self.beta_regularizer.get_config() if self.beta_regularizer else None, + 'momentum': self.momentum, + 'r_max_value': self.r_max_value, + 'd_max_value': self.d_max_value, + 't_delta': self.t_delta} + base_config = super(BatchRenormalization, self).get_config() + return dict(list(base_config.items()) + list(config.items())) diff --git a/tests/keras_contrib/backend/test_backends.py b/tests/keras_contrib/backend/test_backends.py index 51d6f59..15614ad 100644 --- a/tests/keras_contrib/backend/test_backends.py +++ b/tests/keras_contrib/backend/test_backends.py @@ -73,6 +73,7 @@ def check_composed_tensor_operations(first_function_name, first_function_args, class TestBackend(object): + def test_extract(self): for input_shape in [(1, 3, 40, 40), (1, 3, 10, 10)]: for kernel_shape in [2, 5]: @@ -99,6 +100,32 @@ class TestBackend(object): assert zth.shape == ztf.shape assert_allclose(zth, ztf, atol=1e-02) + def test_moments(self): + input_shape = (10, 10, 10, 10) + x_0 = np.zeros(input_shape) + x_1 = np.ones(input_shape) + x_random = np.random.random(input_shape) + + th_axes = [0, 2, 3] + tf_axes = [0, 1, 2] + + for ip in [x_0, x_1, x_random]: + for axes in [th_axes, tf_axes]: + for keep_dims in [True, False]: + ip_th = KTH.variable(ip) + th_mean, th_var = KCTH.moments(ip_th, axes, keep_dims=keep_dims) + + ip_tf = KTF.variable(ip) + tf_mean, tf_var = KCTF.moments(ip_tf, axes, keep_dims=keep_dims) + + th_mean_val = KTH.eval(th_mean) + tf_mean_val = KTF.eval(tf_mean) + th_var_val = KTH.eval(th_var) + tf_var_val = KTF.eval(tf_var) + + assert_allclose(th_mean_val, tf_mean_val, rtol=1e-4) + assert_allclose(th_var_val, tf_var_val, rtol=1e-4) + if __name__ == '__main__': pytest.main([__file__]) diff --git a/tests/keras_contrib/layers/test_normalization.py b/tests/keras_contrib/layers/test_normalization.py index 278929d..9028d3c 100644 --- a/tests/keras_contrib/layers/test_normalization.py +++ b/tests/keras_contrib/layers/test_normalization.py @@ -9,6 +9,117 @@ from keras.models import Sequential, Model from keras import backend as K from keras_contrib import backend as KC +input_1 = np.arange(10) +input_2 = np.zeros(10) +input_3 = np.ones((10)) +input_shapes = [np.ones((10, 10)), np.ones((10, 10, 10))] + + +@keras_test +def basic_batchrenorm_test(): + from keras import regularizers + layer_test(normalization.BatchRenormalization, + kwargs={'mode': 1, + 'gamma_regularizer': regularizers.l2(0.01), + 'beta_regularizer': regularizers.l2(0.01)}, + input_shape=(3, 4, 2)) + layer_test(normalization.BatchRenormalization, + kwargs={'mode': 0}, + input_shape=(3, 4, 2)) + + +@keras_test +def test_batchrenorm_mode_0_or_2(): + for mode in [0, 2]: + model = Sequential() + norm_m0 = normalization.BatchRenormalization(mode=mode, input_shape=(10,), momentum=0.8) + model.add(norm_m0) + model.compile(loss='mse', optimizer='sgd') + + # centered on 5.0, variance 10.0 + X = np.random.normal(loc=5.0, scale=10.0, size=(1000, 10)) + model.fit(X, X, nb_epoch=4, verbose=0) + out = model.predict(X) + out -= K.eval(norm_m0.beta) + out /= K.eval(norm_m0.gamma) + + assert_allclose(out.mean(), 0.0, atol=1e-1) + assert_allclose(out.std(), 1.0, atol=1e-1) + + +@keras_test +def test_batchrenorm_mode_0_or_2_twice(): + # This is a regression test for issue #4881 with the old + # batch normalization functions in the Theano backend. + model = Sequential() + model.add(normalization.BatchRenormalization(mode=0, input_shape=(10, 5, 5), axis=1)) + model.add(normalization.BatchRenormalization(mode=0, input_shape=(10, 5, 5), axis=1)) + model.compile(loss='mse', optimizer='sgd') + + X = np.random.normal(loc=5.0, scale=10.0, size=(20, 10, 5, 5)) + model.fit(X, X, nb_epoch=1, verbose=0) + model.predict(X) + + +@keras_test +def test_batchrenorm_mode_0_convnet(): + model = Sequential() + norm_m0 = normalization.BatchRenormalization(mode=0, axis=1, input_shape=(3, 4, 4), momentum=0.8) + model.add(norm_m0) + model.compile(loss='mse', optimizer='sgd') + + # centered on 5.0, variance 10.0 + X = np.random.normal(loc=5.0, scale=10.0, size=(1000, 3, 4, 4)) + model.fit(X, X, nb_epoch=4, verbose=0) + out = model.predict(X) + out -= np.reshape(K.eval(norm_m0.beta), (1, 3, 1, 1)) + out /= np.reshape(K.eval(norm_m0.gamma), (1, 3, 1, 1)) + + assert_allclose(np.mean(out, axis=(0, 2, 3)), 0.0, atol=1e-1) + assert_allclose(np.std(out, axis=(0, 2, 3)), 1.0, atol=1e-1) + + +@keras_test +def test_batchrenorm_mode_1(): + norm_m1 = normalization.BatchRenormalization(input_shape=(10,), mode=1) + norm_m1.build(input_shape=(None, 10)) + + for inp in [input_1, input_2, input_3]: + out = (norm_m1.call(K.variable(inp)) - norm_m1.beta) / norm_m1.gamma + assert_allclose(K.eval(K.mean(out)), 0.0, atol=1e-1) + if inp.std() > 0.: + assert_allclose(K.eval(K.std(out)), 1.0, atol=1e-1) + else: + assert_allclose(K.eval(K.std(out)), 0.0, atol=1e-1) + + +@keras_test +def test_shared_batchrenorm(): + '''Test that a BN layer can be shared + across different data streams. + ''' + # Test single layer reuse + bn = normalization.BatchRenormalization(input_shape=(10,), mode=0) + x1 = Input(shape=(10,)) + bn(x1) + + x2 = Input(shape=(10,)) + y2 = bn(x2) + + x = np.random.normal(loc=5.0, scale=10.0, size=(2, 10)) + model = Model(x2, y2) + assert len(model.updates) == 5 + model.compile('sgd', 'mse') + model.train_on_batch(x, x) + + # Test model-level reuse + x3 = Input(shape=(10,)) + y3 = model(x3) + new_model = Model(x3, y3) + assert len(model.updates) == 5 + new_model.compile('sgd', 'mse') + new_model.train_on_batch(x, x) + if __name__ == '__main__': pytest.main([__file__])