"); //-->
本篇文章小编给大家分享一下扣丁学堂Python在线教程TensorFlow入门使用 tf.train.Saver()保存模型,希望可以帮到对Python开发感兴趣的小伙伴们。

关于模型保存的一点心得
saver = tf.train.Saver(max_to_keep=3)
在定义 saver 的时候一般会定义最多保存模型的数量,一般来说,如果模型本身很大,我们需要考虑到硬盘大小。如果你需要在当前训练好的模型的基础上进行 fine-tune,那么尽可能多的保存模型,后继 fine-tune 不一定从最好的 ckpt 进行,因为有可能一下子就过拟合了。但是如果保存太多,硬盘也有压力呀。如果只想保留最好的模型,方法就是每次迭代到一定步数就在验证集上计算一次 accuracy 或者 f1 值,如果本次结果比上次好才保存新的模型,否则没必要保存。
如果你想用不同 epoch 保存下来的模型进行融合的话,3到5 个模型已经足够了,假设这各融合的模型成为 M,而最好的一个单模型称为 m_best, 这样融合的话对于M 确实可以比 m_best 更好。但是如果拿这个模型和其他结构的模型再做融合的话,M 的效果并没有 m_best 好,因为M 相当于做了平均操作,减少了该模型的“特性”。
但是又有一种新的融合方式,就是利用调整学习率来获取多个局部最优点,就是当 loss 降不下了,保存一个 ckpt, 然后开大学习率继续寻找下一个局部最优点,然后用这些 ckpt 来做融合,还没试过,单模型肯定是有提高的,就是不知道还会不会出现上面再与其他模型融合就没提高的情况。
如何使用 tf.train.Saver() 来保存模型
之前一直出错,主要是因为坑爹的编码问题。所以要注意文件的路径绝对不不要出现什么中文呀。
import tensorflow as tf
config = tf.ConfigProto()
config.gpu_options.allow_growth = True
sess = tf.Session(config=config)
# Create some variables.
v1 = tf.Variable([1.0, 2.3], name="v1")
v2 = tf.Variable(55.5, name="v2")
# Add an op to initialize the variables.
init_op = tf.global_variables_initializer()
# Add ops to save and restore all the variables.
saver = tf.train.Saver()
ckpt_path = './ckpt/test-model.ckpt'
# Later, launch the model, initialize the variables, do some work, save the
# variables to disk.
sess.run(init_op)
save_path = saver.save(sess, ckpt_path, global_step=1)
print("Model saved in file: %s" % save_path)Model saved in file: ./ckpt/test-model.ckpt-1
注意,在上面保存完了模型之后。应该把 kernel restart 之后才能使用下面的模型导入。否则会因为两次命名 “v1” 而导致名字错误。
import tensorflow as tf
config = tf.ConfigProto()
config.gpu_options.allow_growth = True
sess = tf.Session(config=config)
# Create some variables.
v1 = tf.Variable([11.0, 16.3], name="v1")
v2 = tf.Variable(33.5, name="v2")
# Add ops to save and restore all the variables.
saver = tf.train.Saver()
# Later, launch the model, use the saver to restore variables from disk, and
# do some work with the model.
# Restore variables from disk.
ckpt_path = './ckpt/test-model.ckpt'
saver.restore(sess, ckpt_path + '-'+ str(1))
print("Model restored.")
print sess.run(v1)
print sess.run(v2)INFO:tensorflow:Restoring parameters from ./ckpt/test-model.ckpt-1
Model restored.
[ 1. 2.29999995]
55.5
导入模型之前,必须重新再定义一遍变量。
但是并不需要全部变量都重新进行定义,只定义我们需要的变量就行了。
也就是说,你所定义的变量一定要在 checkpoint 中存在;但不是所有在checkpoint中的变量,你都要重新定义。
import tensorflow as tf
config = tf.ConfigProto()
config.gpu_options.allow_growth = True
sess = tf.Session(config=config)
# Create some variables.
v1 = tf.Variable([11.0, 16.3], name="v1")
# Add ops to save and restore all the variables.
saver = tf.train.Saver()
# Later, launch the model, use the saver to restore variables from disk, and
# do some work with the model.
# Restore variables from disk.
ckpt_path = './ckpt/test-model.ckpt'
saver.restore(sess, ckpt_path + '-'+ str(1))
print("Model restored.")
print sess.run(v1)INFO:tensorflow:Restoring parameters from ./ckpt/test-model.ckpt-1
Model restored.
[ 1. 2.29999995]
tf.Saver([tensors_to_be_saved]) 中可以传入一个 list,把要保存的 tensors 传入,如果没有给定这个list的话,他会默认保存当前所有的 tensors。
以上就是扣丁学堂Python在线教程TensorFlow入门使用 tf.train.Saver()保存模型,想要了解更多Python内容的小伙伴可以登录扣丁学堂官网咨询,想要学好Python的小伙伴可以选择扣丁学堂Python培训机构学习,扣丁学堂不仅有专业的老师和与时俱进的课程体系,还有大量的Python在线教程供学员观看学习,喜欢的小伙伴可以登录官网学习哦。扣丁学堂python学习交流群:816572891。微信号:codingbb
专栏文章内容及配图由作者撰写发布,仅供工程师学习之用,如有侵权或者其他违规问题,请联系本站处理。 联系我们
相关推荐
关于在 KEIL C51 中嵌入汇编以及C51与A51间的相互调用
500Hz信号发生器
请问一下,加载VXWORKS的CPU需要MMU吗
红色飓风FPGA普及行动 第六讲:SoPC硬件系统
半导体晶圆代工市场预计到 2032 年将实现全面增长
800Hz振荡器
2007智能车决赛展示视频清华大学第一代表队三角洲队
800Hz信号发生器
高性能FLASH存储器在DSP电机智能保护中的应用
关于模块逻辑固化的问题
关于AVR I-O 的驱动能力的介绍
ASURO循迹避障测速智能车平衡功能演示
抢先台积电 传三星将在美国推出2纳米
三星推动2纳米技术突破,业界押注台积电3纳米鳍式场效应晶体管
蚊式无人机专为中国间谍任务而设计 — 军事机器人实验室展示极其微型的仿生飞行机器人
一季度全球智能摄像头增速放缓至4.6%!中国厂商占主导,拉美/亚太成市场新引擎!
台积电3纳米FinFET 三星2纳米恐难敌
1kHz信号发生器
2009年3月北京邮电大学第一届校园智能车大赛
高性能数据采集系统芯片LM12H458及其应用
EP1SGX10CF672C6ES的“ES”代表什么?
450音频信号发生器
618恐成今年补贴最后一波 IC设计坦言需求难回天
据报道英特尔将于7月15日开始大规模裁员,并逐步缩减汽车部门
【求助】 用usb接口达成通信功能-有会编程的么,必有重谢!
TDK推出具备扩展工作温度范围且面向全球分销的 SmartAutomotive™ 6轴 IMU
杭州博通软件,从事手机mmi的设计和多种语言改制, 如果你感兴趣, 请发送你的简历到mlu92_1@163.com , 谢谢!
富士通2纳米CPU交台积电代工 锁定AI与数据中心应用
2009年第四届智能车竞赛北京科技大学表演车模(光电组)
和printf一样具有可变参数的C51函数