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

Keras自定义SiameseModel调用save()报前向传播未定义错误

问题描述

模型训练全流程可正常执行完成,调用model.save API保存模型时抛出异常,相关实现代码如下:

class SiameseModel(Model):
    """The Siamese Network model with a custom training and testing loops.
    Computes the triplet loss using the three embeddings produced by the
    Siamese Network.
    The triplet loss is defined as:
       L(A, P, N) = max(‖f(A) - f(P)‖² - ‖f(A) - f(N)‖² + margin, 0)
    """

    def __init__(self, siamese_network, margin=0.5):
        super(SiameseModel, self).__init__()
        self.siamese_network = siamese_network
        self.margin = margin
        self.loss_tracker = metrics.Mean(name="loss")

    def call(self, inputs):
        return self.siamese_network(inputs)

    def train_step(self, data):
        # GradientTape is a context manager that records every operation that
        # you do inside. We are using it here to compute the loss so we can get
        # the gradients and apply them using the optimizer specified in
        # `compile()`.
        with tf.GradientTape() as tape:
            loss = self._compute_loss(data)

        # Storing the gradients of the loss function with respect to the
        # weights/parameters.
        gradients = tape.gradient(loss, self.siamese_network.trainable_weights)

        # Applying the gradients on the model using the specified optimizer
        self.optimizer.apply_gradients(
            zip(gradients, self.siamese_network.trainable_weights)
        )

        # Let's update and return the training loss metric.
        self.loss_tracker.update_state(loss)
        return {"loss": self.loss_tracker.result()}

    def test_step(self, data):
        loss = self._compute_loss(data)

        # Let's update and return the loss metric.
        self.loss_tracker.update_state(loss)
        return {"loss": self.loss_tracker.result()}

    def _compute_loss(self, data):
        # The output of the network is a tuple containing the distances
        # between the anchor and the positive example, and the anchor and
        # the negative example.
        ap_distance, an_distance = self.siamese_network(data)

        # Computing the Triplet Loss by subtracting both distances and
        # making sure we don't get a negative value.
        loss = ap_distance - an_distance
        loss = tf.maximum(loss + self.margin, 0.0)
        return loss

    @property
    def metrics(self):
        # We need to list our metrics here so the `reset_states()` can be
        # called automatically.
        return [self.loss_tracker]


"""
## Training
We are now ready to train our model.
"""

siamese_model = SiameseModel(siamese_network)
siamese_model.compile(optimizer=optimizers.Adam(0.0001))
siamese_model.fit(train_dataset, epochs=10, validation_data=val_dataset)

siamese_model.save("./siamese_model.pt")
报错信息
ValueError: Model <siamese_model.SiameseModel object at 0x17e0c7b50> cannot be saved either because the input shape is not available or because the forward pass of the model is not defined.To define a forward pass, please override `Model.call()`. To specify an input shape, either call `build(input_shape)` directly, or call the model on actual data using `Model()`, `Model.fit()`, or `Model.predict()`. If you have a custom training step, please make sure to invoke the forward pass in train step through `Model.__call__`, i.e. `model(inputs)`, as opposed to `model.call()`.
故障原因

自定义的SiameseModel虽然重写了call方法,但在自定义train_step、test_step的逻辑中,前向传播直接调用了内部嵌套的self.siamese_network子网络,从未通过外层SiameseModel自身的调用入口(即self(inputs),对应Model.__call__方法)执行前向计算。训练过程中Keras无法自动追踪外层模型的拓扑结构、输入形状信息,保存时判定模型没有合法的前向传播定义,因此抛出异常。

解决方法

任意选择以下一种方案即可修复问题:

  • 方案1:保存前取一批真实输入跑一次前向,触发模型自动构建
    # 从训练集取一个批次的样本
    sample_data = next(iter(train_dataset))
    # 调用外层模型执行一次前向,完成结构构建
    _ = siamese_model(sample_data)
    # 再执行保存即可正常运行
    siamese_model.save("./siamese_model.pt")
    
  • 方案2:修改自定义步中的前向逻辑,通过外层模型入口调用
    将_compute_loss方法中的ap_distance, an_distance = self.siamese_network(data)替换为ap_distance, an_distance = self(data),让训练过程中Keras可以正常追踪模型前向拓扑,训练完成后可直接保存。
  • 方案3:手动调用build方法指定输入形状完成构建
    传入和实际输入匹配的形状(batch维度填None即可),手动完成模型初始化:
    # 注意将输入形状替换为你实际使用的输入维度,以下为三元组输入224*224三通道图片的示例
    siamese_model.build(input_shape=(None, 224, 224, 3))
    siamese_model.save("./siamese_model.pt")
    

内容的提问来源于stack exchange,提问作者Exploring

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 18:45:32