Advertisement

TensorFlow模型参数保存与加载,并附有示例代码

阅读量:

在利用TensorFlow搭建训练模型,例如人脸识别系统或场景分类网络时,若能够获取到合适的数据集,并经过长时间的训练后获得令人满意的识别效果,我们通常希望将训练结果保存下来。这样在后续使用过程中可以直接调用已保存的结果,而无需重复进行训练过程。这一需求直接引出了如何实现TensorFlow训练参数的保存与加载的问题。

为了便于用户存储训练成果,TensorFlow引入了tf.train.Saver模块,该模块可用于保留当前会话中所有变量的值(Variables)。在网络模型构建过程中,需要为计划保存的变量设置name属性,具体示例如下:

W = tf.Variable(tf.zeros([784, 10]), name="var_W")

当执行加载恢复操作(restore)时,只需通过变量名即可获取之前训练所保存的数值信息。相关代码如下所示:

W = sess.graph.get_tensor_by_name("var_W:0")

在完成整个训练过程后,仅需执行以下代码即可将训练结果存储至指定路径:

saver = tf.train.Saver()
saver_path = saver.save(sess, "%smodel.ckpt" % (SAVER_DIR))

以上步骤可归纳为以下几个关键环节:

  1. 在网络结构搭建阶段为

全部评论 (0)

还没有任何评论哟~