You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.12 05:36:25