tf.contrib.rnn.LayerNormBasicLSTMCell废弃后TensorFlow 2.1实现RNN的替代方案咨询
完全不需要降级到TensorFlow 1.x版本,TensorFlow 2.1配套的Keras生态有两种成熟的替代方案,可以实现和tf.contrib.rnn.LayerNormBasicLSTMCell完全一致的功能:
方案1:使用官方扩展库预实现接口
TensorFlow官方维护的扩展库专门承接原tf.contrib中被移出的成熟功能,其中的LayerNormLSTMCell和原接口功能100%对齐,支持层归一化、输入dropout、循环dropout等原接口的所有参数。
适配TensorFlow 2.1的对应扩展库版本为0.9.1,直接安装对应版本即可使用,调用示例如下:
import tensorflow as tf import tensorflow_addons as tfa # 定义层归一化LSTM单元,参数和原tf.contrib实现完全对应 cell = tfa.rnn.LayerNormLSTMCell( units=128, activation='tanh', recurrent_activation='sigmoid', dropout=0.1, recurrent_dropout=0.1 ) # 封装为RNN层即可正常接入Keras网络 rnn_layer = tf.keras.layers.RNN(cell, return_sequences=True)
该方案成本最低,接口调用逻辑和原实现几乎没有差异,仅需要调整导入路径即可。
方案2:自定义Cell实现(无额外依赖)
如果不想引入第三方依赖,也可以基于原生Keras的LSTMCell手动叠加层归一化,实现逻辑和原接口完全等价:
import tensorflow as tf class CustomLayerNormLSTMCell(tf.keras.layers.LSTMCell): def __init__(self, units, **kwargs): super().__init__(units, **kwargs) self.layer_norm = tf.keras.layers.LayerNormalization() def call(self, inputs, states, training=None): output, new_states = super().call(inputs, states, training=training) # 在输出前添加层归一化操作 output = self.layer_norm(output) return output, new_states # 使用方式和原生LSTMCell完全一致 cell = CustomLayerNormLSTMCell(units=128, dropout=0.1) rnn_layer = tf.keras.layers.RNN(cell, return_sequences=True)
如果需要对循环状态也做层归一化,仅需要在call方法中对new_states的返回值也添加归一化操作即可,灵活性更高。
两种方案都完全兼容TensorFlow 2.x的eager执行、tf.function编译、分布式训练等所有特性,不存在适配问题。
内容的提问来源于stack exchange,提问作者RealRageDontQuit
相关产品推荐
相关产品推荐

