Numpyro中AR(1)均值切换模型采样参数停滞问题排查求助
我要估计一个随隐状态$S \in {0,1}$切换均值的AR(1)过程$y$,隐状态$S$服从固定转移概率的马尔可夫过程(参考《State-Space Models with Regime Switching: Classical》),模型形式如下:
$y_t - mu_{0/1} = phi * (y_{t-1} - mu_{0/1})+ epsilon_t$
当$state_t=0$时使用$\mu_0$,$state_t=1$时使用$\mu_1$。我使用jax/numpyro的DiscreteHMCGibbs采样器(使用NUTS枚举隐状态得到相同结果),但采样器无法正常工作:所有超参数停滞在初始值,诊断结果显示所有标准差为0。以下是可复现该问题的最小示例代码,请问我的实现存在明显错误吗?
可复现代码
import jax.numpy as jnp import numpyro import numpyro.distributions as dist from numpyro.contrib.control_flow import scan from numpyro.infer import MCMC, NUTS,DiscreteHMCGibbs from jax import random, pure_callback import jax import numpy as np def generate_synthetic_data(T=100, mu=[0, 5], phi=0.5, sigma=1.0, p=np.array([[0.95, 0.05], [0.1, 0.9]])): states = np.zeros(T, dtype=np.int32) y = np.zeros(T) current_state = np.random.choice([0, 1], p=[0.5, 0.5]) states[0] = current_state y[0] = np.random.normal(mu[current_state], sigma) for t in range(1, T): current_state = np.random.choice([0, 1], p=p[current_state,:]) states[t] = current_state y[t] = np.random.normal(mu[current_state] + phi * (y[t-1] - mu[current_state]), sigma) return y, states def mean_switching_AR1_model(y): T = len(y) phi = numpyro.sample('phi', dist.Normal(0, 1)) sigma = numpyro.sample('sigma', dist.Exponential(1)) with numpyro.plate('state_plate', 2): mu = numpyro.sample('mu', dist.Normal(0, 5)) p = numpyro.sample('p', dist.Dirichlet(jnp.ones(2))) probs_init = numpyro.sample('probs_init', dist.Dirichlet(jnp.ones(2))) s_0 = numpyro.sample('s_0', dist.Categorical(probs_init)) def transition_fn(carry, y_t): prev_state = carry state_probs = p[prev_state] state = numpyro.sample('state', dist.Categorical(state_probs)) mu_state = mu[state] y_mean = mu_state + phi * (y_t - mu_state) y_next = numpyro.sample('y_next', dist.Normal(y_mean, sigma), obs=y_t) return state, (state, y_next) _ , (signal, y)=scan(transition_fn, s_0, y[:-1], length=T-1) return (signal, y) # Synthetic data generation T = 1000 mu_true = [0, 3] phi_true = 0.5 sigma_true = 0.25 transition_matrix_true = np.array([[0.95, 0.05], [0.1, 0.9]]) y, states_true = generate_synthetic_data(T, mu=mu_true, phi=phi_true, sigma=sigma_true, p=transition_matrix_true) rng_key = random.PRNGKey(0) nuts_kernel = NUTS(mean_switching_AR1_model) gibbs_kernel = DiscreteHMCGibbs(nuts_kernel, modified=True) # Run MCMC mcmc = MCMC(gibbs_kernel, num_samples=1000, num_warmup=1000) mcmc.run(rng_key, y=y) mcmc.print_summary()
核心错误分析与修正
你的代码存在几个关键问题,直接导致采样器无法有效更新参数:
1. 观测值时序匹配错误
原模型中,$y_t$的均值依赖前一时刻的y值和当前时刻的状态,但你的scan逻辑完全搞反了时序:
- 你传入
y[:-1](即$y_0$到$y_{T-2}$)作为遍历对象,却把每个y_t当作当前观测值,用它来计算下一时刻的均值,这导致似然计算完全错误,参数无法获得有效梯度。
修正方案:
调整scan的输入为y[1:](即$y_1$到$y_{T-1}$),同时让transition_fn的carry包含前一状态和前一时刻的y值,用它们来预测当前观测值:
def transition_fn(carry, y_t): prev_state, prev_y = carry state = numpyro.sample('state', dist.Categorical(p[prev_state])) mu_state = mu[state] # 用前一时刻y值计算当前时刻y的均值 y_mean = mu_state + phi * (prev_y - mu_state) numpyro.sample('y_obs', dist.Normal(y_mean, sigma), obs=y_t) return (state, y_t), state
同时修改scan的初始化和输入:
# 初始carry包含初始状态和第一个观测值,遍历剩余的观测值 _ , signal = scan(transition_fn, (s_0, y[0]), y[1:], length=T-1)
2. 模型返回值冗余
原代码返回(signal, y),但y是输入的观测值,对模型似然没有贡献,只需返回状态序列即可。
3. 先验设置可优化
AR(1)过程的$\phi$通常约束在$(-1,1)$区间内,你当前使用的Normal(0,1)虽然大部分概率落在该区间,但可以改用Uniform(-1,1)或Normal(0, 0.5)来进一步约束,避免采样到不稳定的$\phi$值。
修正后的完整模型代码
def mean_switching_AR1_model(y): T = len(y) # 更紧凑的phi先验 phi = numpyro.sample('phi', dist.Normal(0, 0.5)) sigma = numpyro.sample('sigma', dist.Exponential(1)) with numpyro.plate('state_plate', 2): mu = numpyro.sample('mu', dist.Normal(0, 5)) p = numpyro.sample('p', dist.Dirichlet(jnp.ones(2))) probs_init = numpyro.sample('probs_init', dist.Dirichlet(jnp.ones(2))) s_0 = numpyro.sample('s_0', dist.Categorical(probs_init)) def transition_fn(carry, y_t): prev_state, prev_y = carry state = numpyro.sample('state', dist.Categorical(p[prev_state])) mu_state = mu[state] y_mean = mu_state + phi * (prev_y - mu_state) numpyro.sample('y_obs', dist.Normal(y_mean, sigma), obs=y_t) return (state, y_t), state _ , signal = scan(transition_fn, (s_0, y[0]), y[1:], length=T-1) return signal
采样器调整建议
- 尝试移除
DiscreteHMCGibbs的modified=True参数,或先用numpyro.infer.SVI配合离散枚举完成参数初始化,再进行MCMC采样。 - 将warmup步数增加到2000,确保采样器充分收敛。
内容的提问来源于stack exchange,提问作者Thanatopseustes

