基于DenseNet201的Keras多头模型训练报错:输入名称不匹配求助
问题解决:DenseNet201三头模型输入名称不匹配报错
问题原因
DenseNet201在Keras中默认生成的输入层名称是类似input_5的自动命名,而你训练时传入的输入字典key是img_input,两者不匹配导致报错。AlexNet版本能运行是因为代码里手动指定了输入层名称为img_input,但DenseNet部分没做这个处理。
解决方案
修改DenseNet201的模型初始化代码,手动定义名为img_input的输入层,替换默认的输入层:
修改后的DenseNet模型代码片段
class base_model(): def __init__(self, side_dim, n_bb, n_classes, name_model): self.side_dim = side_dim self.name_model = name_model # base model DenseNet if name_model == 'DenseNet201': # 手动创建命名为img_input的输入层 self.image_input = keras.layers.Input(shape=(self.side_dim, self.side_dim, 3), name='img_input') # 将DenseNet201作为层,接收自定义的输入层 self.base_model = keras.applications.DenseNet201( include_top=False, input_tensor=self.image_input, # 使用自定义输入张量 ) self.flatten = keras.layers.Flatten()(self.base_model.layers[-2].output) self.BatcNorm = keras.layers.BatchNormalization()(self.flatten) print('Base model: DenseNet201 (7.2M params x 201 layers') # 后续的三头输出部分保持不变...
原理说明
- 通过
keras.layers.Input手动创建输入层并指定name='img_input',让模型的输入名称和训练时传入的字典key一致。 - 使用
input_tensor参数将DenseNet201的基础模型挂载到这个自定义输入层上,替代默认生成的输入层。
验证方法
修改后可以打印模型的输入名称确认:
print(myModel.input_names) # 输出应该是['img_input']
这样再运行原来的训练代码,输入名称匹配就不会再报错了。
内容的提问来源于stack exchange,提问作者Saeed
相关产品推荐
相关产品推荐

