From e2fdad61d93a2ec830f68b9cf130b4688a4c8040 Mon Sep 17 00:00:00 2001 From: PatWie Date: Thu, 4 Oct 2018 14:08:40 +0200 Subject: [PATCH] fix layer, when input is not a tensor --- layers.py | 27 ++++++++------------------- 1 file changed, 8 insertions(+), 19 deletions(-) diff --git a/layers.py b/layers.py index da84707..8d5012b 100644 --- a/layers.py +++ b/layers.py @@ -83,11 +83,8 @@ class FlexPooling(Layer): neighborhoods = ops.convert_to_tensor(inputs[1], dtype=tf.int32) if self.data_format == 'expanded': - features = _remove_dim(inputs[0], 2) - neighborhoods = _remove_dim(inputs[1], 2) - else: - features = inputs[0] - neighborhoods = inputs[1] + features = _remove_dim(features, 2) + neighborhoods = _remove_dim(neighborhoods, 2) y, _ = _flex_pooling(features, neighborhoods) @@ -234,13 +231,9 @@ class FlexConvolution(Layer): neighborhoods = ops.convert_to_tensor(inputs[2], dtype=tf.int32) if self.data_format == 'expanded': - features = _remove_dim(inputs[0], 2) - positions = _remove_dim(inputs[1], 2) - neighborhoods = _remove_dim(inputs[2], 2) - else: - features = inputs[0] - positions = inputs[1] - neighborhoods = inputs[2] + features = _remove_dim(features, 2) + positions = _remove_dim(positions, 2) + neighborhoods = _remove_dim(neighborhoods, 2) y = _flex_convolution(features, positions, neighborhoods, self.position_theta, self.position_bias) @@ -339,13 +332,9 @@ class FlexConvolutionTranspose(FlexConvolution): neighborhoods = ops.convert_to_tensor(inputs[2], dtype=tf.int32) if self.data_format == 'expanded': - features = _remove_dim(inputs[0], 2) - positions = _remove_dim(inputs[1], 2) - neighborhoods = _remove_dim(inputs[2], 2) - else: - features = inputs[0] - positions = inputs[1] - neighborhoods = inputs[2] + features = _remove_dim(features, 2) + positions = _remove_dim(positions, 2) + neighborhoods = _remove_dim(neighborhoods, 2) y = _flex_convolution_transpose(features, positions, neighborhoods, self.position_theta, self.position_bias)