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

Numpyro中AR(1)均值切换模型采样参数停滞问题排查求助

均值切换AR(1)模型numpyro采样停滞问题排查

我要估计一个随隐状态$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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 04:17:14