如何正确读取 TimesFM 的 quantile_forecast 输出并提取中位数与预测区间?
【免费下载链接】timesfmTimesFM (Time Series Foundation Model) is a pretrained time-series foundation model developed by Google Research for time-series forecasting.项目地址: https://gitcode.com/GitHub_Trending/ti/timesfm
调用 TimesFM 的model.forecast()后,返回值是二元组(point_forecast, quantile_forecast)。很多开发者在这里踩坑:quantile_forecast最后一个维度的 10 个切面并不是从 q0 开始的 10 个分位数,而是index 0 = 均值(mean),index 1 起才是 q10。如果按下标 0 去取"最低分位",取到的会是均值,预测区间会完全算错。本文基于仓库中 API 参考、SKILL.md 和 run_forecast.py 的示例,说明如何正确解析这两个输出、提取中位数和预测区间,并给出可执行的验证方法。
适用前提:TimesFM 2.5(TimesFM_2p5_200M_torch,PyTorch 后端),模型已按仓库 README.md 说明安装timesfm[torch]。
准备:加载与编译模型
forecast()在模型未编译时会抛出RuntimeError: Model is not compiled,所以必须先用timesfm.ForecastConfig调用model.compile()。与分位数输出直接相关的三个配置项(来自 api_reference.md):
| 配置项 | 默认值 | 作用 |
|---|---|---|
use_continuous_quantile_head | False | 启用 30M 参数连续分位数头,文档标注 True 能给出更准确的预测区间,尤其长 horizon 场景 |
fix_quantile_crossing | False | 对分位数做后处理,保证 q10 ≤ q20 ≤ ... ≤ q90 单调有序 |
infer_is_positive | True | 自动检测输入是否全为正并把预测钳制在 ≥ 0;文档明确:温度、收益率、PnL 等可为负的序列需设为 False |
仓库给出的完整调用(节选自 README.md 的 Code Example,数值保持原样):
import torch import numpy as np import timesfm torch.set_float32_matmul_precision("high") model = timesfm.TimesFM_2p5_200M_torch.from_pretrained("google/timesfm-2.5-200m-pytorch") model.compile( timesfm.ForecastConfig( max_context=1024, max_horizon=256, normalize_inputs=True, use_continuous_quantile_head=True, force_flip_invariance=True, infer_is_positive=True, fix_quantile_crossing=True, ) ) point_forecast, quantile_forecast = model.forecast( horizon=12, inputs=[ np.linspace(0, 1, 100), np.sin(np.linspace(0, 20, 67)), ], # Two dummy inputs ) point_forecast.shape # (2, 12) quantile_forecast.shape # (2, 12, 10): mean, then 10th to 90th quantiles.输出形状与下标含义
设B为输入序列条数(inputs列表长度),H为forecast(horizon=H)指定的预测步长,两个输出的形状如下(来自 api_reference.md 的 Output Shape Reference):
| 输出 | 形状 | 含义 |
|---|---|---|
point_forecast | (B, H) | 中位数预测(0.5 分位数) |
quantile_forecast | (B, H, 10) | 完整分位数分布 |
quantile_forecast[:,:,0] | (B, H) | 均值(Mean) |
quantile_forecast[:,:,1] | (B, H) | 10% 分位数 |
quantile_forecast[:,:,5] | (B, H) | 50% 分位数(即中位数,等于point_forecast) |
quantile_forecast[:,:,9] | (B, H) | 90% 分位数 |
quantile_forecast全部 10 个切面的完整对照(来自 SKILL.md 的 Understanding the Output 一节):
| 下标 | 分位数 | 用途 |
|---|---|---|
| 0 | mean | 均值预测 |
| 1 | 0.1 | 80% 预测区间下界 |
| 2 | 0.2 | 60% 预测区间下界 |
| 5 | 0.5 | 中位数(=point_forecast) |
| 8 | 0.8 | 60% 预测区间上界 |
| 9 | 0.9 | 80% 预测区间上界 |
提取中位数与预测区间
对单条序列(inputs=[values]):
point, quantiles = model.forecast(horizon=24, inputs=[values]) median = quantiles[0, :, 5] # 或直接用 point[0],二者同为中位数 lower_80 = quantiles[0, :, 1] # q10,80% PI 下界 upper_80 = quantiles[0, :, 9] # q90,80% PI 上界 lower_60 = quantiles[0, :, 2] # q20,60% PI 下界 upper_60 = quantiles[0, :, 8] # q80,60% PI 上界 mean = quantiles[0, :, 0] # 均值(注意:下标 0 不是 q10)对批量序列(inputs长度为B),用三维下标一次取全部序列:
point, quantiles = model.forecast(horizon=30, inputs=inputs) results = { col: { "median": point[i].tolist(), "lower_80": quantiles[i, :, 1].tolist(), "upper_80": quantiles[i, :, 9].tolist(), } for i, col in enumerate(column_names) }上式取数方式来自 SKILL.md 的 Batch Forecasting 工作流。仓库中还有一条端到端参考:run_forecast.py 把quantiles[:, 0]到quantiles[:, 9]逐列映射为mean, q10, q20, ..., q90写进forecast_output.csv和forecast_output.json,并注释说明5=50% (median)。该示例的运行产物(含 80%/60% CI 扇形图)可对照 forecast_visualization.png,图中内、外两条红色带分别是 60% 与 80% 区间,中位数曲线即点预测。
注意一个容易混淆的点:point_forecast的文档描述是"median forecast",而quantile_forecast[:,:,0]是均值。两者通常接近但不保证相同,取哪个取决于你的下游用途。
验证结果是否正确
SKILL.md 的 Quality Checklist 给出了每次任务完成后应核对的项,与本文任务相关的有:
import numpy as np # 1. 形状检查:point 为 (n_series, horizon),quantiles 为 (n_series, horizon, 10) assert point_forecast.shape == (2, 12) assert quantile_forecast.shape == (2, 12, 10) # 2. 中位数一致性:point_forecast 应等于 quantile_forecast 的 index 5 assert np.allclose(point_forecast, quantile_forecast[:, :, 5]) # 3. 无 NaN assert not np.isnan(point_forecast).any()如果希望分位数保持单调有序(q10 ≤ q20 ≤ ... ≤ q90),在compile()时设置fix_quantile_crossing=True;文档说明该选项关闭时"quantiles may occasionally cross"。此外 Checklist 还要求输入序列 context 至少 32 个数据点;api_reference.md 对输入的行为约定是:前导 NaN 自动剥离、内部 NaN 线性插值、超过max_context的序列截断取末尾max_context个点、短于max_context的序列自动填充。
常见错误与边界
- 下标错位:SKILL.md 的 Common Mistakes 第一条即指出,
quantiles[..., 0]是均值而非 q0;q10 在下标 1,q90 在下标 9。文档建议直接定义常量IDX_Q10, IDX_Q90 = 1, 9避免手写下标。 - 未编译就预测:会抛
RuntimeError(Model is not compiled),先调用model.compile(ForecastConfig(...))。 - 输入不是列表:传入单个 array 而非 list 会抛
ValueError: inputs must be list,需包一层[array]。 - 负值序列:对可为负的序列(温度、收益率等),
infer_is_positive必须设为 False,否则预测会被钳制到 ≥ 0。 - 版本差异:TimesFM 1.0/2.0 的 API 已归档在
v1/目录,返回的是experimental_quantile_forecast,且月度数据需要传freq参数;TimesFM 2.5 移除了 frequency 标志。本文全部下标与配置说明仅针对 2.5 接口。 - 显存不足:
torch.cuda.OutOfMemoryError时按文档降低per_core_batch_size,或分块调用forecast()。
参数完整默认值与全部行为说明见 api_reference.md;带协变量的forecast_with_covariates()返回同样结构的(point, quantiles)二元组,但需要额外安装timesfm[xreg],不在本文范围内。
【免费下载链接】timesfmTimesFM (Time Series Foundation Model) is a pretrained time-series foundation model developed by Google Research for time-series forecasting.项目地址: https://gitcode.com/GitHub_Trending/ti/timesfm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考