Toto-2.0-2.5B-FT-NPU开发API参考:Toto2Model.forecast接口与输入输出格式详解
2026/8/31 12:45:57 网站建设 项目流程

Toto-2.0-2.5B-FT-NPU开发API参考:Toto2Model.forecast接口与输入输出格式详解

【免费下载链接】Toto-2.0-2.5B-FT-NPU项目地址: https://ai.gitcode.com/z_studio/Toto-2.0-2.5B-FT-NPU

想把时间序列预测模型跑起来,第一步就是要搞懂它的核心 API。本文将带你快速掌握Toto2Model.forecast 接口:它是 Datadog 开源时序基础模型 Toto-2.0-2.5B-FT 在昇腾 NPU 上最核心的调用入口,负责把一段历史序列变成未来多步的概率预测。全文将详解 forecast 接口的输入格式、输出格式、常用参数与实战示例,新手也能照着写。

一、Toto2Model.forecast 接口是什么

Toto-2.0-2.5B-FT 是一个约24.5 亿参数的时序预测基础模型(非 LLM),采用 decoder-only patched transformer 架构。它不像大语言模型那样逐字生成文本,而是一次性把历史序列编码,直接输出未来预测

forecast接口就是模型对外暴露的唯一预测入口,你只需要准备好一个包含历史数据的字典,调用一次即可拿到完整的多分位预测结果。项目中的inference.py(本仓库的推理脚本)已经把这个接口封装好了,并自动处理了模型的加载、缩放与反缩放,可以直接参照学习。

二、forecast 接口输入格式详解

调用model.forecast()时,第一个参数是一个 Python 字典,包含三个必需的字段:

字段形状类型含义
target(batch, n_variates, time)float历史序列值,多维时序数据
target_mask(batch, n_variates, time)bool有效数据掩码,缺失位置为 False
series_ids(batch, n_variates)long序列分组 ID,用于区分不同时间序列

以本仓库最常用的单变量场景为例(batch=1、n_variates=1、上下文 512 点):

inputs = { "target": x, # shape (1, 1, 512) "target_mask": torch.ones_like(x, dtype=torch.bool), # 全部有效 "series_ids": torch.zeros(1, 1, dtype=torch.long), # 单个分组 }

新手最容易忽略的 3 个细节

  • 时间维放在最后一维,与常见(batch, time, feature)布局不同;
  • target_mask必须与target形状完全一致,缺失值处填False,模型会自行处理;
  • 模型内部通过PatchedCausalStdScaler对输入做自动缩放,所以直接喂入原始数值即可,无需手动归一化,输出也会自动反缩放回原始量纲。

三、forecast 接口输出格式详解

forecast()的返回值是一个形状为(9, batch, n_variates, horizon)的张量,其中:

  • 第一维 9:对应 9 个升序分位[0.1, 0.2, ..., 0.9],代表预测的不确定性区间;
  • 最后一维 horizon:预测的未来时间步数。

中位数即点预测:分位数组中的下标 4 对应 0.5 分位,可直接作为最可能的预测值使用。以本仓库实测为例(output/forecast.json),512 点上下文预测未来 96 点,输出形状为(9, 1, 1, 96)

quantiles = model.forecast(inputs, horizon=96, decode_block_size=768) median = quantiles[4, 0, 0, :] # 取出中位数预测,形状 (96,)

第 1 步的 9 个分位输出示例:0.1→85.418、0.3→85.489、0.5→85.531、0.7→85.592、0.9→85.686,分位越宽表示模型对预测越不确定,非常适合做区间告警等场景。

四、快速上手:一次完整的 forecast 调用

在昇腾 NPU 上运行,核心代码只需要四步。完整可运行的版本见仓库的inference.py

from toto2 import Toto2Model model = Toto2Model.from_pretrained("/data/models/Datadog/Toto-2.0-2.5B-FT") model = model.to(device="npu:0").eval() inputs = { # 见上文输入格式 "target": x.to("npu:0"), "target_mask": mask.to("npu:0"), "series_ids": ids.to("npu:0"), } with torch.no_grad(): quantiles = model.forecast(inputs, horizon=96, decode_block_size=768, has_missing_values=False)

提示:本仓库默认使用确定性合成小时序列(趋势 + 24h 日周期 + 168h 周周期)做零样本预测,保证 README 中的输入输出完全可复现;也可以使用--data传入自己的单列 CSV。

五、forecast 接口常用参数与调优建议

forecast接口还有三个高频参数,直接影响预测效果与性能:

参数说明推荐值
horizon预测长度(未来点数)96
decode_block_size分块解码大小,需为 32 的整数倍768
has_missing_values输入是否含缺失值False

调优建议

  • horizon越大预测时间越长,但模型对远端的不确定性也会上升,可观察 0.1/0.9 分位区间是否过宽;
  • decode_block_size建议不小于horizon,此时单次前向即可输出全部预测,避免分块带来的额外开销;
  • 上下文长度(context-length,默认 512)越接近模型的 patch 对齐,效果越稳定,模型按 patch=32 切块,任意长度均可输入。

六、在昇腾 NPU 上运行的关键注意事项

1. 推理引擎选择:Toto 是时序预测模型,vllm-ascend / sglang 的模型注册表不含该架构,唯一适用的引擎是 torch_npu

2. 推荐使用 fp32 精度:实测 fp32 下 NPU 与 CPU 参考结果最大绝对偏差仅0.000168,数值完全对齐;bf16 虽可用但精度下降明显(MAE 约 0.61),且无提速,不建议默认使用。

3. 性能参考:在 Ascend 910B 上,2.5B 模型 fp32 单次 96 步零样本预测约228ms,加载权重约需 30~40 秒,显存约 10GB,单卡即可长时间高吞吐服务。

4. 环境依赖:安装依赖时需使用pip install --no-deps --ignore-requires-python -r requirements.txt(详见requirements.txt注释),避免 pip 误装与 torch_npu 不匹配的 torch 版本。

七、总结

Toto2Model.forecast 接口的使用可以概括为一句话:构造target + target_mask + series_ids三个字段的字典,调用model.forecast(),拿到(9, batch, n_var, horizon)的多分位预测。记住输出第 4 个下标是中位数(点预测),输入无需手动归一化,你就能把 Toto-2.0-2.5B-FT 用起来了。更多细节可参考仓库内的README.md(模型说明与实测数据)、inference.py(完整可运行示例)与AGENT_WORKFLOW.md(适配过程全记录)。

【免费下载链接】Toto-2.0-2.5B-FT-NPU项目地址: https://ai.gitcode.com/z_studio/Toto-2.0-2.5B-FT-NPU

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询