【问题标题】:Add metadata to tensorflow frozen graph pb将元数据添加到 tensorflow 冻结图 pb
【发布时间】:2019-02-12 03:58:33
【问题描述】:

为了分享我们训练有素的 tensorflow 网络,我们将图表冻结到 .pb 文件中。我们还创建了一个 xml 文件,其中包含一些元数据,例如输入张量和输出张量、要应用的预处理类型、训练数据信息等。然后通过加载图形和评估张量等使用 Java 或 C# 为模型提供服务。

为了使共享更容易,我想将此 xml 数据包含在 .pb 文件中的某处。有没有办法做到这一点?一个想法是将它作为 tf.Constant,但我不知道如何将它连接到普通图。

注意这里使用的是freeze_graph.py。新的 SavedModel 格式是否更合适?

【问题讨论】:

标签: python tensorflow


【解决方案1】:

首先,是的,您应该使用新的 SavedModel 格式,因为它是未来 TF 团队将支持的格式,并且也适用于 Keras。您可以向模型添加一个额外的端点,它返回一个带有 XML 数据字符串的常量张量(正如您所提到的)。

这很好,因为它是封闭的——底层的 savemodel 格式无关紧要,因为您的元数据保存在计算图本身中。

查看此问题的答案:Saving a TF2 keras model with custom signature defs。对于 Keras,这个答案并没有让你 100% 成功,因为它不能与 tf.keras.models.load 函数很好地互操作,因为它们将它包装在 tf.Module 中。幸运的是,如果添加 tf.function 装饰器,在 TF2 中使用 tf.keras.Model 也可以:

class MyModel(tf.keras.Model):

  def __init__(self, metadata, **kwargs):
    super(MyModel, self).__init__(**kwargs)
    self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu)
    self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax)
    self.metadata = tf.constant(metadata)

  def call(self, inputs):
    x = self.dense1(inputs)
    return self.dense2(x)

  @tf.function(input_signature=[])
  def get_metadata(self):
    return self.metadata

model = MyModel('metadata_test')
input_arr = tf.random.uniform((5, 5, 1)) # This call is needed so Keras knows its input shape. You could define manually too
outputs = model(input_arr)

然后您可以按如下方式保存和加载您的模型:

tf.keras.models.save_model(model, 'test_model_keras')
model_loaded = tf.keras.models.load_model('test_model_keras')

最后使用model_loaded.get_metadata() 检索您的常量元数据张量。

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2019-01-04
    • 2018-01-08
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-01-26
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多