Keras中Sequential与Functional API模型性能差异排查及修改方案
问题分析与解决方案
你的Functional API模型性能极差的核心原因是输出结构定义错误,以及后续的训练数据格式不匹配,具体问题和修复方案如下:
核心问题
输出层的嵌套列表错误
Sequential模型最后一层是Dense(7),输出是形状为(batch_size, 7)的单一张量;而你的Functional API代码中,out已经是7个Dense(1)张量的列表,但你错误地将模型输出设为[out],导致输出结构变成嵌套列表(外层列表包含一个内层列表,内层有7个张量)。这种结构和训练数据Y_tr的(batch_size,7)形状完全不匹配,Keras无法正确计算损失,训练完全偏离方向。训练数据格式不匹配
即使修复了输出结构,原有的Y_tr是单一张量,而多输出模型需要对应每个输出的独立标签张量,否则Keras无法正确分配损失计算。
修复后的Functional API代码
1. 修正模型定义
import tensorflow as tf from tensorflow import keras import numpy as np import random # 固定随机种子 np.random.seed(42) tf.random.set_seed(42) random.seed(42) targets = ('output1', 'output2', 'output3', 'output4', 'output5', 'output6', 'output7') inp = keras.Input(shape=18) hidden1 = keras.layers.Dense(1000, activation="relu")(inp) hidden2 = keras.layers.BatchNormalization()(hidden1) hidden3 = keras.layers.Dense(1000, activation="relu")(hidden2) hidden4 = keras.layers.BatchNormalization()(hidden3) # 生成7个独立的输出层 out_layers = [keras.layers.Dense(1, name=nn)(hidden4) for nn in targets] # 直接用out_layers作为输出,无需嵌套 modelF = keras.Model(inputs=inp, outputs=out_layers)
2. 调整训练数据格式并编译训练
# 拆分训练/验证标签为对应每个输出的张量(假设Y_tr/Y_val是形状为(batch_size,7)的数组) Y_tr_list = [Y_tr[:, i:i+1] for i in range(7)] Y_val_list = [Y_val[:, i:i+1] for i in range(7)] # 编译模型:给每个输出指定MSE损失,总损失默认是各输出损失的平均值(和Sequential模型的损失计算逻辑一致) modelF.compile( optimizer=keras.optimizers.Adam(learning_rate=1e-4), loss={name: keras.losses.mse for name in targets} ) # 传入拆分后的标签列表进行训练 history = modelF.fit( X_tr, Y_tr_list, epochs=1000, validation_data=(X_val, Y_val_list) )
3. 查看每个输出的独立MSE
训练完成后,history.history中会包含每个输出的训练/验证损失,例如val_output1_loss、val_output2_loss等,这些就是对应每个输出的独立MSE值。
额外说明
如果希望更清晰地映射输出和标签,可以用字典形式传入数据:
# 用字典关联输出名称和对应标签 Y_tr_dict = {name: Y_tr[:, i:i+1] for i, name in enumerate(targets)} Y_val_dict = {name: Y_val[:, i:i+1] for i, name in enumerate(targets)} history = modelF.fit( X_tr, Y_tr_dict, epochs=1000, validation_data=(X_val, Y_val_dict) )
这种修改后的模型结构和Sequential模型的权重数量、计算逻辑完全一致,训练后应该能得到和Sequential模型接近的验证MSE,同时可以单独获取每个输出的MSE。
内容的提问来源于stack exchange,提问作者MM Khan
相关产品推荐
相关产品推荐

