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

带生存分析的贝叶斯线性回归表现劣于标准模型的原因排查

贝叶斯线性回归模型性能异常排查

我开发了一款考虑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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 16:52:02