
2. The Saver class
Saver类是TensorFlow库供给的类,它是保存图形构造和变量的首选办法。
在以下几行代铝闼楝我们定义一个Saver对象,并在train_graph()函数中,经由100次迭代的办法最小化成本函数。然后,在每次迭代中以及优化完成后,将模型保存稻磁逄。每个保存在磁盘上创建二进制文件被称为“检查点”。
- # Create a Saver object
- saver = tf.train.Saver()
- init = tf.global_variables_initializer()
- # Run a session. Go through 100 iterations to minimize the cost
- def train_graph():
- with tf.Session() as sess:
- sess.run(init)
- for i in range(100):
- for (x, y) in zip(x_train, y_train):
- # Feed actual data to the train operation
- sess.run(trainop, feed_dict={X: x, Y: y})
- # Create a checkpoint in every iteration
- saver.save(sess, 'model_iter', global_step=i)
- # Save the final model
- saver.save(sess, 'model_final')
- h_ = sess.run(h_est)
- v_ = sess.run(v_est)
- return h_, v_
如今让我们用上述功能练习模型,并打印出练习的参数。

网友点评
精彩导读
科技快报
品牌展示