【问题标题】:How to correctly restore a OOP tensorflow model?如何正确恢复 OOP tensorflow 模型?
【发布时间】:2018-07-04 09:44:12
【问题描述】:

为了当前的项目,我决定在一个类实例中定义一个 tensorflow 模型。这一切都很好,直到我想恢复它以从最新的检查点继续训练。它是一个简单的线性回归模型,建立在实例的初始化之上。它试图逼近函数f(x) = 3x + 1。

逻辑是:如果还没有检查点,则创建一个新模型,将其训练 20 个 epoch,然后保存它。如果已经有一个检查点,则加载它,并从它继续训练 20 个 epoch。

现在,最初训练网络是可行的。但是在加载后尝试训练它时,它会抛出以下错误:

文件“”,第 1 行,在 runfile('/home/abc/tf_tests/restore_test/restoretest.py', wdir='/home/sku/tf_tests/restore_test')

文件 “/home/abc/anaconda3/envs/tensorflow/lib/python3.5/site-packages/spyder/utils/site/sitecustomize.py”, 第 710 行,在运行文件中 execfile(文件名,命名空间)

文件 “/home/abc/anaconda3/envs/tensorflow/lib/python3.5/site-packages/spyder/utils/site/sitecustomize.py”, 第 101 行,在 execfile 中 exec(编译(f.read(),文件名,'exec'),命名空间)

文件“/home/sku/tf_tests/restore_test/restoretest.py”,第 71 行,在 model.run_training_step(sess, x, y)

NameError:名称“模型”未定义

问题是:如何恢复它并正确进行训练?我发现了一篇关于 OOP here 的有趣文章,但它不涉及保存和恢复模型。

我的代码如下。谢谢你帮助我!

import tensorflow as tf
import numpy as np
import matplotlib.pyplot as plt

class LinearModel(object):

    def __init__(self):
        self.build_model()

    def build_model(self):
        # x is input, y is output
        self.x = tf.placeholder(dtype=tf.float32, name='x')
        self.y = tf.placeholder(dtype=tf.float32, name='y')

        self.w = tf.Variable(0.0, name='w')
        self.b = tf.Variable(0.0, name='b')

        self.global_step = tf.Variable(0, trainable=False, name='global_step', dtype=tf.int32)

        self.y_pred = self.w * self.x + self.b

        # quadratic error as loss
        self.loss = tf.square(self.y - self.y_pred)

        self.train_op = tf.train.AdamOptimizer(0.001).minimize(self.loss)
        self.increment_global_step_op = tf.assign(self.global_step, self.global_step+1)

        return 

    # run a single (x, y) pair through the graph
    def run_training_step(self, sess, x, y):
        _, loss = sess.run([self.train_op, self.loss], feed_dict={self.x:x, self.y:y})
        return loss

    # convenience function for checking the values
    def get_vars(self, sess):
        return sess.run([self.w, self.b])


tf.reset_default_graph()

# training data generation, is a linear function of 3x+1 + noise
tr_input = np.linspace(-5.0, 5.0)
tr_output = 3*tr_input+1+np.random.randn(tr_input.shape[0])


with tf.Session() as sess:

    # check if there are checkpoints
    latest_checkpoint = tf.train.latest_checkpoint('./model_saves')

    # ADDED BY EDIT1
    model = LinearModel()

    # if there are, load them
    if latest_checkpoint:

        saver = tf.train.import_meta_graph('./model_saves/lin_model-20.meta')
        saver.restore(sess, latest_checkpoint)  

    # if not, create a new model
    else:

        ### REMOVED BY EDIT1
        ### model = LinearModel()
        sess.run(tf.global_variables_initializer())

        saver = tf.train.Saver()

    # show vars before doing the training
    w, b = model.get_vars(sess)       
    print("final weight: {}".format(w))
    print("final bias: {}".format(b))

    # train for 20 epochs and save it
    for epoch in range(20):
        for x, y in zip(tr_input, tr_output):
            model.run_training_step(sess, x, y)
        sess.run(model.increment_global_step_op)

    saver.save(sess, './model_saves/lin_model', global_step=model.global_step)       

    # show vars after doing the training
    w_opt, b_opt = model.get_vars(sess)       
    print("final weight: {}".format(w_opt))
    print("final bias: {}".format(b_opt))

EDIT1:

在检查是否存在检查点之前实例化模型时,会导致优化器变量的前置条件错误:

FailedPreconditionError: 尝试使用未初始化的值 beta1_power [[节点:beta1_power/read = IdentityT=DT_FLOAT, _class=["loc:@Adam/Assign"], _device="/job:localhost/replica:0/task:0/device:GPU:0"]] [[节点: Square/_25 = _Recvclient_terminated=false, recv_device="/job:localhost/replica:0/task:0/device:CPU:0", send_device="/job:localhost/replica:0/task:0/device:GPU:0", send_device_incarnation=1, tensor_name="edge_103_Square", tensor_type=DT_FLOAT, _device="/job:localhost/replica:0/task:0/device:CPU:0"]] ...

【问题讨论】:

    标签: python oop tensorflow


    【解决方案1】:

    当您尝试从检查点恢复时,您没有实例化您的 LinearModel 类。这应该有效:

    ...
    latest_checkpoint = tf.train.latest_checkpoint('/home/sku/tf_tests/restore_test/model_saves')
    
    model = LinearModel()
    saver = tf.train.Saver()
    
    if latest_checkpoint:
        saver.restore(sess, latest_checkpoint)
    else:
        sess.run(tf.global_variables_initializer())
    ...
    

    【讨论】:

    • 谢谢,我试过了,但结果是一个优化器变量的FailedPreconditionError。我会将错误编辑到我的问题中。
    • 您是否在任何时候更改过优化器?我修改了答案以适应您的新错误。
    • 不,我没有。如果我尝试使用您修改后的答案,则培训总是从头开始,例如w 和 b 变量的值为零。但是错误消失了!
    • 我再次修改了我的答案。问题是您在使用 import_meta_graph 时再次导入了图表。您不需要这条线,因为您已经在 build_model 函数中定义了图形。
    • 修复了它,现在非常有意义。也许这是一个功能请求的主题,因此您在执行此操作时会收到警告。
    猜你喜欢
    • 2018-02-02
    • 2017-08-21
    • 2018-09-06
    • 2016-05-01
    • 1970-01-01
    • 2018-09-25
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多