【问题标题】:Saving and restoring functions in TensorFlow在 TensorFlow 中保存和恢复函数
【发布时间】:2018-11-11 06:14:55
【问题描述】:

我正在 TensorFlow 中开展一个 VAE 项目,其中编码器/解码器网络内置于函数中。这个想法是能够保存,然后加载训练好的模型并使用编码器功能进行采样。

恢复模型后,我无法让解码器函数运行并将恢复的、经过训练的变量返回给我,出现“未初始化值”错误。我认为这是因为该功能要么创建一个新的,要么覆盖现有的,或者以其他方式。但我不知道如何解决这个问题。这是一些代码:

class VAE(object):    
    def __init__(self, restore=True):
        self.session = tf.Session()
        if restore:
            self.restore_model()
            self.build_decoder = tf.make_template('decoder', self._build_decoder)

@staticmethod
def _build_decoder(z, output_size=768, hidden_size=200,
                  hidden_activation=tf.nn.elu, output_activation=tf.nn.sigmoid):
    x = tf.layers.dense(z, hidden_size, activation=hidden_activation)
    x = tf.layers.dense(x, hidden_size, activation=hidden_activation)
    logits = tf.layers.dense(x, output_size, activation=output_activation)
    return distributions.Independent(distributions.Bernoulli(logits), 2)

def sample_decoder(self, n_samples):
    prior = self.build_prior(self.latent_dim)
    samples = self.build_decoder(prior.sample(n_samples), self.input_size).mean()
    return self.session.run([samples])

def restore_model(self):
    print("Restoring")
    self.saver = tf.train.import_meta_graph(os.path.join(self.save_dir, "turbolearn.meta"))
    self.saver.restore(self.sess, tf.train.latest_checkpoint(self.save_dir))
    self._restored = True

想跑samples = vae.sample_decoder(5)

在我的日常训练中,我会跑步:

        if self.checkpoint:
            self.saver.save(self.session, os.path.join(self.save_dir, "myvae"), write_meta_graph=True)

更新

根据以下建议的答案,我更改了恢复方法

self.saver = tf.train.Saver()
self.saver.restore(self.session, tf.train.latest_checkpoint(self.save_dir))

但是现在在创建 Saver() 对象时得到一个值错误:

ValueError: No variables to save

【问题讨论】:

    标签: python tensorflow machine-learning tensorflow-probability


    【解决方案1】:

    tf.train.import_meta_graph 恢复图形,意味着重建存储到文件中的网络架构。另一方面,对tf.train.Saver.restore 的调用仅将文件中的变量值恢复到会话中的当前图形(如果文件中的某些值属于当前活动图形中不存在的变量,这自然会失败) .

    因此,如果您已经在代码中构建了网络层,则无需调用tf.train.import_meta_graph。否则这可能会给您带来问题。

    不确定您的代码的其余部分如何,但这里有一些建议。首先构建图,然后创建会话,如果适用,最后恢复。你的 init 可能看起来像这样

    def __init__(self, restore=True):
        self.build_decoder = tf.make_template('decoder', self._build_decoder)
        self.session = tf.Session()
        if restore:
            self.restore_model()
    

    但是,如果您只是恢复编码器并重新构建解码器,您可能会最后构建解码器。但是不要忘记在使用前初始化它的变量。

    【讨论】:

    • 感觉我在这里遗漏了一些东西。我尝试按照您的建议替换该行,调用构建标准保护程序对象,然后调用 restore 。有关代码更改,请参阅上面的更新
    • @taylormade201 tf.train.Saver 必须在创建变量(图表)之后创建。因此,在您的代码中,您应该注意创建图形和恢复变量的顺序。我用更详细的建议更新了答案
    猜你喜欢
    • 1970-01-01
    • 2016-08-16
    • 2016-04-02
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-01-16
    相关资源
    最近更新 更多