Toto-2.0-4m-npu API 参考手册:forecast 方法参数、返回值与 3 个调用示例

发布时间:2026/8/21 16:27:01
Toto-2.0-4m-npu API 参考手册:forecast 方法参数、返回值与 3 个调用示例 Toto-2.0-4m-npu API 参考手册forecast 方法参数、返回值与 3 个调用示例【免费下载链接】toto-2.0-4m-npu项目地址: https://ai.gitcode.com/atlasleong/toto-2.0-4m-npuToto-2.0-4m-npu 是 Datadog Toto 2.0 系列多变量时间序列预测基础模型的昇腾 NPU 推理交付版本专为在 Ascend 910B4 上运行而设计开箱即用、无 CPU 回退。本文是一份面向新手的 API 参考手册重点讲解核心入口forecast方法target、target_mask、series_ids等参数的完整含义quantiles返回值的结构与 9 个分位水平以及 3 个可以直接照抄的时间序列预测调用示例让你十分钟内跑通第一个预测任务。Toto-2.0-4m-npu 是什么为昇腾 NPU 而生的时间序列预测模型Toto-2.0-4m 是 Datadog 开源的时间序列基础模型家族成员参数量约 400 万4m采用 decoder-only 分块 Transformer 架构时间轴因果注意力与变量轴全量注意力交替堆叠配 9 分位输出头支持多变量概率预测probabilistic forecasting。Toto-2.0-4m-npu 交付版把模型权重、配置与推理入口全部打包通过torch_npu注册 NPU 后端主前向全程在逻辑设备npu:0上执行。获取完整交付代码含入口脚本与权重可通过 git clone 仓库地址https://gitcode.com/atlasleong/toto-2.0-4m-npu仓库内inference.py即为交付入口model/config.json保存模型超参数。模型关键配置一览项目内容参数量4,144,448约 4mfp32 权重架构decoder-only 分块 Transformer4 层、4 头patch_size32输入(batch, n_variates, time)历史观测序列输出9 分位预测(9, batch, n_variates, horizon)运行环境Ascend 910B4 CANN 8.5.1 torch_npu 2.9.0forecast 方法参数详解一次搞懂全部入参forecast是模型唯一的推理入口第一个参数是一个包含三个张量的字典其后跟着若干控制预测行为的超参数。完整参数说明如下参数类型 / 形状含义必填targetfloat32(batch, n_variates, time)历史观测序列即你想预测的那段时间窗口数据✅target_maskbool(batch, n_variates, time)观测掩码True表示该位置有真实观测值缺失位置填False✅series_idslong(batch, n_variates)分组/序列 id用于区分不同变量属于哪条序列✅horizonint预测步长即向后预测多少个时间点✅decode_block_sizeint解码块大小0表示单次并行解码全部 horizon推荐has_missing_valuesbool输入序列是否含缺失值True会走最大算子兼容路径推荐⚠️ 小提示在昇腾 NPU 上官方推荐has_missing_valuesTrue这一设置会选取算子兼容性最好的计算路径这正是交付版inference.py中的标准写法。forecast 返回值解读quantiles 张量与 9 个分位水平forecast只返回一个张量quantiles形状为(9, batch, n_variates, horizon)其中第一维对应 9 个分位水平quantile levels: [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]第 0 层0.1 分位与第 8 层0.9 分位构成预测区间下界与上界第 4 层0.5 分位即中位数预测常作为点预测使用。以交付版实测为例输入(1, 1, 512)的序列horizon96返回(9, 1, 1, 96)共 864 个元素全体均值FORECAST0.019318。真实运行标记输出如下INPUT_DEVICEnpu:0 MODEL_DEVICEnpu:0 OUTPUT_DEVICEnpu:0 CPU_FALLBACKfalse FORECAST0.019318 EXIT_CODE03 个可直接照抄的 forecast 调用示例示例 1最小调用预测未来 96 步最简单的方式构造一个随机输入调用forecast即可得到分位数输出。import torch from toto2 import Toto2Model model Toto2Model.from_pretrained(model, local_files_onlyTrue).eval() target torch.randn(1, 1, 512) # (batch, n_variates, time) target_mask torch.ones_like(target, dtypetorch.bool) series_ids torch.zeros(1, 1, dtypetorch.long) quantiles model.forecast( {target: target, target_mask: target_mask, series_ids: series_ids}, horizon96, ) print(quantiles.shape) # (9, 1, 1, 96)示例 2NPU 上的标准写法处理缺失值 确定性推理这是交付版inference.py的实际写法先把输入搬到npu:0在torch.no_grad()下调用并显式指定decode_block_size0与has_missing_valuesTrue。import torch import torch_npu # 注册 NPU 后端 device torch.device(npu:0) model model.to(device).eval() inputs { target: torch.randn(1, 1, 512).npu(), target_mask: torch.ones(1, 1, 512, dtypetorch.bool).npu(), series_ids: torch.zeros(1, 1, dtypetorch.long).npu(), } with torch.no_grad(): quantiles model.forecast( dict(inputs), horizon96, decode_block_size0, has_missing_valuesTrue, ) # 中位数预测 median quantiles[4] # (1, 1, 96)示例 3批量多变量预测4 条序列 × 3 个变量Toto 支持一次预测多条序列、多个相关变量只需把 batch 维和变量维加大其余代码几乎不变。target torch.randn(4, 3, 512) # 4 批、3 个变量、512 个历史点 target_mask torch.ones_like(target, dtypetorch.bool) series_ids torch.arange(4).unsqueeze(1).repeat(1, 3) # 每批一个 id quantiles model.forecast( {target: target, target_mask: target_mask, series_ids: series_ids}, horizon48, ) print(quantiles.shape) # (9, 4, 3, 48)常见问题与注意事项为什么返回值第一维是 9因为模型有 9 个分位输出头返回 0.10.9 共 9 个分位水平用于刻画预测不确定性而不是单一数值。decode_block_size填 0 还是 768交付版用0单次并行解码全部 horizon原仓库快速开始示例用768在 NPU 上推荐跟随交付版写0。NPU 上的 fp64 降级Ascend 910 不支持 fp64模型内部请求 float64 时会自动降级为 fp32前向结果已通过精度门禁CPU/NPU 最大绝对误差约 3.3e-6。结果可复现吗固定种子 42 固定权重快照下NPU 重复前向差异为 0.0输出完全确定。会不会偷偷回退到 CPU不会。inference.py显式import torch_npuNPU 不可用时直接报错绝不 CPU 回退。总结Toto-2.0-4m-npu 的forecast方法 API 非常简洁传入一个包含target、target_mask、series_ids的字典配好horizon、decode_block_size、has_missing_values三个参数就能拿到(9, batch, n_variates, horizon)的分位数预测结果。本文给出的 3 个调用示例覆盖了最小调用、NPU 标准写法和批量多变量场景直接复制即可上手。更多细节可阅读仓库内的inference.py与model/config.json祝你的时间序列预测之旅顺利【免费下载链接】toto-2.0-4m-npu项目地址: https://ai.gitcode.com/atlasleong/toto-2.0-4m-npu创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考