【问题标题】:How to add some new variables to a loaded checkpoints in tensorflow?如何将一些新变量添加到张量流中加载的检查点?
【发布时间】:2018-05-21 23:50:43
【问题描述】:

我在张量流中训练了一个大图,并通过以下函数将它们保存在检查点中,

def save_model(sess, saver, param_folder, saved_ckpt):
    print("Saving model to disk...")
    address = os.path.join(param_folder, 'model')
    if not os.path.isdir(address):
        os.makedirs(address)
    address = os.path.join(address, saved_ckpt)
    save_path = saver.save(sess, address)
    saver.export_meta_graph(filename=address+'.meta')
    print("Model saved in file: %s" % save_path)

现在,为了加载图表,我使用了以下函数。

def load_model(sess, saver, param_folder, saved_ckpt):
    print("loding model from disk...")
    address = os.path.join(param_folder, 'model')
    if not os.path.isdir(address):
        os.makedirs(address)
    address = os.path.join(address, saved_ckpt)
    print("meta graph address :", address)
    saver = tf.train.import_meta_graph(address+'.meta')
    saver.restore(sess, address)

TensorFlow 的一个很棒的功能是它会自动将保存的权重分配给检查点所需的图。但是,当我想将图形(保存在检查点中的图形)加载到与我保存的图形略有不同/扩展的图形中时,就会出现问题。比如,假设我在上一个图中添加了一个额外的神经网络,并且想要从上一个检查点加载权重,这样我就不必从一开始就训练模型。

简而言之,我的问题是,如何将之前保存的子图加载到更大的(或者你可以说是父图)图中?

【问题讨论】:

标签: python tensorflow save


【解决方案1】:

我也遇到了这个问题,我用@rvinas评论。所以只是为了让下一个读者更容易。

当您加载保存的变量时,您可以在 restore_dict 中添加/删除/编辑它们,如下所示:

def load_model(sess, saver, param_folder, saved_ckpt):
    print("loding model from disk...")
    address = os.path.join(param_folder, 'model')
    if not os.path.isdir(address):
        os.makedirs(address)
    address = os.path.join(address, saved_ckpt)
    print("meta graph address :", address)
    # remove the next two lines
    # saver = tf.train.import_meta_graph(address+'.meta')
    # saver.restore(sess, address)
    # instead put this block:

    reader = tf.train.NewCheckpointReader(address)
    restore_dict = dict()
    for v in tf.trainable_variables():
      tensor_name = v.name.split(':')[0]
      if reader.has_tensor(tensor_name):
        print('has tensor ', tensor_name)
        restore_dict[tensor_name] = v
        # put the logic of the new/modified variable here and assign to the restore_dict, i.e. 
        # restore_dict['my_var_scope/my_var'] = get_my_variable()

希望对您有所帮助。

【讨论】:

    猜你喜欢
    • 2022-01-06
    • 1970-01-01
    • 2017-12-15
    • 1970-01-01
    • 2021-10-22
    • 1970-01-01
    • 2021-10-25
    • 2017-08-22
    • 1970-01-01
    相关资源
    最近更新 更多