You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何通过继承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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.29 18:54:03