From ececa1b7afe09e3566ba2f19004aa24a655dd52f Mon Sep 17 00:00:00 2001 From: lameeus Date: Wed, 25 Oct 2017 13:45:52 +0200 Subject: [PATCH] Made indepent of backend. Now only checks channel priority (as suggested by @arodiss) --- keras_contrib/callbacks/dead_relu_detector.py | 6 ++---- tests/keras_contrib/callbacks/dead_relu_detector_test.py | 6 ++---- 2 files changed, 4 insertions(+), 8 deletions(-) diff --git a/keras_contrib/callbacks/dead_relu_detector.py b/keras_contrib/callbacks/dead_relu_detector.py index 868fc16..f545ee2 100644 --- a/keras_contrib/callbacks/dead_relu_detector.py +++ b/keras_contrib/callbacks/dead_relu_detector.py @@ -63,14 +63,12 @@ class DeadReluDetector(Callback): act_len = len(shape_act) # should work for both Conv and Flat - if K.backend() == 'tensorflow': + if K.image_data_format() == 'channels_last': # features in last axis axis_filter = -1 - elif K.backend() == 'theano': + else: # features before the convolution axis, for weight_len the input and output have to be subtracted axis_filter = -1 - (weight_len - 2) - else: - raise ValueError('Unknown backend: {}'.format(K.backend())) total_featuremaps = shape_act[axis_filter] diff --git a/tests/keras_contrib/callbacks/dead_relu_detector_test.py b/tests/keras_contrib/callbacks/dead_relu_detector_test.py index 89c6bbe..552dbd0 100644 --- a/tests/keras_contrib/callbacks/dead_relu_detector_test.py +++ b/tests/keras_contrib/callbacks/dead_relu_detector_test.py @@ -96,12 +96,10 @@ def test_DeadDeadReluDetector_conv(): n_samples = 9 # (5, 5) kernel, 4 input featuremaps and 10 output featuremaps - if K.backend() == 'tensorflow': + if K.image_data_format() == 'channels_last': input_shape = (n_samples, 5, 5, 4) - elif K.backend() == 'theano': - input_shape = (n_samples, 4, 5, 5) else: - raise ValueError('Unknown backend: {}'.format(K.backend())) + input_shape = (n_samples, 4, 5, 5) # ignore batch size input_shape_conv = tuple(input_shape[1:])