TensorFlow利用saver保存和提取参数的实例
2019/7/15 0:28:12
本文主要是介绍TensorFlow利用saver保存和提取参数的实例,对大家解决编程问题具有一定的参考价值,需要的程序猿们随着小编来一起学习吧!
在训练循环中,定期调用 saver.save() 方法,向文件夹中写入包含了当前模型中所有可训练变量的 checkpoint 文件。
saver.save(sess, FLAGS.train_dir, global_step=step)
global_step是训练的第几步
保存参数:
import tensorflow as tf W = tf.Variable([[1, 2, 3]], dtype=tf.float32) b = tf.Variable([[1]], dtype=tf.float32) saver = tf.train.Saver() sess = tf.InteractiveSession() tf.global_variables_initializer().run() # 必须要指定文件夹,保存到ckpt文件 save_path = saver.save(sess, "winycg/1.ckpt") print(save_path)
一次 saver.save() 后可以在文件夹中看到新增的四个文件,实际上每调用一次保存操作会创建后3个数据文件并创建一个检查点(checkpoint)文件,简单理解就是权重等参数被保存到 .chkp.data 文件中,以字典的形式;图和元数据被保存到 .chkp.meta 文件中,可以被 tf.train.import_meta_graph 加载到当前默认的图。
读取参数:
import tensorflow as tf import numpy as np W = tf.Variable(np.arange(3).reshape(1, 3), dtype=tf.float32) b = tf.Variable(np.arange(1).reshape(1, 1), dtype=tf.float32) saver = tf.train.Saver() sess = tf.InteractiveSession() # 读取参数时不需要global_variables_initializer() save_path = saver.restore(sess, "parameter/1.ckpt") print("weights:", sess.run(W)) print("bias:", sess.run(b))
weights: [[ 1. 2. 3.]]
bias: [[ 1.]]
以上这篇TensorFlow利用saver保存和提取参数的实例就是小编分享给大家的全部内容了,希望能给大家一个参考,也希望大家多多支持找一找教程网。
这篇关于TensorFlow利用saver保存和提取参数的实例的文章就介绍到这儿,希望我们推荐的文章对大家有所帮助,也希望大家多多支持为之网!
- 2024-05-08有遇到过吗?同样的规则 Excel 中 比Python 结果大
- 2024-03-30开始python成长之路
- 2024-03-29python optparse
- 2024-03-29python map 函数
- 2024-03-20invalid format specifier python
- 2024-03-18pool.map python
- 2024-03-18threads in python
- 2024-03-14python Ai 应用开发基础训练,字符串,字典,文件
- 2024-03-13id3 algorithm python
- 2024-03-13sum array elements python