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
相关产品推荐
相关产品推荐

