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

如何在Numpyro中高效实现GARCH模型?自研GARCH-M模型优化求助

在Numpyro中优化GARCH-M模型实现的建议

我想了解如何在Numpyro中最佳实现GARCH模型。曾查阅Numpyro官方时间序列预测教程,但内容较为晦涩(模型符号与变量名称难以对应,对入门者来说模型过于复杂)。我编写了如下代码用于估计GARCH-M模型,但性能表现似乎不佳,希望能得到改进建议。

所使用的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 01:15:45