如何在Numpyro中高效实现GARCH模型?自研GARCH-M模型优化求助
在Numpyro中优化GARCH-M模型实现的建议
我想了解如何在Numpyro中最佳实现GARCH模型。曾查阅Numpyro官方时间序列预测教程,但内容较为晦涩(模型符号与变量名称难以对应,对入门者来说模型过于复杂)。我编写了如下代码用于估计GARCH-M模型,但性能表现似乎不佳,希望能得到改进建议。
所使用的GARCH-M模型公式如下:
选择GARCH-M模型是因为它需要循环操作,无法一次性推断所有ε序列,必须推断时变波动率。
我的实现代码:
# preliminaries import numpyro import numpyro.distributions as dist import jax.numpy as jnp from numpyro.infer import MCMC, NUTS from numpyro.contrib.control_flow import scan import numpy as np def my_model(y=None): μ = numpyro.sample("μ", dist.Normal(0, 4)) ω = numpyro.sample("ω", dist.HalfCauchy(2)) α = numpyro.sample("α", dist.Uniform()) β = numpyro.sample("β", dist.Uniform()) λ = numpyro.sample("λ", dist.Normal(0, 4)) σ2_0 = ω / (1 - α - β) def gjr_var(state, new): # state is past variance and past shock σ2_t, r_t = state ɛ_t = (r_t - μ - λ*σ2_t**0.5)/σ2_t**0.5 σ2_tp = ω + α * ɛ_t**2 + β * σ2_t r_tp = numpyro.sample('r', dist.Normal(μ + λ*σ2_tp**0.5, σ2_tp**0.5), obs=new) return (σ2_tp, r_tp), r_tp _, r = scan(gjr_var, (σ2_0, y[0]), y[1:]) return r
核心改进建议
1. 参数约束与先验优化
- 平稳性约束:GARCH模型要求α + β < 1(保证方差过程平稳),当前Uniform先验无法保证这一点,容易出现α+β≥1导致的数值爆炸。建议通过变量变换+约束实现:
# 用softplus变换确保α、β为正 α_raw = numpyro.sample("α_raw", dist.Normal(0, 1)) α = jax.nn.softplus(α_raw) β_raw = numpyro.sample("β_raw", dist.Normal(0, 1)) β = jax.nn.softplus(β_raw) # 添加平稳性约束,惩罚α+β≥0.99的情况(留余量避免边界数值问题) numpyro.factor("stationarity", jnp.where(α + β < 0.99, 0, -jnp.inf)) - ω的先验:HalfCauchy(2)可能过于宽泛,建议根据数据的方差尺度调整,比如用
dist.HalfNormal(jnp.std(y)**2 if y is not None else 1),让先验贴合数据特征。 - 初始方差σ2_0:原代码用
ω/(1-α-β)初始化,当α+β接近1时会出现数值无穷大。建议将σ2_0作为独立参数采样,比如σ2_0 = numpyro.sample("σ2_0", dist.HalfNormal(jnp.std(y)**2 if y is not None else 1)),避免依赖α和β的组合。
2. 状态更新逻辑修正
- 残差计算:添加极小值
1e-6避免除以0的情况:ɛ_t = (r_t - μ - λ*jnp.sqrt(σ2_t)) / jnp.sqrt(σ2_t + 1e-6) - 扫描循环的观测处理:原代码中
numpyro.sample('r', ...)的名称固定为'r',Numpyro无法区分不同时间步的观测。调整扫描逻辑,直接将整个y序列传入,从第一个时间步开始处理,避免重复命名问题。
3. 推断效率提升
- NUTS采样配置:设置
target_accept_prob=0.9,提高带约束参数的采样稳定性:nuts_kernel = NUTS(improved_garch_m, target_accept_prob=0.9) mcmc = MCMC(nuts_kernel, num_warmup=1000, num_samples=2000) - 参数变换:用softplus将无约束变量转换为正参数,让NUTS的采样空间更规整,减少采样时的数值问题。
4. 模型验证与诊断
- 采样完成后,用
numpyro.diagnostics.summary(mcmc.get_samples())查看参数的R-hat值,确保所有参数的R-hat<1.01,确认收敛。 - 绘制参数的迹图(比如用
arviz.plot_trace),检查采样链的稳定性。 - 生成后验预测样本,对比观测值与预测值,验证模型拟合效果。
优化后的示例代码
# preliminaries import numpyro import numpyro.distributions as dist import jax.numpy as jnp import jax from numpyro.infer import MCMC, NUTS from numpyro.contrib.control_flow import scan import numpy as np def improved_garch_m(y=None): μ = numpyro.sample("μ", dist.Normal(0, 4)) # 根据数据尺度调整ω的先验 y_scale = jnp.std(y)**2 if y is not None else 1.0 ω = numpyro.sample("ω", dist.HalfNormal(y_scale)) # 变量变换保证α、β为正,添加平稳性约束 α_raw = numpyro.sample("α_raw", dist.Normal(0, 1)) α = jax.nn.softplus(α_raw) β_raw = numpyro.sample("β_raw", dist.Normal(0, 1)) β = jax.nn.softplus(β_raw) numpyro.factor("stationarity", jnp.where(α + β < 0.99, 0.0, -jnp.inf)) λ = numpyro.sample("λ", dist.Normal(0, 4)) # 独立采样初始方差,避免数值问题 σ2_0 = numpyro.sample("σ2_0", dist.HalfNormal(y_scale)) def step(state, obs_r): σ2_t = state # 计算残差,添加小epsilon避免除以0 ɛ_t = (obs_r - μ - λ * jnp.sqrt(σ2_t)) / jnp.sqrt(σ2_t + 1e-6) # 更新方差 σ2_tp = ω + α * ɛ_t**2 + β * σ2_t # 采样当前观测(或记录观测) numpyro.sample("r", dist.Normal(μ + λ * jnp.sqrt(σ2_tp), jnp.sqrt(σ2_tp)), obs=obs_r) return σ2_tp, None # 扫描整个观测序列 scan(step, σ2_0, y)
内容的提问来源于stack exchange,提问作者Shffl
相关产品推荐
相关产品推荐

