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

在Numpyro中处理不同长度时间序列的层次时间序列建模问题

在Numpyro中处理不同长度时间序列的层次时间序列建模问题

我明白你现在遇到的问题——你已经能成功拟合所有时间序列长度一致的层次随机游走模型,但当尝试用掩码(mask)处理不同长度/含缺失值的时间序列时,模型跑不通。咱们一步步来解决这个问题,先从回顾你的工作代码开始,再修正掩码部分的实现。


一、回顾:等长时间序列的层次随机游走模型

首先整理你已经跑通的等长数据生成和模型代码,确保逻辑清晰:

1. 生成等长模拟数据

import numpy as np
import matplotlib.pyplot as plt
import jax
from jax import random
import jax.numpy as jnp
import numpyro as ny
from numpyro.contrib.control_flow import scan
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS, Predictive

ny.set_host_device_count(4)

# 生成数据
np.random.seed(42)
T = 800  # 训练时间点
h = 200  # 预测时间点
N = 10   # 序列数量

# 全局超参数
mu_alpha = -0.15
log_mu_sigma = np.log(5)
mu_beta = 0.5

# 生成每个序列的局部参数
alpha = np.random.normal(mu_alpha, 0.35, size=N)
sigma = np.random.gamma(100, np.exp(log_mu_sigma)/100, size=N)
beta = np.random.normal(mu_beta, 0.25, size=N)

# 生成带漂移和趋势的随机游走
trend = beta[:,None] * np.arange(T+h)/365.25
rw_with_drift = np.random.normal(alpha[:,None]+trend, sigma[:,None]).cumsum(-1)
y = jnp.array(rw_with_drift)

# 可视化
fig, ax = plt.subplots()
ax.plot(np.arange(T+h), y.T, alpha=0.25)
ax.set(xlabel='Time', ylabel='Y', title='RW with Drift and Trend')
ax.axvline(T, color='k', ls='--', label='Training Period Ends')
ax.legend()
plt.show()

2. 等长序列的层次模型

def hierarchical_rw_model(y=None, future=0):
    N = 0 if y is None else y.shape[0]
    T = 0 if y is None else y.shape[1]
    # 初始水平:取每个序列的第一个时间点
    level_init = 0 if y is None else y[:,0]

    # 全局超参数
    mu_alpha = ny.sample("mu_alpha", dist.Normal(0,1))
    sig_alpha = ny.sample("sig_alpha", dist.Exponential(2.5))
    mu_beta = ny.sample("mu_beta", dist.Normal(0,1))
    sig_beta = ny.sample("sig_beta", dist.Exponential(1))
    mu_sigma = ny.sample("mu_sigma", dist.Normal(0,1))

    # 转置数据以便scan处理(scan默认沿第一个轴迭代)
    test = y[:,1:].T if (y is not None and T>1) else None

    # 每个序列的局部参数
    with ny.plate("time_series", N):
        alpha = ny.sample("alpha", dist.Normal(mu_alpha, sig_alpha))
        beta = ny.sample("beta", dist.Normal(mu_beta, sig_beta))
        sigma = ny.sample("sigma", dist.HalfNormal(jnp.exp(mu_sigma)))

    def transition_fn(y_t_minus_1, t):
        # 预测均值:当前水平 + 漂移项 + 趋势项
        exp_val = y_t_minus_1 + alpha + beta * t / 365.25
        # 观测采样
        y_ = ny.sample("y", dist.Normal(exp_val, sigma), obs=test[t] if test is not None else None)
        # 更新下一个时间点的水平
        return y_, y_

    # 用scan迭代时间步
    if T + future > 0:
        _, ys = scan(
            f=transition_fn,
            init=level_init,
            xs=jnp.arange(1, T + future)
        )
        # 保存预测结果
        if future > 0:
            ny.deterministic("y_forecast", ys[-future:])

# 拟合模型
kernel = NUTS(hierarchical_rw_model)
mcmc = MCMC(kernel, num_warmup=500, num_samples=1000, num_chains=2)
mcmc.run(random.PRNGKey(0), y=y[:,:T])

二、处理不同长度/含缺失值的时间序列

你的问题核心在于掩码的正确传递和初始值的处理:

  1. 原始代码中直接用y[:,0]作为初始水平,但如果序列的前几个时间点是缺失的(nan),这个初始值会无效;
  2. 在scan中使用掩码时,需要针对每个时间点的每个序列单独应用掩码,而不是全局掩码;
  3. ny.handlers.mask适合全局掩码场景,而你的场景需要逐时间步、逐序列的掩码控制。

1. 生成含缺失值的不同长度序列

我整理了你的数据生成代码,优化了可视化(只画每个序列的有效部分):

np.random.seed(42)
T = 800
h = 200
N = 10
mu_alpha = -0.15
log_mu_sigma = np.log(5)
mu_beta = 0.5

alpha = np.random.normal(mu_alpha, 0.35, size=N)
sigma = np.random.gamma(100, np.exp(log_mu_sigma)/100, size=N)
beta = np.random.normal(mu_beta, 0.25, size=N)

trend = beta[:,None] * np.arange(T+h)/365.25
rw_with_drift = np.random.normal(alpha[:,None]+trend, sigma[:,None]).cumsum(-1)

# 生成掩码:每个序列前eliminate_vals[i]个时间点设为缺失
eliminate_vals = np.random.choice(np.arange(101), N)
for i in range(N):
    rw_with_drift[i,:eliminate_vals[i]] = np.nan

# 转换为jax数组和掩码(True表示有效数据)
y = jnp.array(rw_with_drift)
mask = jnp.array(~np.isnan(rw_with_drift))

# 可视化:只绘制每个序列的有效部分
fig, ax = plt.subplots()
for i in range(N):
    valid_times = np.where(mask[i])[0]
    ax.plot(valid_times, y[i, valid_times], alpha=0.25)
ax.set(xlabel='Time', ylabel='Y', title='RW with Drift (Variable Length)')
ax.axvline(T, color='k', ls='--', label='Training Period Ends')
ax.legend()
plt.show()

2. 修正后的层次模型(支持掩码/不同长度)

重点修改了初始值处理、掩码传递和水平更新逻辑:

def hierarchical_rw_masked_model(y=None, mask=None, future=0):
    N = 0 if y is None else y.shape[0]
    T = 0 if y is None else y.shape[1]

    # --- 处理初始水平:取每个序列的第一个有效时间点的值 ---
    if y is not None and mask is not None:
        # 找到每个序列第一个有效时间点的索引
        first_valid_idx = jnp.argmax(mask, axis=1)
        # 提取对应的初始值
        level_init = y[jnp.arange(N), first_valid_idx]
    else:
        level_init = jnp.zeros(N) if y is not None else 0

    # 全局超参数(和之前一致)
    mu_alpha = ny.sample("mu_alpha", dist.Normal(0,1))
    sig_alpha = ny.sample("sig_alpha", dist.Exponential(2.5))
    mu_beta = ny.sample("mu_beta", dist.Normal(0,1))
    sig_beta = ny.sample("sig_beta", dist.Exponential(1))
    mu_sigma = ny.sample("mu_sigma", dist.Normal(0,1))

    # 转置数据和掩码,方便scan按时间步迭代
    test_y = y[:,1:].T if (y is not None and T>1) else None
    test_mask = mask[:,1:].T if (mask is not None and T>1) else None

    # 每个序列的局部参数
    with ny.plate("time_series", N):
        alpha = ny.sample("alpha", dist.Normal(mu_alpha, sig_alpha))
        beta = ny.sample("beta", dist.Normal(mu_beta, sig_beta))
        sigma = ny.sample("sigma", dist.HalfNormal(jnp.exp(mu_sigma)))

    def transition_fn(y_t_minus_1, t_and_mask):
        t, current_mask = t_and_mask
        # 预测均值
        exp_val = y_t_minus_1 + alpha + beta * t / 365.25

        # --- 关键:只对mask为True的序列进行观测约束 ---
        y_obs = test_y[t] if test_y is not None and t < len(test_y) else None
        y_ = ny.sample(
            "y", 
            dist.Normal(exp_val, sigma), 
            obs=y_obs,
            mask=current_mask
        )

        # --- 更新水平:缺失序列保持预测值,有效序列用观测值 ---
        updated_level = jnp.where(current_mask, y_, exp_val)
        return updated_level, y_

    # 准备scan的输入:每个时间步的t和对应的掩码
    if T + future > 0:
        time_steps = jnp.arange(1, T + future)
        # 处理掩码:训练阶段用传入的掩码,预测阶段全为False
        if test_mask is not None:
            train_masks = test_mask
            future_masks = jnp.zeros((future, N), dtype=bool)
            full_masks = jnp.concatenate([train_masks, future_masks], axis=0)
        else:
            full_masks = jnp.zeros((len(time_steps), N), dtype=bool)
        
        # 组合时间步和掩码作为scan输入
        scan_xs = (time_steps, full_masks)

        # 运行scan
        _, ys = scan(
            f=transition_fn,
            init=level_init,
            xs=scan_xs
        )

        # 保存预测结果
        if future > 0:
            ny.deterministic("y_forecast", ys[-future:])

# 拟合模型(只用训练数据的掩码)
train_mask = mask[:,:T]
kernel = NUTS(hierarchical_rw_masked_model)
mcmc = MCMC(kernel, num_warmup=500, num_samples=1000, num_chains=2)
mcmc.run(random.PRNGKey(0), y=y[:,:T], mask=train_mask, future=h)

3. 模型关键修改点解释

  1. 初始水平处理:用jnp.argmax(mask, axis=1)找到每个序列的第一个有效时间点,避免使用nan作为初始值;
  2. 逐时间步掩码:在sample时显式传入mask参数,Numpyro会自动忽略mask为False的序列的观测约束;
  3. 水平更新逻辑:缺失序列(mask为False)的水平不使用观测值,而是用预测的exp_val继续迭代,保证随机游走的连续性;
  4. 预测阶段掩码:预测时没有观测数据,掩码设为全False,模型按随机游走逻辑生成预测值。

三、验证模型结果

你可以用Arviz查看后验和预测结果:

import arviz as az

# 转换为Arviz的InferenceData
idata = az.from_numpyro(mcmc)

# 查看全局超参数的后验
az.plot_trace(idata, var_names=["mu_alpha", "mu_beta", "mu_sigma"])
plt.show()

# 生成预测样本
predictive = Predictive(hierarchical_rw_masked_model, mcmc.get_samples(), return_sites=["y_forecast"])
forecast_samples = predictive(random.PRNGKey(1), y=y[:,:T], mask=train_mask, future=h)["y_forecast"]

# 可视化单个序列的预测结果
seq_idx = 0
valid_train_times = np.where(train_mask[seq_idx])[0]
forecast_times = np.arange(T, T+h)

fig, ax = plt.subplots()
ax.plot(valid_train_times, y[seq_idx, valid_train_times], label='Training Data')
# 画预测均值和95% HPDI
ax.plot(forecast_times, forecast_samples.mean(axis=0)[:, seq_idx], label='Forecast Mean', color='r')
az.plot_hpd(forecast_times, forecast_samples[..., seq_idx], ax=ax, fill_kwargs={"alpha":0.2}, color='r')
ax.set(xlabel='Time', ylabel='Y', title=f'Forecast for Sequence {seq_idx}')
ax.legend()
plt.show()

内容来源于stack exchange

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.08 11:14:38