使用PyMC3拟合分段函数遇语法问题求助
解决PyMC3拟合分段函数的语法错误
- 错误根源:PyMC3的张量变量无法直接与NumPy数组用
>/<等原生运算符做元素级比较,必须用PyMC3内置的pm.math模块函数执行张量兼容的操作。 - 额外注意:离散变量
switch_point不能用默认的NUTS采样器,需指定Metropolis采样器。
修正后的完整代码:
import pymc3 as pm import numpy as np x = np.linspace(0, 10, 100) y = np.piecewise(x, [x < 5, x >=5], [lambda x: 2*x + 1, lambda x: -3*x + 26]) + np.random.normal(0, 1, 100) with pm.Model() as model: alpha = pm.Normal('alpha', mu=0, sd=10) beta1 = pm.Normal('beta1', mu=0, sd=10) beta2 = pm.Normal('beta2', mu=0, sd=10) switch_point = pm.DiscreteUniform('switch_point', lower=0, upper=10) # 用pm.math.lt实现张量兼容的元素级小于比较,替代原生< mu = pm.math.switch(pm.math.lt(x, switch_point), alpha + beta1*x, alpha + beta2*x) sigma = pm.HalfNormal('sigma', sd=1) likelihood = pm.Normal('likelihood', mu=mu, sd=sigma, observed=y) # 为离散变量指定Metropolis采样器 trace = pm.sample(2000, tune=1000, step=[pm.Metropolis(vars=[switch_point])])
关键修改说明:
- 替换比较逻辑:将
switch_point > x改为pm.math.lt(x, switch_point),确保比较操作在PyMC3的张量计算图中正常运行。 - 指定采样器:通过
step参数为离散的切换点变量配置Metropolis采样器,解决NUTS采样器不支持离散变量的问题。
内容的提问来源于stack exchange,提问作者J.gra
相关产品推荐
相关产品推荐

