如何创建自定义tf.keras.Model类 调用summary()提示未构建如何解决
Keras自定义动态迁移学习模型报错解决方案
你的动态组装模型的思路是可行的,但现有写法存在核心逻辑错误,修正后可正常使用。
核心错误原因
- 混淆了Keras层实例和张量的用法
你在__init__中定义的self.x从始至终都是计算图中的张量(tf.keras.Input生成的是输入张量,后续赋值也是张量经过层计算后的输出张量),不是可调用的层实例。call方法需要接收输入张量,传入你定义的层得到输出,而非拿预先生成的静态张量调用输入。 - 没有正确追踪模型层参数
你没有把用到的层(比如ResNet50基础模型)保存为类的实例属性,TensorFlow无法正确追踪这些层的参数,也无法识别模型的拓扑结构,因此调用summary()时会抛出未构建的报错。
修正方案
方案1:保持tf.keras.Model子类写法
import tensorflow as tf from tensorflow import keras class ModelMaker(tf.keras.Model): def __init__(self, img_height, img_width, trained='None'): super().__init__() self.trained = trained # 仅将用到的层定义为实例属性,不提前计算张量 self.preprocess = None self.base_model = None if trained == 'ResNet50': self.preprocess = tf.keras.applications.resnet50.preprocess_input IMG_SHAPE = (img_height, img_width, 3) base_model = tf.keras.applications.ResNet50(input_shape=IMG_SHAPE, include_top=False, weights='imagenet') # 仅批量归一化层设为可训练 for layer in base_model.layers: layer.trainable = isinstance(layer, keras.layers.BatchNormalization) self.base_model = base_model def call(self, inputs): x = inputs if self.trained == 'ResNet50': x = self.preprocess(x) x = self.base_model(x) return x
调用方法:
# 实例化后先指定输入形状完成构建 model = ModelMaker(224, 224, trained='ResNet50') model.build(input_shape=(None, 224, 224, 3)) # 可正常查看结构 model.summary()
方案2:更适配动态组装需求的Functional API写法
你要实现的动态拼接层的需求,用Keras函数式API实现更简单,无需处理子类的构建逻辑,生成的模型可直接查看结构:
def build_model(img_height, img_width, trained='None'): inputs = tf.keras.Input(shape=(img_height, img_width, 3), name="input_layer") x = inputs if trained == 'ResNet50': x = tf.keras.applications.resnet50.preprocess_input(x) base_model = tf.keras.applications.ResNet50(input_shape=(img_height, img_width, 3), include_top=False, weights='imagenet') for layer in base_model.layers: layer.trainable = isinstance(layer, keras.layers.BatchNormalization) x = base_model(x) # 后续扩展可直接在这里添加全局池化、分类全连接层等逻辑 # x = tf.keras.layers.GlobalAveragePooling2D()(x) # outputs = tf.keras.layers.Dense(分类数, activation='softmax')(x) model = tf.keras.Model(inputs=inputs, outputs=x) return model # 直接生成模型即可查看结构 model = build_model(224,224,trained='ResNet50') model.summary()
内容的提问来源于stack exchange,提问作者er3016
相关产品推荐
相关产品推荐

