Skip to content

Commit a407dd5

Browse files
siju-samueldhruvaray
authored andcommitted
[KERAS]Minimum & AlphaDropout op support (apache#5380)
1 parent f448dac commit a407dd5

File tree

2 files changed

+8
-2
lines changed

2 files changed

+8
-2
lines changed

python/tvm/relay/frontend/keras.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -186,8 +186,11 @@ def _convert_merge(inexpr, keras_layer, _):
186186
elif merge_type == 'Subtract':
187187
assert len(inexpr) == 2, "Subtract merge takes 2 inputs."
188188
ret = _op.subtract(ret, inexpr[1])
189-
elif merge_type in ['Add', 'Multiply', 'Maximum']:
190-
op_map = {'Add': _op.add, 'Multiply': _op.multiply, 'Maximum': _op.maximum}
189+
elif merge_type in ['Add', 'Multiply', 'Minimum', 'Maximum']:
190+
op_map = {'Add': _op.add,
191+
'Multiply': _op.multiply,
192+
'Minimum': _op.minimum,
193+
'Maximum': _op.maximum}
191194
for i in range(1, len(inexpr)):
192195
ret = op_map[merge_type](ret, inexpr[i])
193196
elif merge_type == 'Average':
@@ -902,6 +905,7 @@ def _default_skip(inexpr, keras_layer, _): # pylint: disable=unused-argument
902905
# 'TimeDistributed' : _default_skip,
903906

904907
'Average' : _convert_merge,
908+
'Minimum' : _convert_merge,
905909
'Maximum' : _convert_merge,
906910
'Dot' : _convert_merge,
907911
'Permute' : _convert_permute,
@@ -910,6 +914,7 @@ def _default_skip(inexpr, keras_layer, _): # pylint: disable=unused-argument
910914

911915
'InputLayer' : _default_skip,
912916
'Dropout' : _default_skip,
917+
'AlphaDropout' : _default_skip,
913918
'SpatialDropout2D' : _default_skip,
914919
'SpatialDropout1D' : _default_skip,
915920
'GaussianDropout' : _default_skip,

tests/python/frontend/keras/test_forward.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,7 @@ def test_forward_merge(self, keras):
125125
keras.layers.Subtract(),
126126
keras.layers.Multiply(),
127127
keras.layers.Maximum(),
128+
keras.layers.Minimum(),
128129
keras.layers.Average(),
129130
keras.layers.Concatenate()]
130131
for merge_func in merge_funcs:

0 commit comments

Comments
 (0)