Keras与TensorFlow模型训练结果无法复现的问题求助
看起来你已经做了不少复现性相关的设置,但还是遇到了结果不一致的问题,我帮你梳理几个优先级从高到低的问题点:
最核心的错误:损失函数与任务不匹配
你的模型最后一层用了sigmoid激活,明显是做二分类任务,但你却指定了categorical_crossentropy作为损失函数。这个损失函数是给**多分类任务(标签为独热编码格式)**设计的,完全不适合当前的二分类场景(你的标签是单值0/1的float32类型)。
这个错误会直接导致训练过程极度不稳定,损失和精度波动极大,甚至每次运行结果完全不同。你需要立刻把损失函数改成binary_crossentropy,这应该能解决大部分问题。环境变量设置的顺序是否正确
你设置了TF_DETERMINISTIC_OPS等关键环境变量,但一定要确保这些设置是在导入TensorFlow/Keras之前完成的。如果你的代码里是先设置环境变量再导入TensorFlow,那没问题;但如果顺序反了,这些环境变量根本不会生效,依然会存在非确定性操作。
建议把环境变量设置放在代码最开头,比如:import os # 先配置所有确定性相关的环境变量 os.environ['TF_ENABLE_ONEDNN_OPTS'] = '0' os.environ["TF_DETERMINISTIC_OPS"] = "1" os.environ["TF_CUDNN_DETERMINISTIC"] = "1" # 再导入TensorFlow及其他依赖库 import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers import numpy as np import random from sklearn.model_selection import train_test_splitAdam优化器的潜在非确定性
即使设置了全局种子,某些旧版本的TensorFlow中,Adam优化器的部分并行操作可能依然存在非确定性。你可以试试这两个方案:- 升级到TensorFlow 2.10及以上的稳定版本,新版本对确定性模式下的Adam优化器支持更完善;
- 手动指定Adam优化器的
deterministic=True参数(需要TensorFlow版本支持),比如:optimizer = tf.keras.optimizers.Adam(deterministic=True) model.compile(loss="binary_crossentropy", optimizer=optimizer, metrics=["accuracy"])
验证集划分的确定性强化
你用了validation_split=0.1让Keras自动划分验证集,虽然全局种子理论上能控制这个划分,但为了100%消除不确定性,你可以手动划分训练集和验证集:# 先划分原始数据为「训练+验证集」和测试集 temp_data, test, temp_target, testtarget = train_test_split( df.loc[:, "input_alarm_1":"input_alarm_"+str(num_backtracking_events)], df.loc[:, "output_failure"], test_size=0.25, random_state=SEED ) # 再从「训练+验证集」中拆分出训练集和验证集 training, val, trainingtarget, valtarget = train_test_split( temp_data, temp_target, test_size=0.1/0.75, # 因为已经拆分了25%给测试集,这里要按比例取10%的原始训练数据 random_state=SEED ) # 训练时用validation_data指定验证集,替代validation_split history = model.fit(training, trainingtarget, batch_size=10, epochs=50, validation_data=(val, valtarget))这样能完全掌控验证集的划分逻辑,避免任何潜在的随机风险。
冗余设置的小优化
tf.keras.utils.set_random_seed(SEED)已经包含了random.seed(SEED)、np.random.seed(SEED)、tf.random.set_seed(SEED)这三个操作的功能,所以你后面重复写的这三行可以删掉,虽然不影响结果,但能让代码更简洁。
另外,建议你检查下TensorFlow版本,某些非常旧的版本在确定性模式下存在已知bug,升级到最新稳定版(比如2.15或更高)能避免这类问题。
备注:内容来源于stack exchange,提问作者yūrei

