如何将TensorFlow 1.0的tf.placeholder转换为TensorFlow 2.0的tf.keras.Input
TensorFlow 1.x占位符转TensorFlow 2.x keras输入的实现方法
直接等价转换代码
你可以直接用以下代码替换原有tf.placeholder写法:
tf.keras.Input(shape=(n, p), dtype=tf.float32)
写法说明
tf.keras.Input的shape参数不需要显式声明batch维度,原有TF1写法里的第一维None代表可变batch size,这个是tf.keras.Input的默认支持特性,无需额外配置dtype参数用法和原tf.placeholder完全一致,如果你的场景默认使用float32数据类型,可以省略该参数,简化为:
tf.keras.Input(shape=(n, p))
- 原有TF1中需要用
feed_dict给占位符传值的逻辑,在TF2 keras体系下直接在model.fit、model.predict等方法中传入对应数据即可,不需要单独处理传参逻辑。
内容的提问来源于stack exchange,提问作者lindo
相关产品推荐
相关产品推荐

