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),仅供参考