TensorFlow中自定义MSE函数致Checkpoint体积递增问题咨询
get_mse()会让Checkpoint体积随训练轮次递增? 这个问题的核心原因很直白:你每次在训练循环里调用get_mse(true_ph, pred_ph),都会在TensorFlow的计算图中新增一套完整的MSE计算节点。随着训练轮次增加,计算图会越来越臃肿,Saver保存时会把整个计算图的结构和变量状态都包含进去,自然导致Checkpoint文件体积不断变大。
而用预先定义好的MSE节点时,整个训练过程中计算图不会有任何变化,所以每次保存的Checkpoint体积都是稳定的。
具体分析你的代码问题
你的get_mse函数内部会创建squared_difference、compute_weighted_loss等一系列运算节点,而且默认会把生成的loss加入到ops.GraphKeys.LOSSES集合中。当你在tf.Session()的循环里每次调用这个函数时,TensorFlow都会在当前计算图中追加这些新节点——相当于每训练一轮,你的计算图就多一套MSE计算逻辑,保存的文件自然越来越大。
两种解决方法
方法1:提前在计算图构建阶段定义好自定义MSE节点
把get_mse的调用移到tf.Session()外面,只创建一次节点,循环里只重复运行它:
import os import numpy as np import tensorflow as tf from tensorflow.python.ops.losses.losses_impl import Reduction, compute_weighted_loss from tensorflow.python.framework import ops from tensorflow.python.ops import math_ops def get_mse( labels, predictions, weights=1.0, name="mse", scope=None, loss_collection=ops.GraphKeys.LOSSES, reduction=Reduction.SUM_BY_NONZERO_WEIGHTS): with ops.name_scope(scope, name, (predictions, labels, weights)) as scope: predictions = math_ops.to_float(predictions) labels = math_ops.to_float(labels) predictions.get_shape().assert_is_compatible_with(labels.get_shape()) losses = math_ops.squared_difference(predictions, labels) return compute_weighted_loss( losses, weights, scope, loss_collection, reduction=reduction) true = np.random.random((100,1)) pred = np.random.random((100,1)) variable_to_save = tf.Variable(true) true_ph = tf.placeholder(tf.float32, shape = [None, 1], name='labels') pred_ph = tf.placeholder(tf.float32, shape = [None, 1], name='predictions') MSE = tf.reduce_mean(tf.square(pred_ph - true_ph), name='get_mse') # 关键:提前定义好自定义MSE节点,只创建一次 custom_mse = get_mse(true_ph, pred_ph) with tf.Session() as sess: init = tf.global_variables_initializer() sess.run(init) saver = tf.train.Saver(max_to_keep=100) for epoch in range(10): # 运行预先创建好的节点,不会新增图结构 mse = sess.run(custom_mse, feed_dict={'labels:0': true, 'predictions:0':pred}) saver.save(sess, os.getcwd(), global_step=epoch)
方法2:修改get_mse函数,不将loss加入全局集合
如果你确实需要在循环里动态调用这个函数(虽然这里没必要),可以把loss_collection参数设为None,这样生成的loss不会被加入到全局的loss集合中,Saver就不会追踪这些临时生成的节点:
def get_mse( labels, predictions, weights=1.0, name="mse", scope=None, # 将loss_collection设为None,避免每次添加到全局集合 loss_collection=None, reduction=Reduction.SUM_BY_NONZERO_WEIGHTS): with ops.name_scope(scope, name, (predictions, labels, weights)) as scope: predictions = math_ops.to_float(predictions) labels = math_ops.to_float(labels) predictions.get_shape().assert_is_compatible_with(labels.get_shape()) losses = math_ops.squared_difference(predictions, labels) return compute_weighted_loss( losses, weights, scope, loss_collection, reduction=reduction) # 后续循环可以保持原来的调用方式,不会导致图膨胀 with tf.Session() as sess: init = tf.global_variables_initializer() sess.run(init) saver = tf.train.Saver(max_to_keep=100) for epoch in range(10): mse = sess.run(get_mse(true_ph, pred_ph), feed_dict={'labels:0': true, 'predictions:0':pred}) saver.save(sess, os.getcwd(), global_step=epoch)
总结
TensorFlow 1.x的计算图是静态的,不要在会话循环里重复创建运算节点——这不仅会导致Checkpoint体积膨胀,还会拖慢训练速度,因为每次都要构建新的计算逻辑。如果是TensorFlow 2.x的动态图模式,这个问题就不会出现,但在1.x版本里必须注意图结构的复用。
内容的提问来源于stack exchange,提问作者Atr Cheema

