TensorFlow2.6多输入模型添加validation_data时报错如何解决?
问题解决方案
你遇到的报错是TensorFlow 2.6版本的已知兼容性问题,该版本对validation_data传入多输入元组的解析逻辑存在bug,可按照以下方案依次排查调整:
- 方案1:拆分验证数据为单独参数传入,不使用
validation_data元组写法,适配TF2.6版本的参数解析逻辑
调整后的代码示例:model.fit( x=[train_data1, train_data2], y=train_target, validation_x=[val_data1, val_data2], validation_y=val_target, # 其余你原本设置的epochs、batch_size等参数保持不变 ) - 方案2:如果使用
tf.data.Dataset封装数据,调整验证集的构造格式,确保是(输入组, 标签)的二元结构
构造代码示例:# 验证集构造 val_input_ds = tf.data.Dataset.from_tensor_slices((val_data1, val_data2)) val_label_ds = tf.data.Dataset.from_tensor_slices(val_target) val_dataset = tf.data.Dataset.zip((val_input_ds, val_label_ds)).batch(你的批次大小) # 训练写法 model.fit( x=[train_data1, train_data2], y=train_target, validation_data=val_dataset, # 其余参数保持不变 ) - 方案3:升级TensorFlow版本至2.7及以上,该版本已修复多输入场景下
validation_data的解析bug,你原本的写法可直接正常运行。
内容的提问来源于stack exchange,提问作者Ottpocket
相关产品推荐
相关产品推荐

