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

PyMC贝叶斯线性模型样本外数据预测报错求助

贝叶斯线性回归PyMC样本外预测形状不匹配问题解决

问题场景

用PyMC实现贝叶斯线性回归模型,训练完成后想在测试集上验证模型性能,但调用pm.set_data切换数据集时触发形状不匹配错误,复现代码及错误信息如下:

复现代码

import numpy as np
import pymc as pm
import arviz as az
import matplotlib.pyplot as plt

def run_model():
    # Generate synthetic data
    np.random.seed(42)
    x = np.linspace(0, 10, 100)
    a_true = 2.5  # True slope
    b_true = 1.0  # True intercept
    y_true = a_true * x + b_true
    y = y_true + np.random.normal(0, 1, size=x.size)  # Add some noise

    # Split into training and test sets
    x_train, x_test = x[:80], x[80:]
    y_train, y_test = y[:80], y[80:]

    # Define and fit the model
    with pm.Model() as linear_model:
        # Define x as a pm.Data variable to allow updating with pm.set_data
        x_shared = pm.Data("x", x_train)

        # Priors for slope and intercept
        a = pm.Normal("a", mu=0, sigma=10)
        b = pm.Normal("b", mu=0, sigma=10)
        sigma = pm.HalfNormal("sigma", sigma=1)

        # Expected value of y
        mu = a * x_shared + b

        # Likelihood
        y_obs = pm.Normal("y_obs", mu=mu, sigma=sigma, observed=y_train)

        # Sample from the posterior
        trace = pm.sample(1000, tune=1000, return_inferencedata=True, chains=1)

    # Predict on training data
    with linear_model:
        pm.set_data({"x": x_train})  # Update data to training
        post_pred_train = pm.sample_posterior_predictive(trace)

    # Predict on test data
    with linear_model:
        pm.set_data({"x": x_test})  # Update data to testing
        post_pred_test = pm.sample_posterior_predictive(trace)

    # Plot results
    plt.figure(figsize=(10, 5))

    # Plot training data
    plt.scatter(x_train, y_train, c="blue", label="Training data")
    plt.plot(x_train, y_true[:80], "k--", label="True function")

    # Plot posterior predictive for training data
    plt.plot(
        x_train,
        post_pred_train["y_obs"].mean(axis=0),
        label="Posterior predictive (train)",
        color="red",
    )

    # Plot test data
    plt.scatter(x_test, y_test, c="green", label="Test data")

    # Plot posterior predictive for test data
    plt.plot(
        x_test,
        post_pred_test["y_obs"].mean(axis=0),
        label="Posterior predictive (test)",
        color="orange",
    )

    plt.legend()
    plt.xlabel("x")
    plt.ylabel("y")
    plt.title("Bayesian Linear Regression with PyMC")
    plt.show()

    # Summary of the model parameters
    print(az.summary(trace, var_names=["a", "b", "sigma"]))


# Only execute if run as the main module
if __name__ == '__main__':
    run_model()

错误信息

ValueError: shape mismatch: objects cannot be broadcast to a single shape.  Mismatch is between arg 0 with shape (80,) and arg 1 with shape (20,).
Apply node that caused the error: normal_rv{"(),()->()"}(RNG(<Generator(PCG64) at 0x1F23323B5A0>), [80], Composite{((i0 * i1) + i2)}.0, ExpandDims{axis=0}.0)
Toposort index: 4
Inputs types: [RandomGeneratorType, TensorType(int64, shape=(1,)), TensorType(float64, shape=(None,)), TensorType(float64, shape=(1,))]
Inputs shapes: ['No shapes', (1,), (20,), (1,)]
Inputs strides: ['No strides', (8,), (8,), (8,)]
Inputs values: [Generator(PCG64) at 0x1F23323B5A0, array([80], dtype=int64), 'not shown', array([0.97974278])]
Outputs clients: [[output[1](normal_rv{"(),()->()"}.0)], [output[0](y_obs)]]

错误原因

模型中仅将输入x定义为pm.Data可切换变量,但观测变量y_obs绑定了训练集的形状(80,)。当调用pm.set_data将x切换为测试集(20,)时,预测的均值mu形状变为(20,),但y_obs的观测数据仍保留原始的(80,)形状,导致sample_posterior_predictive采样时出现形状不匹配冲突。

解决方法

将观测变量y同样定义为pm.Data变量,切换数据集时同步更新y的形状(预测时可传入任意同测试集形状的数组,比如全NaN或零数组,因为后验预测采样不需要真实观测值)。

修改后的完整代码

import numpy as np
import pymc as pm
import arviz as az
import matplotlib.pyplot as plt

def run_model():
    # Generate synthetic data
    np.random.seed(42)
    x = np.linspace(0, 10, 100)
    a_true = 2.5  # True slope
    b_true = 1.0  # True intercept
    y_true = a_true * x + b_true
    y = y_true + np.random.normal(0, 1, size=x.size)  # Add some noise

    # Split into training and test sets
    x_train, x_test = x[:80], x[80:]
    y_train, y_test = y[:80], y[80:]

    # Define and fit the model
    with pm.Model() as linear_model:
        # 将x和y都定义为pm.Data变量,支持后续切换
        x_shared = pm.Data("x", x_train)
        y_shared = pm.Data("y", y_train)

        # Priors for slope and intercept
        a = pm.Normal("a", mu=0, sigma=10)
        b = pm.Normal("b", mu=0, sigma=10)
        sigma = pm.HalfNormal("sigma", sigma=1)

        # Expected value of y
        mu = a * x_shared + b

        # Likelihood
        y_obs = pm.Normal("y_obs", mu=mu, sigma=sigma, observed=y_shared)

        # Sample from the posterior
        trace = pm.sample(1000, tune=1000, return_inferencedata=True, chains=1)

    # Predict on training data
    with linear_model:
        pm.set_data({"x": x_train, "y": y_train})
        post_pred_train = pm.sample_posterior_predictive(trace)

    # Predict on test data
    with linear_model:
        # 传入和测试集同形状的y,这里用零数组即可,不影响预测结果
        pm.set_data({"x": x_test, "y": np.zeros_like(x_test)})
        post_pred_test = pm.sample_posterior_predictive(trace)

    # Plot results
    plt.figure(figsize=(10, 5))

    # Plot training data
    plt.scatter(x_train, y_train, c="blue", label="Training data")
    plt.plot(x_train, y_true[:80], "k--", label="True function")

    # Plot posterior predictive for training data
    plt.plot(
        x_train,
        post_pred_train["y_obs"].mean(axis=0),
        label="Posterior predictive (train)",
        color="red",
    )

    # Plot test data
    plt.scatter(x_test, y_test, c="green", label="Test data")

    # Plot posterior predictive for test data
    plt.plot(
        x_test,
        post_pred_test["y_obs"].mean(axis=0),
        label="Posterior predictive (test)",
        color="orange",
    )

    plt.legend()
    plt.xlabel("x")
    plt.ylabel("y")
    plt.title("Bayesian Linear Regression with PyMC")
    plt.show()

    # Summary of the model parameters
    print(az.summary(trace, var_names=["a", "b", "sigma"]))


# Only execute if run as the main module
if __name__ == '__main__':
    run_model()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 15:49:51