m
mengsihan/NeoQuasar-Kronos-mini-NPU
模型介绍
文件和版本
Pull Requests
讨论
分析

Kronos-mini — 金融K线(OHLCV)时间序列预测 · 昇腾NPU适配

#NPU

模型简介

  • 模型: NeoQuasar/Kronos-mini
  • 配套Tokenizer: NeoQuasar/Kronos-Tokenizer-2k
  • revision: main (config SHA f4e68697d9d5aed55cef5c96aabc3376bcad9f81)
  • 任务类型: time-series-forecasting(金融K线 OHLCV 序列预测)
  • 架构: Kronos 两阶段框架 —— 专用 tokenizer(层级离散量化)将连续多维 K 线量化为层次化离散 token,再由自回归 Transformer 预测。
  • 参数量: 4.1M(d_model=256, n_layers=4, n_heads=4, s1_bits=10, s2_bits=10,context=2048)
  • License: MIT(arXiv 2508.02739)

数据契约

推理脚本按官方 KronosPredictor 契约消费 5分钟金融K线数据,必须包含列:

open, high, low, close, volume, amount(可选 timestamps;缺失 timestamps 时自动按 5min 频率生成)。

  • 列顺序固定为 open → high → low → close → volume → amount,脚本会严格校验缺列。
  • volume/amount 缺失时按官方行为补零或 volume * mean(price) 填充;存在 NaN 时报错。
  • 时间特征:minute, hour, weekday, day, month 由 timestamps 提取,作为时间嵌入输入。
  • 预处理:按窗口 mean/std 标准化并裁剪到 [-5, 5],推理后逆标准化还原价格。
  • 本次验证使用真实可追溯序列:shiyu-coder/Kronos 仓库 finetune_csv/data/HK_ali_09988_kline_5min_all.csv (香港阿里 09988,2019-11-26 起 5min K线)。取连续 520 行窗口:lookback=400,pred_len=120。

环境

  • 硬件: 单卡昇腾 Ascend910B(本验证设备 Ascend910_9362,npu:0)
  • 软件: Python 3.11, torch 2.9.0+cpu, torch_npu 2.9.0.post1+gitee7ba04, transformers 4.57.6, CANN 8.5.1
  • 依赖见 requirements.txt

安装

pip install -r requirements.txt
# 下载权重到本地目录(提交仓不含权重)
hf download NeoQuasar/Kronos-mini --local-dir models/Kronos-mini
hf download NeoQuasar/Kronos-Tokenizer-2k --local-dir models/Kronos-Tokenizer-2k

NPU 推理

默认命令(需先下载权重到 models/,数据放到 data/ 或 --data 指定):

python inference.py
# 可选参数: --model-dir --tokenizer-dir --data --lookback 400 --pred-len 120 --device npu --seed 42

inference.py 为自包含脚本(内联模型/分词器/量化器定义),从本地 models/ 加载同一份权重, 把模型与全部参与计算的 Tensor 显式迁移到 npu:0。异常时返回非零退出码。

真实结果(npu:0)

输入窗口:HK 阿里 09988,400 根 5min K线(2019-11-26 09:35 → 2019-12-04 09:50), 预测未来 120 个 5min K线(2019-12-04 09:55 → 2019-12-05 15:20)。

指标数值
close MAE2.209
close RMSE2.688
同步推理耗时1.28 s

预测样例(close,前5个预测点):

时间close 预测
2019-12-04 09:55186.54
2019-12-04 10:00186.68
2019-12-04 10:05186.74
2019-12-04 10:10186.91
2019-12-04 10:15186.85

预测为概率采样路径均值(sample_count=1, top_p=0.9, T=1.0);零样本预测,未在该序列上微调。

一致性(CPU-NPU)

同一权重、同一数据窗口、同一预处理、同一 seed、eval 模式,比较确定性核心计算:

  • tokenizer.encode + decode_s1 输出的 s1 logits:max_abs_error = 1.43e-05,PASS
  • decode_s2 输出的 s2 logits:max_abs_error = 1.53e-05,PASS
  • 判定阈值:atol=1e-4, rtol=1e-3(FP32)

完整预测输出(含随机采样 torch.multinomial)因 CPU/NPU 的 RNG 实现不同会有采样路径差异, close 相关系数 0.92,最大相对偏差约 1.3%;模型核心计算(logits)在两种设备上数值一致, 未发生 CPU fallback(模型参数与全部中间 Tensor 均位于 npu:0)。

性能(npu:0)

batch=1, lookback=400, pred_len=120, 6通道, dtype=float32,预热 3 次、计时 10 次:

项目数值
首次编译(含构图)1.270 s
稳定平均延迟0.776 s
最小 / 最大0.725 / 0.793 s
p50 / p90 / p950.787 / 0.791 / 0.792 s
吞吐1.29 series/s
峰值显存83.9 MiB

证据图

以下图片由 xterm.js 依据本次真实执行日志生成(非手工绘制):

  • assets/agent_workflow.png — 侦察、数据下载、模型加载、NPU 推理、一致性、性能与提交校验流水线日志
  • assets/npu_device_call.png — npu-smi、NPU 可用性、设备名、模型参数 device/dtype、输入输出 Tensor device
  • assets/model_result.png — 默认 python inference.py 的真实完整输出(含窗口摘要、真实预测、同步耗时与状态)

限制

  • 随机采样路径在 CPU/NPU 间 RNG 实现不同,完整预测序列会有采样差异;数值一致性以确定性 logits 为准。
  • 本仓不包含权重与数据集,运行前需按「安装」下载权重并准备 OHLCV 数据。
  • 时间序列数值一致性为固定窗口验证,不是完整 benchmark 排名。