在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])
二、处理不同长度/含缺失值的时间序列
你的问题核心在于掩码的正确传递和初始值的处理:
- 原始代码中直接用
y[:,0]作为初始水平,但如果序列的前几个时间点是缺失的(nan),这个初始值会无效; - 在
scan中使用掩码时,需要针对每个时间点的每个序列单独应用掩码,而不是全局掩码; 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. 模型关键修改点解释
- 初始水平处理:用
jnp.argmax(mask, axis=1)找到每个序列的第一个有效时间点,避免使用nan作为初始值; - 逐时间步掩码:在
sample时显式传入mask参数,Numpyro会自动忽略mask为False的序列的观测约束; - 水平更新逻辑:缺失序列(mask为False)的水平不使用观测值,而是用预测的
exp_val继续迭代,保证随机游走的连续性; - 预测阶段掩码:预测时没有观测数据,掩码设为全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
相关产品推荐
相关产品推荐

