如何用TensorFlow Probability学习Beta分布参数?损失出现NaN问题
解决Beta分布参数拟合时损失为NaN的问题
问题根源
- 参数无约束:Beta分布的
concentration1(alpha)和concentration0(beta)必须为正数,但原代码直接用普通tf.Variable初始化,优化过程中参数可能被更新为非正数,导致log_prob计算出现NaN。 - 样本边界值:Scipy的
beta.rvs可能生成极接近0或1的浮点数,当参数较小时,log_prob计算会出现无穷大,进而导致损失NaN。 - 优化策略不合适:SGD学习率设置偏小,且未针对参数特性做优化调整。
修复后的代码
from scipy.stats import beta import tensorflow as tf import tensorflow_probability as tfp tfd = tfp.distributions # 生成样本并过滤边界值,避免log_prob计算异常 beta_sample_data = beta.rvs(5, 5, size=1000) beta_sample_data = beta_sample_data[(beta_sample_data > 1e-8) & (beta_sample_data < 1 - 1e-8)] # 用Softplus变换确保参数始终为正 def positive_variable(initial_value, name): return tfp.util.TransformedVariable( initial_value=initial_value, bijector=tfp.bijectors.Softplus(), name=name ) beta_train = tfd.Beta( concentration1=positive_variable(1., name='alpha'), concentration0=positive_variable(1., name='beta'), name='beta_train' ) def nll(x_train, distribution): return -tf.reduce_mean(distribution.log_prob(x_train)) @tf.function def get_loss_and_grads(x_train, distribution): with tf.GradientTape() as tape: loss = nll(x_train, distribution) grads = tape.gradient(loss, distribution.trainable_variables) return loss, grads def beta_dist_optimisation(data, distribution): train_loss_results = [] train_rate_results = [] # 改用Adam优化器,提升拟合效率 optimizer = tf.keras.optimizers.Adam(learning_rate=0.1) num_steps = 1000 # 增加迭代步数确保收敛 for i in range(num_steps): loss, grads = get_loss_and_grads(data, distribution) optimizer.apply_gradients(zip(grads, distribution.trainable_variables)) # 获取变换后的实际参数值 alpha_value = distribution.concentration1().numpy() beta_value = distribution.concentration0().numpy() train_loss_results.append(loss.numpy()) train_rate_results.append((alpha_value, beta_value)) if i % 100 == 0: print(f"Step {i:04d}: Loss: {loss:.3f}, Alpha: {alpha_value:.3f}, Beta: {beta_value:.3f}") return train_loss_results, train_rate_results sample_data = tf.cast(beta_sample_data, tf.float32) train_loss_results, train_rate_results = beta_dist_optimisation(sample_data, beta_train)
关键修改说明
- 参数约束:使用
TransformedVariable结合Softplus变换,强制alpha和beta始终为正,从根源避免参数非法导致的NaN。 - 样本预处理:过滤掉极接近0和1的样本,消除边界值对
log_prob计算的干扰。 - 优化器调整:替换为Adam优化器并调优学习率,同时增加迭代步数,提升参数收敛效果。
- 变量取值修正:
TransformedVariable需通过调用()获取变换后的实际参数值,而非原代码的.value()。
内容的提问来源于stack exchange,提问作者AI92
相关产品推荐
相关产品推荐

