TensorFlow多输入输出模型训练需递增input编号问题咨询
问题原因
该问题是Keras输入层的默认命名规则导致的:你没有给模型的输入层设置显式名称时,Keras会按创建顺序给输入层自动分配input_序号格式的默认名称,每在同一个Python进程中创建一次输入层,序号就会自动累加。因为你的模型有2个输入,所以每次重建模型后,输入层名称的序号都会递增2,和你之前写死在数据集里的input_17、input_18无法匹配,才会要求你修改编号。
可行解决方案
方案1:给输入层设置固定名称(最推荐)
在定义模型的输入层时,通过name参数指定固定的自定义名称,后续数据集字典的key和该名称保持一致即可,无论重建多少次模型都不会出现不匹配问题:
# 示例:定义输入层时添加name参数 input_time = tf.keras.Input(shape=(你的输入形状,), name="input_time") input_a = tf.keras.Input(shape=(你的输入形状,), name="input_a") # 构造数据集时直接用你定义的固定名称做key train_dataset = tf.data.Dataset.from_tensor_slices( ( {"input_time": timetr, "input_a": atr}, {"ed": wtr, "sd": wbtr}, ) )
方案2:动态获取模型输入名称构造数据集
如果不想修改模型定义,可以在构造数据集前,直接从已构建的模型对象中获取输入层的实际名称,自动适配编号:
# 构建完模型后获取输入层名称列表,顺序和你定义输入层的顺序一致 input_names = model.input_names # 直接用动态获取的名称构造字典 train_dataset = tf.data.Dataset.from_tensor_slices( ( {input_names[0]: timetr, input_names[1]: atr}, {"ed": wtr, "sd": wbtr}, ) ) train_dataset = train_dataset.batch(10)
方案3:重启运行环境
如果是在Jupyter等交互环境下运行代码,每次运行全量训练代码前重启内核,进程重置后Keras的输入层序号会从1开始重新计数,也能避免编号递增的问题,该方案适合临时测试场景,不推荐长期使用。
内容的提问来源于stack exchange,提问作者user16931025
相关产品推荐
相关产品推荐

