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
This commit is contained in:
Somshubra Majumdar
2017-02-22 09:07:13 -08:00
committed by Michael Oliver
parent 4ddb91c1b2
commit bdc64bf2f3
5 changed files with 383 additions and 0 deletions
@@ -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)
+9
View File
@@ -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
+230
View File
@@ -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()))
@@ -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__])
@@ -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__])