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
额外注意事项
- 确保模型已构建输入形状:在调用
model.save()前,可通过以下方式确保模型知晓输入形状:- 调用一次
model.fit()完成至少一个batch的训练 - 手动调用
model.build(input_shape),例如:# 假设输入图像尺寸为(224,224,3) model.build([(None, 224, 224, 3), (None, 224, 224, 3)])
- 调用一次
- 训练时使用
self(inputs)而非直接调用子模块:自定义train_step中必须通过self([x1,x2], training=True)触发前向传播,而非直接调用self.encoder或self.projection_head,这样Keras才能跟踪模型的完整前向路径。
内容的提问来源于stack exchange,提问作者HuckleberryFinn
相关产品推荐
相关产品推荐

