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)
还没有任何评论哟~
