【问题标题】:Restoring saved model with dropout applied from java program使用从 java 程序应用的 dropout 恢复保存的模型
【发布时间】:2020-04-02 07:45:41
【问题描述】:

我有一个应用了 dropout 的预训练模型,我想从 java 程序中恢复它。

对于我的应用程序,在推理步骤中,我需要打开 dropout 并多次重复向模型输入输入并获得一系列预测。

我做了什么:

  1. 加载模型并初始化会话
model = SavedModelBundle.load ("path_to_model", "serve");
sess  = model.session();
  1. 向模型馈送(重复多次,例如 3 次)
for (i = 0; i < 3; i++) 
    t_pred = sess.runner().feed("x", x).fetch("y").run().get(0);

假设:

  • 第一次:获取数组 A1 = [y1, y2, y3]
  • 第二次:获取数组A2 =[z1, z2, z3]

...

我想要相同的推理,但 A2 与 A1 不同。 我知道辍学面具会随着时间的推移而改变。 我想我需要“种子”变量,就像我们在 python API 中所拥有的那样。但我找不到任何参考资料。

我尝试了什么: 要获得相同的预测列表,我需要加载模型并多次初始化会话。

for (i = 0; i < 3; i++) 
    model = SavedModelBundle.load ("path_to_model", "serve");
    sess  = model.session();
    t_pred = sess.runner().feed("x", x).fetch("y").run().get(0);

但这不是最优的,因为加载模型需要时间并且可能导致与内存相关的问题。

我该如何解决这个问题?

提前谢谢你!

【问题讨论】:

    标签: java tensorflow seed dropout


    【解决方案1】:

    最后,我解决了这个问题。

    我错误地认为在重新打开会话时会重新初始化会话:Session s = modelBundle.session();

    它被重新初始化并包含一个图表。

    byte[] metaGraph = Files.readAllBytes(Paths.get(save_path));
    Graph g = new Graph();
    Session sess = new Session(g);
    

    但它会导致错误:

    “尝试使用未初始化的值”

    我通过更改在 python 中保存模型的方式修复了这个错误。

    以前,我用过:

    builder = tf.saved_model.builder.SavedModelBuilder(save_folder)
    builder.add_meta_graph_and_variables(sess,[tf.saved_model.tag_constants.SERVING])
    save_path = builder.save()
    

    它似乎没有保存种子(初始化为局部变量)然后导致模型不保存状态。

    我改成:

    with tf.gfile.GFile(save_path, 'wb') as f:
       f.write(out_graph_def.SerializeToString())
    

    而且效果很好^^。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2013-02-24
      • 1970-01-01
      • 2020-02-04
      • 1970-01-01
      • 2015-12-31
      • 2023-03-16
      • 2019-04-07
      • 2020-01-16
      相关资源
      最近更新 更多