Keras模型训练遭遇鞍点问题的解决咨询
Keras模型训练损失停滞问题及优化建议
问题描述
我的Keras类VGG-16模型训练过程中损失停滞在0.0025,疑似陷入鞍点,尝试多种方法均无法进一步降低损失,且当前模型预测效果不佳,寻求优化方案。目前最多训练7轮,其中6轮损失及验证损失无变化,不确定0.0025是否为数据集全局最小值,也不确定是否需要延长训练时长。
已尝试方案
- 使用Adam和RMSProp优化器,搭配或不搭配循环学习率(范围0.001至0.1),损失始终维持在0.0989
- 训练4-5轮无进展后换用SGD,损失稳步降至0.0025后停滞;后续尝试带循环学习率的SGD,结果一致
- 调整网络容量(提升4个全连接层至4096单元或降低),无改善
- 尝试不同批次大小
模型代码
# imports pip install -q -U tensorflow_addons import tensorflow_addons as tfa import tensorflow as tf from tensorflow import keras from keras import layers def get_model(input_shape): input = keras.input(shape=input_shape) x = layers.Conv2D(filters=64, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=64, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.MaxPooling2D(pool_size=(2, 2) strides=none, paddings="same")(x) x = layers.Conv2D(filters=128, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=128, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.MaxPooling2D(pool_size=(2, 2) strides=none, paddings="same")(x) x = layers.Conv2D(filters=256, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=256, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=256, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=256, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.MaxPooling2D(pool_size=(2, 2) strides=none, paddings="same")(x) x = layers.Conv2D(filters=512, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=512, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=512, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.Conv2D(filters=512, kernel_size= (3, 3), activation='relu', paddings="same")(input) x = layers.MaxPooling2D(pool_size=(2, 2) strides=none, paddings="same")(x) x = layers.Flatten()(x) x = layers.Dense(4096, activation='relu')(x) x = layers.Dense(2048, activation='relu')(x) x = layers.Dense(1024, activation='relu')(x) x = layers.Dense(512, activation='relu')(x) output = layers.Dense(9, activation='sigmoid')(x) return keras.models.Model(inputs=input, outputs=output) # define learning rate range lr_range = [.001, .1] epochs = 100 batch_size = 32 steps_per_epoch = len(training_data)/batch_size clr = tfa.optimizers.CyclicalLearningRate(initial_learning_rate=lr_range[0], maximal_learning_rate=lr_range[1], scale_fn=lambda x: 1/(2.**(x-1)), step_size=2 * steps_per_epoch ) optimizer = tf.keras.optimizers.Adam(clr) model = get_model((224, 224, 3)) model.compile(optimzer=optimzer, loss='mean_squared_error') # used tf.dataset objects for model input model.fit(train_ds, validation_data=valid_ds, batch_size=batch_size, epochs=epochs)
优化建议
1. 修复致命模型结构错误
你的卷积层全部直接连接到原始input,没有形成链式递进结构!每一组卷积都应该以上一层的输出x作为输入,而不是重复使用input。修正后的第一组卷积示例:
x = layers.Conv2D(filters=64, kernel_size=(3, 3), activation='relu', padding="same")(input) x = layers.Conv2D(filters=64, kernel_size=(3, 3), activation='relu', padding="same")(x) # 这里用x代替input x = layers.MaxPooling2D(pool_size=(2, 2), strides=None, padding="same")(x)
后续所有卷积组都需要做同样修改,否则网络无法完成逐层特征提取,之前的训练完全无效。
2. 修正代码拼写与参数错误
- 将所有
paddings改为padding(Keras官方参数名为padding) - 将
strides=none改为strides=None(Python中None需大写) - 修正
model.compile中的拼写错误:optimzer=optimzer改为optimizer=optimizer
3. 匹配损失函数与任务类型
如果是多标签分类任务,当前sigmoid输出+mean_squared_error损失的组合是可行的;如果是多分类任务,应改为softmax输出+SparseCategoricalCrossentropy(标签为整数)或CategoricalCrossentropy(标签为独热编码)损失。
4. 调整训练策略
- 修复结构后,先使用固定小学习率(比如0.0001)训练,观察损失变化,循环学习率可在模型稳定后再尝试
- 加入正则化:在全连接层后添加
layers.Dropout(0.5),或在卷积/全连接层中加入kernel_regularizer=keras.regularizers.L2(1e-4),防止过拟合 - 启用早停机制:添加
keras.callbacks.EarlyStopping(monitor='val_loss', patience=5, restore_best_weights=True),监控验证损失,自动停止无效训练并保留最优权重 - 延长训练时长:修复结构后的模型需要至少30-50轮训练才能看到明显效果
5. 规范数据预处理
确保输入图像做了标准化处理,比如归一化到[0,1]区间,或使用ImageNet数据集的均值和方差做标准化:
preprocess_input = keras.applications.vgg16.preprocess_input train_ds = train_ds.map(lambda x, y: (preprocess_input(x), y)) valid_ds = valid_ds.map(lambda x, y: (preprocess_input(x), y))
内容的提问来源于stack exchange,提问作者junfanbl
相关产品推荐
相关产品推荐

