拓冰建站拓冰建站
首页 / 资讯中心 / 正文

如何正确读取 TimesFM 的 quantile_forecast 输出并提取中位数与预测区间?

如何正确读取 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 均值meanindex 1 起才是 q10。如果按下标 0 去取最低分位取到的会是均值预测区间会完全算错。本文基于仓库中 API 参考、SKILL.md 和 run_forecast.py 的示例说明如何正确解析这两个输出、提取中位数和预测区间并给出可执行的验证方法。适用前提TimesFM 2.5TimesFM_2p5_200M_torchPyTorch 后端模型已按仓库 README.md 说明安装timesfm[torch]。准备加载与编译模型forecast()在模型未编译时会抛出RuntimeError: Model is not compiled所以必须先用timesfm.ForecastConfig调用model.compile()。与分位数输出直接相关的三个配置项来自 api_reference.md配置项默认值作用use_continuous_quantile_headFalse启用 30M 参数连续分位数头文档标注 True 能给出更准确的预测区间尤其长 horizon 场景fix_quantile_crossingFalse对分位数做后处理保证 q10 ≤ q20 ≤ ... ≤ q90 单调有序infer_is_positiveTrue自动检测输入是否全为正并把预测钳制在 ≥ 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_context1024, max_horizon256, normalize_inputsTrue, use_continuous_quantile_headTrue, force_flip_invarianceTrue, infer_is_positiveTrue, fix_quantile_crossingTrue, ) ) point_forecast, quantile_forecast model.forecast( horizon12, 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(horizonH)指定的预测步长两个输出的形状如下来自 api_reference.md 的 Output Shape Reference输出形状含义point_forecast(B, H)中位数预测0.5 分位数quantile_forecast(B, H, 10)完整分位数分布quantile_forecast[:,:,0](B, H)均值Meanquantile_forecast[:,:,1](B, H)10% 分位数quantile_forecast[:,:,5](B, H)50% 分位数即中位数等于point_forecastquantile_forecast[:,:,9](B, H)90% 分位数quantile_forecast全部 10 个切面的完整对照来自 SKILL.md 的 Understanding the Output 一节下标分位数用途0mean均值预测10.180% 预测区间下界20.260% 预测区间下界50.5中位数point_forecast80.860% 预测区间上界90.980% 预测区间上界提取中位数与预测区间对单条序列inputs[values]point, quantiles model.forecast(horizon24, inputs[values]) median quantiles[0, :, 5] # 或直接用 point[0]二者同为中位数 lower_80 quantiles[0, :, 1] # q1080% PI 下界 upper_80 quantiles[0, :, 9] # q9080% PI 上界 lower_60 quantiles[0, :, 2] # q2060% PI 下界 upper_60 quantiles[0, :, 8] # q8060% PI 上界 mean quantiles[0, :, 0] # 均值注意下标 0 不是 q10对批量序列inputs长度为B用三维下标一次取全部序列point, quantiles model.forecast(horizon30, inputsinputs) 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并注释说明550% (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_crossingTrue文档说明该选项关闭时quantiles may occasionally cross。此外 Checklist 还要求输入序列 context 至少 32 个数据点api_reference.md 对输入的行为约定是前导 NaN 自动剥离、内部 NaN 线性插值、超过max_context的序列截断取末尾max_context个点、短于max_context的序列自动填充。常见错误与边界下标错位SKILL.md 的 Common Mistakes 第一条即指出quantiles[..., 0]是均值而非 q0q10 在下标 1q90 在下标 9。文档建议直接定义常量IDX_Q10, IDX_Q90 1, 9避免手写下标。未编译就预测会抛RuntimeErrorModel 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),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门