如何通过继承tf.keras.Model类封装TensorFlow模型代码并解决运行报错
问题原因
你的报错是因为混淆了Keras Functional API和子类化API的使用规则,核心错误有3处:
- 所有可训练层、预训练模型实例必须放在
__init__方法中初始化,不能放在call方法里每次前向传播重复创建,会导致计算图构建异常、参数无法复用。 - 不能在
call方法中定义Input层,子类化模型的输入就是call方法的inputs参数,额外定义的Input会生成无关联的符号张量,和实际运行传入的张量冲突,触发类型转换错误。 - 调用ResNet模型时不需要取
.output属性,直接把输入张量传给ResNet实例即可得到输出,.output是Functional API构建静态计算图的专属用法,不适用于子类化的动态调用逻辑。
正确实现代码
import tensorflow as tf from tensorflow.keras.layers import Dense, GlobalAveragePooling2D from tensorflow.keras.applications import ResNet50V2 class StereoModel(tf.keras.Model): def __init__(self): super(StereoModel, self).__init__() # 所有层/模型统一在初始化阶段创建 self.base_model = ResNet50V2( input_shape=(720, 540, 2), include_top=False, weights=None, classes=4 ) self.global_avg_pool = GlobalAveragePooling2D() self.output_dense = Dense(4, activation='sigmoid') def call(self, inputs, training=None): # 直接用方法入参作为输入,依次调用各层 x = self.base_model(inputs, training=training) x = self.global_avg_pool(x) return self.output_dense(x)
使用说明
实例化后建议先调用build方法指定输入形状,避免延迟初始化导致的问题:
model = StereoModel() # 指定批次维度为None,支持动态批次大小 model.build(input_shape=(None, 720, 540, 2)) # 打印模型结构 model.summary()
注意:
call方法的training参数需要传递给ResNet实例,保证BatchNormalization、Dropout等层在训练和推理阶段的行为符合预期。
内容的提问来源于stack exchange,提问作者coderide
相关产品推荐
相关产品推荐

