如何在Keras中无需自定义类实现两个Dense层间的循环连接?
无需自定义类实现Keras循环连接模型
你要的循环连接结构完全可以用Keras的Functional API实现,不用写自定义类。核心就是把输出层的结果和原始输入拼接后,作为hidden1层的输入,形成闭环。
完整代码
import tensorflow as tf from tensorflow import keras # 定义输入层 inputs = keras.layers.Input(shape=(1,)) # 提前定义所有用到的层实例 concat_layer = keras.layers.Concatenate() hidden1 = keras.layers.Dense(10, name='hidden1') hidden2 = keras.layers.Dense(10, name='hidden2') outputs_layer = keras.layers.Dense(1, name='outputs') # 初始化反馈信号(第一次迭代用零向量替代输出) current_feedback = keras.layers.Lambda(lambda x: tf.zeros_like(x))(inputs) # 构建循环连接路径 concat_result = concat_layer([inputs, current_feedback]) hidden_out1 = hidden1(concat_result) hidden_out2 = hidden2(hidden_out1) final_output = outputs_layer(hidden_out2) # 如果需要多次循环迭代,直接用for循环重复上述步骤即可 # 比如循环3次: # for _ in range(3): # concat_result = concat_layer([inputs, final_output]) # hidden_out1 = hidden1(concat_result) # hidden_out2 = hidden2(hidden_out1) # final_output = outputs_layer(hidden_out2) # 定义模型 model = keras.Model(inputs=inputs, outputs=final_output)
说明
- 第一次迭代时用零向量作为初始反馈,之后每一轮的输出都会作为下一轮的反馈信号,和原始输入拼接后传入hidden1层,完美实现你要的循环结构。
- 所有操作都用Keras内置层完成,完全不需要自定义类。
内容的提问来源于stack exchange,提问作者Thomas Wagenaar
相关产品推荐
相关产品推荐

