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

TensorFlow会话初始化与神经网络模型可复现性问题排查

我来帮你排查这个可复现性问题,你遇到的情况其实是TensorFlow 1.x中很容易踩的几个坑,咱们逐个解决:

1. 随机种子设置时机完全错了!

你现在是在创建tf.Session()之后才调用tf.set_random_seed(777),这时候模型里的权重初始化等随机操作已经完成了,种子根本没生效!

正确的做法是在所有图构建代码的最开头就设置全局随机种子,还要同时固定NumPy的种子(避免数据处理环节的随机性):

import tensorflow as tf
import numpy as np

# 先把全局种子固定死,放在所有代码最前面
tf.set_random_seed(777)
np.random.seed(777)

# 之后再写你的循环、模型构建代码

另外,每次循环创建新模型前,记得重置默认图,避免旧图的残留变量干扰随机性:

for j in range(m_neuron): 
    n_neuron = 2**(j+1) 
    y_con[j] = n_neuron 
    print('n_layer: ',n_layer,'n_neuron:',n_neuron) 
    
    # 重置默认图,确保每次模型都是全新的
    tf.reset_default_graph()
    sess = tf.Session() 
    # 这里不要再重复设置tf.set_random_seed了!

2. Batch Normalization的动态更新在搞鬼

你的模型用了BN层,它的moving_mean和moving_variance是训练时动态累积更新的。而你用了早停逻辑(diff_tmp <= tol就停止训练),即使参数相同,每次训练的收敛步数可能有细微差异,导致最终的BN统计量不一样,直接影响预测结果。

解决办法二选一:

  • 去掉早停,用固定的n_epoch步数训练,确保BN的更新次数完全一致;
  • 把早停阈值设置得更严格(比如把tol调小几个数量级),或者增加连续满足阈值才停止的逻辑(比如连续3步满足才停),减少偶然停止的差异;
  • 最重要的:检查你的predict方法,一定要把BN的mode占位符设为False,确保预测时用的是训练好的累积统计量,而不是当前batch的临时统计量:
    # 你的Solver类的predict方法应该类似这样
    def predict(self, x_data):
        return self.sess.run(self.model.hypothesis, feed_dict={
            self.model.x: x_data,
            self.model.mode: False  # 必须设为False!
        })
    

3. R²计算里有个致命的索引错误!

我一眼看到你这段代码的问题:

for k in range(n_output):
    r2_train_tmp = m1_solver.evaluate_r2(y_train_scaled[:,i], y_train_predict[:,i])
    # ... 其他R²计算

这里你遍历的是k(输出维度),但取的却是外层循环的i(层数变量)的索引!这导致你每次计算的都是第i个输出的R²,而不是第k个,结果肯定乱掉,看起来像是模型不稳定,其实是你取错了维度!

赶紧把所有[:,i]改成[:,k]:

for k in range(n_output):
    r2_train_tmp = m1_solver.evaluate_r2(y_train_scaled[:,k], y_train_predict[:,k])
    r2_valid_tmp = m1_solver.evaluate_r2(y_valid_scaled[:,k], y_valid_predict[:,k])
    r2_test_tmp = m1_solver.evaluate_r2(y_test_scaled[:,k], y_test_predict[:,k])
    r2_train[j,i,k] = r2_train_tmp[0]
    r2_valid[j,i,k] = r2_valid_tmp[0]
    r2_test [j,i,k] = r2_test_tmp[0]

4. 会话初始化的小细节

每次循环里,你创建会话后初始化变量是对的,但要确保tf.global_variables_initializer()覆盖了所有变量(包括BN的moving_mean/variance),这个你当前代码是没问题的,但关闭会话后记得彻底释放资源:

sess.close()
tf.reset_default_graph()  # 放在close之后,确保下一次循环的图是干净的

按照上面的步骤修改后,你的模型应该就能实现完全可复现的结果了。

内容的提问来源于stack exchange,提问作者Keuntae Park

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 08:07:42