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

基于变分推断的MLP回归训练中验证损失随机尖峰问题求解

变分推断BNN回归任务中验证损失随机尖峰的解决方法

问题描述

我用变分推断训练一个基于Flipout的贝叶斯MLP,在仅含1个特征的小数据集上执行回归任务。训练损失持续下降,但验证损失出现随机尖峰,不清楚如何解决。实现代码如下:

import tensorflow_probability as tfp
import tensorflow as tf
 
from tensorflow.keras.layers import Input
from tensorflow.keras.layers import Dense
from tensorflow.keras.models import Model
from tensorflow.keras.optimizers import Adam

def create_flipout_bnn_model(train_size):
  def normal_sp(params): 
      return tfd.Normal(loc=params[:,0:1], scale=1e-3 + tf.math.softplus(0.05 * params[:,1:2]))

  kernel_divergence_fn=lambda q, p, _: tfp.distributions.kl_divergence(q, p) / (train_size)
  bias_divergence_fn=lambda q, p, _: tfp.distributions.kl_divergence(q, p) / (train_size)


  inputs = Input(shape=(1,),name="input layer")


  hidden = tfp.layers.DenseFlipout(30,
                           kernel_divergence_fn=kernel_divergence_fn,
                           activation="relu",name="DenseFlipout_layer_1")(inputs)
  hidden = tfp.layers.DenseFlipout(30,
                           kernel_divergence_fn=kernel_divergence_fn,
                           activation="relu",name="DenseFlipout_layer_2")(hidden)
  hidden = tfp.layers.DenseFlipout(30,
                           kernel_divergence_fn=kernel_divergence_fn,
                           activation="relu",name="DenseFlipout_layer_3")(hidden)
  params = tfp.layers.DenseFlipout(2,
                           kernel_divergence_fn=kernel_divergence_fn,
                           name="DenseFlipout_layer_5")(hidden)
  dist = tfp.layers.DistributionLambda(normal_sp,name = 'normal_sp')(params) 

  model = Model(inputs=inputs, outputs=dist)

 
  return model

batch_size  = train_size
 
callback = tf.keras.callbacks.EarlyStopping(monitor='val_loss', patience=1830,restore_best_weights=True)
flipout_BNN = create_flipout_bnn_model(train_size=train_size)
flipout_BNN.compile(optimizer=Adam(learning_rate=0.002 ),jit_compile=True,
                  loss=NLL,metrics= [tf.keras.metrics.RootMeanSquaredError()]
                 ) 
flipout_BNN.summary()
history_flipout_BNN = flipout_BNN.fit(X_train, y_train, epochs=30000, verbose=0, batch_size=batch_size,validation_data=(X_val,y_val),callbacks=[callback] )

解决办法

  • 降低学习率:当前Adam的学习率0.002对小数据集和贝叶斯模型来说偏高,容易导致参数更新幅度过大,引发验证损失波动。建议尝试0.0005或0.0001这类更小的学习率。
  • 缩小模型规模:3层各30神经元的结构对于单特征小数据集过于复杂,易引发过拟合和不稳定。可减少层数(比如1-2层)或降低每层神经元数量(10-15个),简化模型后验证损失的波动会减少。
  • 调整KL散度权重:当前KL散度除以train_size,对于小数据集来说正则化强度不足。可以尝试去掉除以train_size的操作,或者乘以一个小常数(如0.1)增强正则化,约束参数分布的波动。
  • 修改批量大小:全批量训练(batch_size = train_size)虽然稳定,但梯度更新缺乏多样性。尝试用更小的批量(比如2-8,根据数据集实际大小调整),引入适度梯度噪声有助于模型泛化,缓解验证尖峰。
  • 优化早停设置:patience=1830过大,早停触发太晚,模型训练后期易出现不稳定。建议把patience调低到50-200,一旦验证损失连续多个epoch无下降就停止,及时保留最优权重。
  • 平滑输出尺度计算:输出层scale的计算用了0.05 * params,系数过小会让scale的变化过于敏感。尝试增大系数到0.5,让scale的更新更平缓,避免因尺度突变导致验证损失尖峰。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 06:10:39