带生存分析的贝叶斯线性回归表现劣于标准模型的原因排查
贝叶斯线性回归模型性能异常排查
我开发了一款考虑x与y双变量非对称不确定性及截尾数据(上下限)的贝叶斯线性回归函数,采用emcee MCMC包估计斜率与截距的后验分布。但在模拟数据测试中,该模型表现反而不如scipy.stats.linregress的标准线性回归或未考虑不确定性与截尾的基础emcee贝叶斯线性回归。
模型代码
def bayesian_linear_regression_survival_analysis(x, y, xerr_low, xerr_high, yerr_low, yerr_high, m_guess, d_guess, nwalkers=8, nsteps=int(0.5e4)): """ Perform Bayesian linear regression considering asymmetric uncertainties and censored data in both x and y. """ def model(theta, x): slope, intercept = theta return slope * x + intercept def survival_function(observed_value, model_value, sigma): """ Survival function for censored data. Returns the probability of observing data above (for lower limits) or below (for upper limits) the given value. """ return stats.norm.sf(observed_value, loc=model_value, scale=sigma) def estimate_sigma(limit_value, model_value, fixed_sigma=1e-6): # Calculate the sigma as the absolute difference between the limit value and the model value sigma = np.abs(limit_value - model_value) # Ensure that the sigma is not an array with ambiguous truth values sigma = np.where(sigma > 0, sigma, fixed_sigma) return sigma def log_likelihood(theta, x, y, xerr_low, xerr_high, yerr_low, yerr_high): slope, intercept = theta model_line = model(theta, x) sigma_x = estimate_sigma(x, model_line) sigma_y = estimate_sigma(y, model_line) # Weight sigma_x and sigma_y based on the relative contributions of the lower and upper errors weight_x = np.where(x > model_line, xerr_high / (xerr_high + xerr_low), xerr_low / (xerr_high + xerr_low)) weighted_sigma_x = weight_x * sigma_x weight_y = np.where(y > model_line, yerr_high / (yerr_high + yerr_low), yerr_low / (yerr_high + yerr_low)) weighted_sigma_y = weight_y * sigma_y # Calculate the total variance, including the contribution from sigma_x and sigma_y point_variance = weighted_sigma_y**2 + (slope * weighted_sigma_x)**2 # Calculate the residuals between observed data and the model prediction residuals = y - model_line # Calculate weights based on how much the errors contribute to the point variance weight = (weighted_sigma_y**2 + (slope * weighted_sigma_x)**2) / point_variance # Calculate the weighted log-likelihood log_likelihood = -0.5 * np.sum(weight * ((residuals)**2 / point_variance + np.log(point_variance))) # Apply survival function for censored data (upper/lower limits) for i in range(len(x)): if xerr_low[i] == 0: # Lower limit on x sigma_x_i = estimate_sigma(x[i], model_line[i]) log_likelihood += np.log(survival_function(x[i], model_line[i], sigma_x_i) + 1e-10) if xerr_high[i] == 0: # Upper limit on x sigma_x_i = estimate_sigma(x[i], model_line[i]) log_likelihood += np.log(survival_function(x[i], model_line[i], sigma_x_i) + 1e-10) if yerr_low[i] == 0: # Lower limit on y sigma_y_i = estimate_sigma(y[i], model_line[i]) log_likelihood += np.log(survival_function(y[i], model_line[i], sigma_y_i) + 1e-10) if yerr_high[i] == 0: # Upper limit on y sigma_y_i = estimate_sigma(y[i], model_line[i]) log_likelihood += np.log(survival_function(y[i], model_line[i], sigma_y_i) + 1e-10) return log_likelihood def log_prior(theta): slope, intercept = theta if -10.0 < slope < 10.0 and -10.0 < intercept < 10.0: return 0.0 return -np.inf def log_probability(theta, x, y, xerr_low, xerr_high, yerr_low, yerr_high): lp = log_prior(theta) if not np.isfinite(lp): return -np.inf return lp + log_likelihood(theta, x, y, xerr_low, xerr_high, yerr_low, yerr_high) initial = np.array([m_guess, d_guess]) ndim = 2 sampler = emcee.EnsembleSampler(nwalkers, ndim, log_probability, args=(x, y, xerr_low, xerr_high, yerr_low, yerr_high)) pos = initial + 1e-4 * np.random.randn(nwalkers, ndim) sampler.run_mcmc(pos, nsteps, progress=True) samples = sampler.get_chain(discard=100, thin=15, flat=True) return samples, sampler
测试代码
nwalkers, nsteps = 8, int(1e4) # Initialize random data # np.random.seed(42) N = 20 x = np.random.uniform(0, 10, N) true_slope = 2 true_intercept = 5 y = true_slope * x + true_intercept + np.random.normal(0, 3, N) # True values for comparison x_truth = np.arange(-1, 12, 0.2) y_truth = x_truth * true_slope + true_intercept # Initialize errors yerr_low = np.random.uniform(0.5, 3.0, N) yerr_upp = np.random.uniform(0.5, 3.0, N) xerr_low = np.random.uniform(0.1, 3.0, N) xerr_upp = np.random.uniform(0.1, 3.0, N) large_error_threshold_x = 2.5 large_error_threshold_y = 4.0 # Identify indices where the errors exceed the threshold upper_limit_indices_x = np.where(xerr_upp > large_error_threshold_x)[0] lower_limit_indices_x = np.where(xerr_low > large_error_threshold_x)[0] upper_limit_indices_y = np.where(yerr_upp > large_error_threshold_y)[0] lower_limit_indices_y = np.where(yerr_low > large_error_threshold_y)[0] # Remove indices from lower limits that are already in upper limits lower_limit_indices_x = np.setdiff1d(lower_limit_indices_x, upper_limit_indices_x) lower_limit_indices_y = np.setdiff1d(lower_limit_indices_y, upper_limit_indices_y) # Apply the upper and lower limits y[upper_limit_indices_y] += yerr_upp[upper_limit_indices_y] # Shifted upward for upper limits y[lower_limit_indices_y] -= yerr_low[lower_limit_indices_y] # Shifted downward for lower limits x[upper_limit_indices_x] += xerr_upp[upper_limit_indices_x] x[lower_limit_indices_x] -= xerr_low[lower_limit_indices_x] # Set yerr_low to 0 for upper limits and yerr_upp to 0 for lower limits in y, same for x yerr_low[upper_limit_indices_y] = 0 yerr_upp[lower_limit_indices_y] = 0 xerr_low[upper_limit_indices_x] = 0 xerr_upp[lower_limit_indices_x] = 0 # Perform standard linear regression standard_slope, standard_intercept, _, _, _ = stats.linregress(x, y) # Perform standard Bayesian linear regression using emcee def log_likelihood(theta, x, y): slope, intercept = theta model_line = slope * x + intercept sigma2 = np.var(y) return -0.5 * np.sum((y - model_line) ** 2 / sigma2 + np.log(sigma2)) def log_prior(theta): slope, intercept = theta if -10.0 < slope < 10.0 and -10.0 < intercept < 10.0: return 0.0 return -np.inf def log_probability(theta, x, y): lp = log_prior(theta) if not np.isfinite(lp): return -np.inf return lp + log_likelihood(theta, x, y) initial = np.array([standard_slope, standard_intercept]) ndim = 2 sampler = emcee.EnsembleSampler(nwalkers, ndim, log_probability, args=(x, y)) pos = initial + 1e-4 * np.random.randn(nwalkers, ndim) sampler.run_mcmc(pos, nsteps, progress=True) emcee_samples = sampler.get_chain(discard=100, thin=15, flat=True) emcee_median_slope = np.median(emcee_samples[:, 0]) emcee_median_intercept = np.median(emcee_samples[:, 1]) emcee_slope_percentiles = np.percentile(emcee_samples[:, 0], [16, 84]) emcee_intercept_percentiles = np.percentile(emcee_samples[:, 1], [16, 84]) # Perform Bayesian linear regression with survival analysis samples, sampler = bayesian_linear_regression_survival_analysis( x, y, xerr_low, xerr_upp, yerr_low, yerr_upp, standard_slope, standard_intercept, nwalkers, nsteps) survival_median_slope = np.median(samples[:, 0]) survival_median_intercept = np.median(samples[:, 1]) survival_slope_percentiles = np.percentile(samples[:, 0], [16, 84]) survival_intercept_percentiles = np.percentile(samples[:, 1], [16, 84]) # Plot the data and the regression results fig, ax = plt.subplots(figsize=(10, 6)) # Truth ax.plot(x_truth, y_truth, color='k', label=f"Truth: $y = ({true_slope:.2f})x + ({true_intercept:.2f})$", ls='--') # Plot all data points ax.errorbar(x, y, xerr=[xerr_low, xerr_upp], yerr=[yerr_low, yerr_upp], fmt='o', mfc='silver', color='k', ecolor='silver', elinewidth=1) # Plot the standard linear regression line y_plot_standard = standard_slope * x_truth + standard_intercept ax.plot(x_truth, y_plot_standard, label=f"Lin. Regression Fit: $y = ({standard_slope:.2f})x + ({standard_intercept:.2f})$", color='green') # Plot the emcee fit line y_plot_emcee = emcee_median_slope * x_truth + emcee_median_intercept ax.plot(x_truth, y_plot_emcee, label=f"Emcee Fit: $y = ({emcee_median_slope:.2f})x + ({emcee_median_intercept:.2f})$", color='blue') # Plot the Bayesian linear regression with survival analysis y_plot_survival = survival_median_slope * x_truth + survival_median_intercept ax.plot(x_truth, y_plot_survival, color='red', label=f"Survival Analysis Fit: $y = ({survival_median_slope:.2f})x + ({survival_median_intercept:.2f})$") # Fill between 1-sigma for survival analysis y_lower = survival_slope_percentiles[0] * x_truth + survival_intercept_percentiles[0] y_upper = survival_slope_percentiles[1] * x_truth + survival_intercept_percentiles[1] # ax.fill_between(x_truth, y_lower, y_upper, color='tab:red', alpha=0.3, label='1-sigma region (Survival)') ax.set_xlabel("x") ax.set_ylabel("y") ax.legend() # Plot the corner plot for the MCMC samples fig2 = corner.corner(samples, labels=["Slope", "Intercept"], truths=[survival_median_slope, survival_median_intercept], quantiles=[0.16, 0.5, 0.84], show_titles=True, title_fmt=".2f", title_kwargs={"fontsize": 12})
测试结果

问题求助
理论上该模型应更适配含不确定性与截尾的数据,但实际表现反而不如简单模型。请问我的对数似然计算、不确定性处理逻辑、截尾数据的生存函数应用是否存在问题?恳请专业建议与见解。
内容的提问来源于stack exchange,提问作者Alessandro Peca
相关产品推荐
相关产品推荐

