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

TensorFlow Keras子类化模型保存失败:如何重写Model.call()?

解决NNCLR模型无法保存的问题

问题根源

原NNCLR示例的自定义keras.Model类未重写call()方法,导致Keras无法识别模型的前向传播路径和输入输出形状,进而无法完成序列化保存。

解决方案:重写call()方法

修改NNCLR类,添加call()方法,明确模型在训练和推理场景下的前向逻辑,同时保证Keras能正确捕获输入输出结构。

修改后的NNCLR类示例

class NNCLR(keras.Model):
    def __init__(self, encoder, projection_head, temperature=0.1):
        super().__init__()
        self.encoder = encoder
        self.projection_head = projection_head
        self.temperature = temperature
        self.loss_tracker = keras.metrics.Mean(name="loss")

    @property
    def metrics(self):
        return [self.loss_tracker]

    # 重写call方法,明确前向传播逻辑
    def call(self, inputs, training=False):
        # 训练阶段:输入为两个增强视图的列表
        if training and isinstance(inputs, list) and len(inputs) == 2:
            x1, x2 = inputs
            # 编码并生成投影特征
            z1 = self.projection_head(self.encoder(x1, training=training))
            z2 = self.projection_head(self.encoder(x2, training=training))
            return z1, z2
        # 推理阶段:输入为单张图像,返回编码器的特征(可根据需求改为投影头输出)
        else:
            return self.encoder(inputs, training=training)

    def train_step(self, data):
        x1, x2 = data
        # 改用model.__call__方式调用前向传播(即self(inputs)),而非直接调用子模块
        z1, z2 = self([x1, x2], training=True)

        # 保留原训练逻辑的loss计算、梯度更新部分
        with tf.GradientTape() as tape:
            loss = self.nnclr_loss(z1, z2)

        # 计算梯度并更新权重
        trainable_vars = self.trainable_variables
        gradients = tape.gradient(loss, trainable_vars)
        self.optimizer.apply_gradients(zip(gradients, trainable_vars))

        # 更新损失指标
        self.loss_tracker.update_state(loss)
        return {"loss": self.loss_tracker.result()}

    def nnclr_loss(self, z1, z2):
        # 保留原损失计算逻辑
        z1 = tf.math.l2_normalize(z1, axis=1)
        z2 = tf.math.l2_normalize(z2, axis=1)
        cross_view_loss = -tf.reduce_mean(tf.matmul(z1, z2, transpose_b=True) / self.temperature)
        return cross_view_loss

额外注意事项

  1. 确保模型已构建输入形状:在调用model.save()前,可通过以下方式确保模型知晓输入形状:
    • 调用一次model.fit()完成至少一个batch的训练
    • 手动调用model.build(input_shape),例如:
      # 假设输入图像尺寸为(224,224,3)
      model.build([(None, 224, 224, 3), (None, 224, 224, 3)])
      
  2. 训练时使用self(inputs)而非直接调用子模块:自定义train_step中必须通过self([x1,x2], training=True)触发前向传播,而非直接调用self.encoder或self.projection_head,这样Keras才能跟踪模型的完整前向路径。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 07:31:07