tf.keras调用model.summary()首层输出为multiple问题排查
问题说明
使用TensorFlow tf.keras构建VGG16为骨干的图像分类模型时,调用model.summary()打印结构发现首层输出形状显示为multiple,初步判断和首层关联多输入有关,但无法定位代码问题,需要定位错误并给出可行实现。
核心错误原因
问题完全出在自定义的嵌套模型展平逻辑,以及VGG基模型的重复调用操作:
- 实例化VGG16基模型时,模型内部会自动创建专属输入层
input_1。代码中feature_batch = vgg(image_batch)这行直接调用VGG模型推理真实图像batch,已经让input_1绑定了第一个输入节点;后续又手动新建inp = layers.Input((224,224,3)),通过cnn = vgg(inp)把VGG接到这个新的Input层上,等于VGG自带的input_1同时绑定了两个不同的输入来源。 - 自定义的
flatten_model函数逻辑存在本质缺陷:它只是简单把嵌套模型里的所有层抽出来拼成列表,再塞给Sequential模型,完全没有重建层与层之间的连接拓扑,还把已经绑定了两个输入的VGG内部input_1直接作为Sequential的第一层。这个Input层没有被Sequential正确绑定唯一的输入规格,最终summary就会显示输出形状为multiple。 - 代码还有一处逻辑bug:256维的中间Dense层错误使用了
softmax激活函数,softmax会将输出压缩为和为1的概率分布,放在中间层会严重损失特征表达能力,应该替换为relu激活。
可行修正方案
方案1:直接使用函数式API构建的完整模型(最推荐)
你之前通过函数式API写的model = Model(inp, out)本身就是拓扑完整、可直接训练的模型,完全不需要额外写展平嵌套模型的逻辑,直接删除flatten_model相关代码,使用原模型即可,调用model.summary()就不会出现multiple的显示问题。
核心修正后的代码片段如下:
# 删掉所有flatten_model相关代码,直接使用函数式API构建的模型 vgg = tf.keras.applications.VGG16(input_shape=IMG_SHAPE, weights='imagenet', include_top=False, pooling='max') # 删除无意义的提前推理代码,避免提前绑定输入 # image_batch, label_batch = next(iter(x_train)) # feature_batch = vgg(image_batch) # print(feature_batch.shape) for layer in vgg.layers: layer.trainable = False inp = layers.Input((224,224,3)) # 接入VGG对应的预处理层,无需手动额外处理数据 x = tf.keras.applications.vgg16.preprocess_input(inp) cnn = vgg(x) x = layers.BatchNormalization()(cnn) x = layers.Dropout(0.2)(x) # 修正中间层激活函数为relu x = layers.Dense(256, activation='relu')(x) x = layers.BatchNormalization()(x) x = layers.Dropout(0.2)(x) out = layers.Dense(291, activation='softmax')(x) model = Model(inp, out) # 直接打印原模型的summary,不会出现multiple问题 model.summary()
方案2:正确构建无嵌套的Sequential模型
如果确实需要使用纯Sequential结构、不存在嵌套子模型,不要直接抽离子模型的层列表,应该按顺序逐层添加,同时丢弃VGG内部自带的Input层,让Sequential自动创建绑定唯一规格的输入层:
model_flat = Sequential() # 先添加输入层,明确输入规格 model_flat.add(layers.Input((224,224,3))) model_flat.add(tf.keras.applications.vgg16.preprocess_input) # 逐层添加VGG的卷积、池化层,跳过VGG自带的Input层 for layer in vgg.layers[1:]: model_flat.add(layer) # 依次添加后续自定义层 model_flat.add(layers.BatchNormalization()) model_flat.add(layers.Dropout(0.2)) model_flat.add(layers.Dense(256, activation='relu')) model_flat.add(layers.BatchNormalization()) model_flat.add(layers.Dropout(0.2)) model_flat.add(layers.Dense(291, activation='softmax')) model_flat.summary()
修正后首层输出形状会正常显示为(None, 224, 224, 3),不会再出现multiple标识。
内容的提问来源于stack exchange,提问作者JackDH
相关产品推荐
相关产品推荐

