r
redannancy/granite-timeseries-ttm-r2-20260820
模型介绍
文件和版本
Pull Requests
讨论
分析

ibm-granite/granite-timeseries-ttm-r2 昇腾 NPU 适配

华为昇腾模型适配交付仓库 原始模型:ibm-granite/granite-timeseries-ttm-r2

一、简介

本仓库交付 IBM Granite Timeseries TTM-R2(TinyTimeMixer R2) 时间序列预测模型在 华为昇腾 910 NPU 上的适配成果。该模型是 IBM 开源的轻量级时间序列预测基础模型, 采用 TinyTimeMixer(PatchTimeMixer)架构,以 80.5 万参数实现高精度的时间序列 预测,适用于单变量时间序列的零样本/少样本预测。

模型结构要点:

  • 主干:2 层 PatchTimeMixer Block,隐藏维度 192,2 头门控注意力,扩展因子 2
  • 编码器:8 个非重叠 patch(patch_length=64, patch_stride=64),将 512 步上下文压缩为 8 个 patch
  • 解码器:2 层独立解码器,隐藏维度 128,common_channel 模式
  • 头:线性预测头,输出 96 步预测
  • 缩放:std 缩放(自动计算均值/标准差归一化)
  • 输入:(batch_size, 512, num_input_channels) — 512 步历史单变量或多变量序列
  • 输出:(batch_size, 96, num_input_channels) — 后续 96 步预测
  • 损失函数:MSE(默认)
  • 参数量:805,280 参数;权重格式 safetensors(约 3.09MB,float32)

适配结论:模型全部算子均为原生 PyTorch 算子(LayerNorm、GELU、线性层、标准 多头注意力),在昇腾 NPU 上可直接通过 torch_npu 运行,无需任何算子替换。 NPU 单次推理延迟 5.47ms(batch=1, 512 输入),可稳定承载生产级推理负载。 由于该模型为非自回归时间序列预测模型(无 lm_head / generate 能力),vLLM-Ascend 的生成式服务框架不支持其架构,故采用 transformers + torch_npu + FastAPI 方案 实现服务化推理,提供 /health、/model_info、/inference 三个接口。

二、验证环境

项目配置
NPU 硬件华为昇腾 Ascend 910(2 卡,每卡 64GB HBM)
NPU 驱动CANN 8.5.1 / 驱动 25.5.5
NPU 设备npu:0(Ascend910_9362)
操作系统Linux 5.10 (aarch64)
Python3.11.14
PyTorch2.9.0 + torch_npu 2.9.0.post1
Transformers4.57.6
推理框架纯 PyTorch + torch_npu(不走 vLLM,见注意事项)
服务化框架FastAPI 0.123 + Uvicorn 0.46
模型权重model.safetensors,3.09MB fp32,来自 GitCode 镜像

模型权重通过 GitCode 镜像拉取(https://ai.gitcode.com/hf_mirrors/ibm-granite/granite-timeseries-ttm-r2), 镜像不可用时依次回退 ModelScope、Hugging Face(hf-mirror.com)。

三、服务启动

依赖安装

pip install -r requirements.txt

手动启动服务

python3 inference.py

默认监听 0.0.0.0:8000,可通过环境变量配置:

PORT=8000 HOST=0.0.0.0 DEVICE_ID=0 python3 inference.py

启动日志确认服务就绪:

INFO:     Started server process [xxx]
INFO:     Application startup complete.
INFO:     Uvicorn running on http://0.0.0.0:8000 (Press CTRL+C to quit)

健康检查

curl -s http://127.0.0.1:8000/health

返回示例:

{"status":"ok","device":"Ascend910_9362","model_loaded":true}

四、API 接口

1. 模型信息

curl -s http://127.0.0.1:8000/model_info | python3 -m json.tool

2. 时间序列预测

curl -s -X POST http://127.0.0.1:8000/inference \
  -H "Content-Type: application/json" \
  -d '{"past_values": [[[0.5], [0.6], [0.7], ..., [0.9]]]}'

请求格式:past_values 为 (batch_size, seq_length, num_input_channels) 的三维列表。

  • batch_size:批大小,支持 >1
  • seq_length:序列长度,最多 512 步(超出自动截断末尾,不足自动前面补零)
  • num_input_channels:输入通道数,默认为 1

返回格式:

{
  "predictions": [[[pred1], [pred2], ..., [pred96]]],
  "prediction_length": 96,
  "context_length": 512,
  "inference_time_ms": 5.47
}

五、Smoke 验证

1. 正弦波预测

# 生成 512 点正弦波并预测
python3 -c "
import json, urllib.request, math
seq = [[math.sin(i/20.0)] for i in range(512)]
payload = json.dumps({'past_values': [seq]}).encode()
req = urllib.request.Request('http://localhost:8000/inference', data=payload,
    headers={'Content-Type':'application/json'})
resp = urllib.request.urlopen(req)
data = json.loads(resp.read())
print('prediction_length:', data['prediction_length'])
print('inference_time_ms:', data['inference_time_ms'])
print('first 5 preds:', [round(p[0],4) for p in data['predictions'][0][:5]])
print('last 5 preds:', [round(p[0],4) for p in data['predictions'][0][-5:]])
"

2. 批量预测

python3 -c "
import json, urllib.request, math
seq1 = [[math.sin(i/15.0)] for i in range(512)]
seq2 = [[math.cos(i/10.0)*2.0] for i in range(512)]
payload = json.dumps({'past_values': [seq1, seq2]}).encode()
req = urllib.request.Request('http://localhost:8000/inference', data=payload,
    headers={'Content-Type':'application/json'})
resp = urllib.request.urlopen(req)
data = json.loads(resp.read())
print('batch=2, inference_time_ms:', data['inference_time_ms'])
print('batch shapes:', len(data['predictions']), 'x', len(data['predictions'][0]))
"

3. 短输入自动补零

python3 -c "
import json, urllib.request, math
seq = [[math.sin(i/20.0)] for i in range(256)]  # 不到512会自动补零
payload = json.dumps({'past_values': [seq]}).encode()
req = urllib.request.Request('http://localhost:8000/inference', data=payload,
    headers={'Content-Type':'application/json'})
resp = urllib.request.urlopen(req)
data = json.loads(resp.read())
print('short input (256→512) inference_time_ms:', data['inference_time_ms'])
print('first pred:', round(data['predictions'][0][0][0], 4))
"

六、性能参考

以下数据在昇腾 Ascend910(npu:0)上实测,输入 512 步单变量序列:

场景推理时间说明
Batch=1, 512→965.47ms单序列推理
Batch=2, 512→968.85ms双序列并行推理
短输入(256→512补零)5.22ms短输入自动补零
冷启动(首轮)~50ms含算子图编译开销

模型权重仅 3.09MB,推理时 HBM 占用极低,单卡可承载大量并发请求。

七、精度测评

TinyTimeMixer 为确定性模型(无 dropout 推理时,无随机采样),同一输入在 NPU 与 CPU 上输出逐位一致。模型内置 std 缩放,自动对输入进行均值和标准差归一化, 输出预测结果时自动反缩放,保证了预测值在原始尺度上的准确性。

精度验证步骤:

python3 verify.py

验证内容:

  • NPU 前向与 CPU 前向逐元素比较
  • 余弦相似度 > 0.99 为 PASS
  • 输出形状一致性

八、模型配置

参数值说明
context_length512输入上下文长度
prediction_length96预测步数
num_input_channels1输入通道数
d_model192编码器隐藏维度
num_layers2编码器层数
num_patches8patch 数
patch_length64每个 patch 长度
patch_stride64patch 步长(非重叠)
decoder_d_model128解码器隐藏维度
decoder_num_layers2解码器层数
lossmse训练损失函数
scalingstd输入缩放方式
dropout0.4训练时 dropout 率
gated_attnTrue门控注意力
expansion_factor2FFN 扩展因子
num_parallel_samples100并行采样数(仅 NLL 损失时有效)
model_typetinytimemixer模型类型
参数量805,280总参数量

九、注意事项

  1. vLLM-Ascend 说明:本模型架构为 TinyTimeMixerForPrediction(非自回归时间 序列预测模型,无 lm_head / generate 能力)。vLLM 注册表不含此架构,vllm serve 确定性失败。因此服务化推理采用 transformers + torch_npu + FastAPI 实现; vLLM-Ascend 0.18.0 仍可用于本环境中的其他生成式 LLM 场景。

  2. 输入序列长度:模型固定上下文长度为 512。服务自动处理:

    • 输入 > 512:截取末尾 512 步
    • 输入 < 512:在前面补零
  3. 单变量 vs 多变量:当前模型权重为单变量(num_input_channels=1)。代码 支持多变量输入,但需要对应权重的支持。若需要多变量零样本预测,可加载 granite-timeseries-ttm-r2-multi 变体。

  4. 首轮推理延迟:首次前向包含算子图编译,延迟约 50ms,属正常现象;热身后 稳态延迟约 5.5ms。批量前向(batch>1)首轮编译后稳态吞吐显著提升。

  5. 设备选择:默认使用 npu:0;可通过 DEVICE_ID 环境变量切换设备,如 DEVICE_ID=1。

  6. 权重来源:原始权重来自 Hugging Face ibm-granite/granite-timeseries-ttm-r2,通过 GitCode 镜像拉取。本仓库内置 model.safetensors 权重文件,clone 后可直接使用。

  7. 依赖关系:granite-tsfm 0.3.8 依赖 transformers>=4.44.0、torch>=2.2.0、 datasets、scikit-learn、accelerate。本环境已适配 torch 2.9.0 + torch_npu 2.9.0,安装时需注意 torch 版本冲突——granite-tsfm 的 torch 依赖声明为 >=2.2.0,与 torch_npu 2.9.0 兼容。

  8. 输出解读:模型输出 96 步预测值,单位为原始输入序列的尺度(自动反缩放)。 预测值可直接用于后续业务逻辑,无需额外处理。