Advertisement

TensorFlow和Keras模型的保存与加载参数用于继续训练

阅读量:

TensorFlow框架应用

在TensorFlow框架中,实现模型参数的存储与调用主要依赖于Saver()这一功能模块。完成首次训练过程后,可通过调用save函数进行参数保存;而在进行预测操作之前,或需要继续训练时,则需通过load函数加载已存储的参数信息。

复制代码
    def __init__():
    	self.sess = tf.Session()
    	# 定义好网络结构...
    	self.sess.run(tf.global_variables_initializer())
    def check_path(self, path):
    if not os.path.exists(path):
        os.mkdir(path)
    def save(self):
    self.check_path('model')
    saver=tf.train.Saver(tf.global_variables(),max_to_keep=10)
    print("model: ",saver.save(self.sess,'model/modle.ckpt'))
    
    def load(self):
    saver=tf.train.Saver(tf.global_variables())

全部评论 (0)

还没有任何评论哟~