【问题标题】:How to save a non serializable model in Tensorflow如何在 Tensorflow 中保存不可序列化的模型
【发布时间】:2021-11-21 12:55:19
【问题描述】:

我创建了一个带有自定义层的模型,其中包含矩阵运算和类似的东西。我现在想在训练后保存我的模型。我试过了:

model.save("model.h5", save_format='tf')

但是出现了错误:

NotImplementedError: Saving the model to HDF5 format requires the model to be 
a Functional model or a Sequential model. It does not work for subclassed models, 
because such models are defined via the body of a Python method, which isn't safely serializable.
Consider saving to the Tensorflow SavedModel format (by setting save_format="tf") or using `save_weights`.

我发现了一些有用的东西:

checkpoint_path = "checkpoints"

ckpt = tf.train.Checkpoint(model=model,
                           optimizer=optimizer)

ckpt_manager = tf.train.CheckpointManager(ckpt, checkpoint_path, max_to_keep=5)

# if a checkpoint exists, restore the latest checkpoint.
if ckpt_manager.latest_checkpoint:
  ckpt.restore(ckpt_manager.latest_checkpoint)

我的问题是:通过这种方式,我可以做与保存可序列化模型(如序列模型)相同的操作,还是将此检查点用于其他目的?

【问题讨论】:

    标签: python tensorflow tensorflow2.0


    【解决方案1】:

    实际上,您可以使用两种格式来保存模型。您可以简单地使用旧的 Keras H5 格式 model.save("test", save_format='h5') 保存模型,或者通过显式设置 model.save("test", save_format='tf') 或简单地使用 model.save("test") 来使用 Tensorflow SavedModel format,因为在调用 @987654327 时默认使用 tf 格式@。使用model.save("model.h5", save_format='tf'),您似乎正在尝试同时使用这两种格式,但这看起来并不奏效。使用tf 格式保存模型应该可以。更多信息可以找到here。例如以下模型,只有在使用model.savemodel.save("test", save_format='tf')时才能保存:

    import tensorflow as tf
    
    class SomeModel(tf.keras.Model):
    
      def __init__(self):
        super(SomeModel, self).__init__()
        self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu, )
        self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax)
    
      def call(self, inputs):
        x = self.dense1(inputs)
        return self.dense2(x)
    
    model = SomeModel()
    model.compute_output_shape(input_shape=(1,1))
    model.save("model")
    

    在这个子类模型上调用model.save("test", save_format='h5')model.save("test.h5") 甚至model.save("model.h5", save_format='tf') 将导致错误。

    Checkpoints 在您需要中断训练或崩溃并且您想从保存状态恢复训练模型时特别有用。在推理过程中,您可以轻松加载模型的最新检查点并进行预测,而无需重新编译。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2023-02-09
      • 2020-02-26
      • 1970-01-01
      • 1970-01-01
      • 2021-05-06
      • 2023-03-06
      • 1970-01-01
      相关资源
      最近更新 更多