作家
登录

如何保存和恢复TensorFlow训练的模型

作者: 来源: 2017-11-02 09:27:06 阅读 我要评论

1   
  • v = -2 
  • # Generate training data with noise 
  • x_train = np.linspace(-2,4,201)   
  • noise = np.random.randn(*x_train.shape) * 0.4   
  • y_train = (x_train - h) ** 2 + v + noise 
  • # Visualize the data  
  • plt.rcParams['figure.figsize'] = (10, 6)   
  • plt.scatter(x_train, y_train)   
  • plt.xlabel('x_train')   
  • plt.ylabel('y_train') 
  • 2. The Saver class

    Saver类是TensorFlow库供给的类,它是保存图形构造和变量的首选办法。

    在以下几行代铝闼楝我们定义一个Saver对象,并在train_graph()函数中,经由100次迭代的办法最小化成本函数。然后,在每次迭代中以及优化完成后,将模型保存稻磁逄。每个保存在磁盘上创建二进制文件被称为“检查点”。

    1. # Create a Saver object 
    2. saver = tf.train.Saver() 
    3.  
    4. init = tf.global_variables_initializer() 
    5.  
    6. # Run a session. Go through 100 iterations to minimize the cost 
    7. def train_graph():   
    8.     with tf.Session() as sess: 
    9.         sess.run(init) 
    10.         for i in range(100): 
    11.             for (x, y) in zip(x_train, y_train): 
    12.  
    13.                 # Feed actual data to the train operation 
    14.                 sess.run(trainop, feed_dict={X: x, Y: y}) 
    15.  
    16.             # Create a checkpoint in every iteration 
    17.             saver.save(sess, 'model_iter', global_step=i
    18.  
    19.         # Save the final model 
    20.         saver.save(sess, 'model_final') 
    21.         h_ = sess.run(h_est) 
    22.         v_ = sess.run(v_est) 
    23.     return h_, v_ 

    如今让我们用上述功能练习模型,并打印出练习的参数。

    1. result = train_graph()   

        推荐阅读

        开发人员该如何对超级计算机进行编程?

      有些编程技巧针对的是当前或者将来的超等计算机,固然它们已经存在了很长时光,但如今很多开辟人员并没有对这>>>详细阅读


      本文标题:如何保存和恢复TensorFlow训练的模型

      地址:http://www.17bianji.com/lsqh/38364.html

    关键词: 探索发现

    乐购科技部分新闻及文章转载自互联网,供读者交流和学习,若有涉及作者版权等问题请及时与我们联系,以便更正、删除或按规定办理。感谢所有提供资讯的网站,欢迎各类媒体与乐购科技进行文章共享合作。

    网友点评
    自媒体专栏

    评论

    热度

    精彩导读
    栏目ID=71的表不存在(操作类型=0)