TDgpt 预测分析算法开发指南:基于 AbstractForecastService 为 TDengine 扩展 FORECAST 能力
【免费下载链接】TDengineHigh-performance, scalable time-series database designed for Industrial IoT (IIoT) scenarios项目地址: https://gitcode.com/GitHub_Trending/tde/TDengine
TDgpt 是 TDengine 的外置式时序数据分析智能体,通过FORECAST函数为 SQL 查询注入时间序列预测能力。本文以预测分析算法开发为主题,完整讲解预测算法的输入/输出约定、AbstractForecastService父类属性、示例算法开发、部署注册、SQL 调用与单元测试全流程,并深入tools/tdgpt源码剖析类的动态加载机制与参数校验规则,帮助开发者按约定将自己开发的预测算法快速接入 TDgpt,实现"开发一次、SQL 随处调用"。
预测算法在 TDgpt 中的定位
在 TDengine 中,时序数据预测通过FORECAST函数对外提供:基于输入的(历史)时间序列数据,调用指定(或默认)预测算法给出后续时间序列的预测数据。查询执行时,TDengine 中的 Vnode 会将涉及时序数据高级分析的部分直接转发到 Anode(分析节点),等待分析完成后将结果组装进查询执行流程。
TDgpt 是一个开放系统,Anode 采用 Python 类动态加载模式,启动时扫描特定目录内满足约定条件的代码文件并加载到系统中。得益于 TDgpt 与 TDengine 主进程taosd的松散耦合,Anode 算法升级对taosd完全没有影响——应用侧通常只需调整 SQL 中algo参数即可完成分析能力升级。
向 TDgpt 添加自定义预测算法只需三步(详见 开发者指南):
- 开发完成符合要求的预测分析算法类;
- 将代码文件放入对应目录,然后重启 Anode;
- 使用 SQL 命令将 Anode 添加到 TDengine(首次部署)并刷新算法缓存,随后即可用 SQL 调用。
输入约定:execute 与 self.list
execute是预测分析算法的核心方法。框架调用该方法之前,会在对象属性参数self.list中设置完毕用于预测的历史时间序列数据。
从源码看,这一约定定义在 taosanalytics/base.py 的AbstractAnalyticsService基类中:
class AbstractAnalyticsService(AnalyticsService, ABC): def __init__(self): self.list = None self.ts_list = None def set_input_list(self, input_list: list, input_ts_list: list = None): """set the input list""" self.list = input_list self.ts_list = input_ts_listset_input_list由框架在调用execute前触发,将历史序列数值写入self.list,对应的时间戳列表写入self.ts_list。预测算法作者只需要在execute中读取这两个成员即可拿到全部输入数据,无需关心数据从何而来。
此外,AbstractForecastService还提供了set_input_data方法,可额外设置两类动态回归(dynamic regression)所需的数据:
self.past_dynamic_real = [] # 过去一段时间的动态实况数据 self.dynamic_real = [] # 当前时刻的动态实况数据普通预测算法可忽略这两个变量,直接使用self.list作为唯一输入源。
输出约定与返回结构
execute方法执行完成后返回一个字典对象,预测返回结果如下:
return { "mse": mse, # 预测算法的拟合数据最小均方误差 (minimum squared error) "res": res # 结果数组 [时间戳数组,预测结果数组,预测结果执行区间下界数组,预测结果执行区间上界数组] }res是一个列表,其内部结构为:
| 数组位置 | 内容 |
|---|---|
res[0] | 预测结果时间戳数组,长度为self.rows,从self.start_ts起按self.time_step递增 |
res[1] | 预测结果数值数组 |
res[2] | 预测结果置信区间下界数组 |
res[3] | 预测结果置信区间上界数组 |
时间戳数组在 base.py 的AbstractStatsForecastService.execute中有标准生成方式可供参考:
timestamps = [ self.start_ts + index * self.time_step for index in range(self.rows) ]当self.return_conf为0时,上界和下界与预测值相同(算法可不返回置信区间,直接复用预测值数组即可)。
父类 AbstractForecastService 属性说明
预测算法的父类AbstractForecastService位于 tools/tdgpt/taosanalytics/base.py,其对象属性如下:
| 属性名称 | 说明 | 默认值 |
|---|---|---|
period | 输入时间序列的周期性,多少个数据点表示一个完整的周期。如果没有周期性,设置为0即可 | 0 |
start_ts | 预测结果的开始时间 | 0 |
time_step | 预测结果的两个数据点之间时间间隔 | 0 |
rows | 预测结果的数量 | 0 |
return_conf | 预测结果中是否包含置信区间;为0时上界和下界与预测值相同 | 1 |
conf | 置信水平,取值须满足0 <= conf < 1.0(与 SQL 侧conf一致,常用0.95) | 0.95 |
除此之外,从 base.py 的构造函数还可以看到以下补充属性:
| 属性名称 | 说明 | 默认值 |
|---|---|---|
type | 算法类型标识,预测算法固定为"forecast" | "forecast" |
precision | 预测结果时间戳的精度单位 | "ms" |
tz | 本地时区信息,用于预测结果时间戳列表生成 | 系统本地时区 |
past_dynamic_real | 动态回归使用的历史实况数据列表 | [] |
dynamic_real | 动态回归使用的当前实况数据列表 | [] |
set_params 的参数校验规则
AbstractForecastService.set_params在 base.py 中对参数做了严格校验,自定义算法若直接调用父类实现(或想复用校验逻辑),需注意以下规则:
- 必填参数:
start_ts、time_step、rows三个键必须同时出现在参数 dict 中,缺失任何一个都会抛出ValueError("params are missing, start_ts, time_step, rows are all required"); time_step必须大于0,否则报错time_step should be greater than 0;rows必须大于0,否则报错forecast rows is not specified yet;period可选,默认为0,且不能小于0;conf可选,默认为0.95,取值必须满足0 <= conf < 1.0,否则报错invalid value of conf, should between 0 and 1.0;return_conf可选,默认为1;precision可选,默认"ms";tz可选,用于指定预测时间戳使用的时区。
父类还提供了get_params()方法,返回算法的当前参数快照,供SHOW ANODES FULL等管理接口展示。
开发一个示例预测算法
下面开发一个示例预测算法:对任意输入时间序列,固定返回预测值1。该示例虽然逻辑简单,但完整覆盖了类命名、类属性、execute实现、置信区间处理和参数透传的全部约定。
from taosanalytics.base import AbstractForecastService # 算法实现类名称 需要以下划线 "_" 开始,并以 Service 结束 class _MyForecastService(AbstractForecastService): """ 定义类,从 AbstractForecastService 继承并实现其定义的抽象方法 execute """ # 定义算法调用关键词,全小写 ASCII 码 name = 'myfc' # 该算法的描述信息 (建议添加) desc = """return the forecast time series data""" def __init__(self): """类初始化方法""" super().__init__() def execute(self): """ 算法逻辑的核心实现""" res = [] """这个预测算法固定返回 1 作为预测值,预测值的数量是用户通过 self.rows 指定""" ts_list = [self.start_ts + i * self.time_step for i in range(self.rows)] res.append(ts_list) # 设置预测结果时间戳列 """生成全部为 1 的预测结果 """ res_list = [1] * self.rows res.append(res_list) """检查用户输入,是否要求返回预测置信区间上下界""" if self.return_conf: """对于没有计算预测置信区间上下界的算法,直接返回预测值作为上下界即可""" bound_list = [1] * self.rows res.append(bound_list) # 预测结果置信区间下界 res.append(bound_list) # 预测结果执行区间上界 """返回结果""" return {"res": res, "mse": 0} def set_params(self, params): """该算法无需任何输入参数,直接调用父类函数,不处理算法参数设置逻辑""" return super().set_params(params)对该示例的逐段说明:
- 类命名:
_MyForecastService以下划线_开始、以Service结束。这是 Anode 动态加载机制的硬性约定,不符合该命名规则的类不会被识别为可加载算法。 - 类属性
name:算法的调用关键词,必须为全小写 ASCII 字符。SQL 中algo=myfc即对应此值;SHOW ANODES FULL显示的算法名也是它。 - 类属性
desc:算法描述信息,建议填写,会出现在算法列表中。 - 时间戳生成:
ts_list = [self.start_ts + i * self.time_step for i in range(self.rows)],这是标准的时间戳列生成方式,与父类AbstractStatsForecastService的实现一致。 - 置信区间:
self.return_conf为真时返回上下界数组;对无置信区间计算能力的算法,直接返回预测值数组作为上下界即可满足输出约定。 mse:示例固定返回0。真实算法应返回拟合数据的最小均方误差,作为预测质量的量化指标(AbstractStatsForecastService中通过np.nanmean((values - fitted) ** 2)计算)。
算法类加载的底层机制
Anode 启动时通过 service_registry.py 的register_all_services扫描algo/fc/(内置预测算法)与algo/custom/fc/(用户自定义预测算法)目录。加载逻辑的关键点包括:
- 跳过
__init__.py、__pycache__和非.py文件; - 对模块内每个类,只加载类名以
_开头且不是基类名的类(if class_name in ServiceRegistry._base_class_name or (not class_name.startswith("_")): continue); - 校验类确实定义在该模块内(
algo_cls.__module__ == module.__name__),避免误加载从其他模块导入的类; - 实例化后以
algo_cls.name为键注册到服务注册表self.services,重名会抛出RuntimeError。
因此,开发好的算法类只要满足命名与继承约定、name不与其他算法冲突,放进对应目录重启即可自动加载。
将算法部署到 Anode 并注册到 TDengine
目录放置
将开发完成的代码文件保存在./lib/taosanalytics/algo/fc/目录下(对应源码仓库为 tools/tdgpt/taosanalytics/algo/fc/),然后重启 taosanode 服务。仓库中该目录已内置了arima.py、holtwinters.py、prophet.py、ets.py、theta.py、timemoe.py、chronos.py等算法,可作为参照实现。
重启 Anode 服务
Linux 环境下使用 systemd 管理 Anode 服务(详见 Anode 管理):
systemctl restart taosanoded systemctl status taosanoded首次部署需先注册 Anode
如果是第一次启动 Anode,请按照 Anode 管理 中的步骤先将该 Anode 添加到 TDengine 系统中:
-- node_url 为 Anode 的 IP 和 PORT 组成的字符串 CREATE ANODE '127.0.0.1:6035';验证算法已加载
在 TDengine 命令行接口中执行SHOW ANODES FULL,可以看到新加入的算法出现在 forecast 类型列表中:
SHOW ANODES FULL;若算法列表未更新,可执行UPDATE ALL ANODES刷新分析算法缓存。
通过 SQL 调用自定义预测算法
算法注册成功后,应用即可通过 SQL 语句调用该预测算法。使用FORECAST函数,通过algo参数指定算法名:
-- 对 col 列进行预测,通过指定 algo 参数为 myfc 来调用新添加的预测类 SELECT _flow, _fhigh, _frowts, FORECAST(col_name, "algo=myfc") FROM foo;返回结果中的_frowts(预测时间戳)、_flow(置信区间下界)、_fhigh(置信区间上界)与execute返回值res中的各数组一一对应:res[0]时间戳列、res[2]下界列、res[3]上界列,而res[1]预测值本身由FORECAST函数直接返回。
SQL 侧还可以传递算法参数,例如:
SELECT _flow, _fhigh, _frowts, FORECAST(col_name, "algo=myfc,conf=0.95,rows=100,time_step=86400000") FROM foo;这些参数最终会被解析并传入set_params(参数解析与校验入口见 handlers/forecast.py 中的handle_forecast与add_forecast_params)。值得注意的是,若 SQL 中不指定algo,默认使用holtwinters算法(源码algo = req_json["algo"].lower() if "algo" in req_json else "holtwinters")。
编写单元测试
单元测试依赖 Pythonunittest包。在测试目录taosanalytics/test中的forecast_test.py中增加单元测试用例或添加新的测试文件。仓库对应的测试文件为 tools/tdgpt/tests/forecast_test.py,通过loader.get_service(name)从服务注册表获取算法实例后直接驱动执行:
def test_myfc(self): """ 测试 myfc 类 """ s = loader.get_service("myfc") # 设置用于预测分析的数据 s.set_input_list(self.get_input_list(), None) # 检查预测结果应该全部为 1 r = s.set_params( {"rows": 10, "start_ts": 171000000, "time_step": 86400 * 30, "start_p": 0} ) r = s.execute() expected_list = [1] * 10 self.assertEqlist(r["res"][1], expected_list)测试要点说明:
loader.get_service("myfc")返回注册表中name == 'myfc'的服务实例副本;set_input_list模拟框架注入历史数据;set_params传入rows、start_ts、time_step等预测参数,其中86400 * 30即按月为步长(以秒计的时间戳场景);execute后断言r["res"][1](预测值数组)全部为1。
对更复杂的算法,还可以像仓库内置算法一样断言mse值、置信区间上下界、时间戳数组长度与递增规律等,保证算法在不同参数组合下的行为可回归验证。
进阶扩展建议
- 复用统计模型基类:若算法基于 StatsForecast 等统计模型实现,可直接继承
AbstractStatsForecastService(同样位于 base.py),只需实现_fit_model,父类会统一完成置信区间格式化与 MSE 计算。 - 参考内置实现:algo/fc/ 目录下的
arima.py、holtwinters.py、prophet.py等是完整的预测算法示例,展示了真实算法如何利用self.period、self.conf等属性。 - 动态模型目录:除代码文件外,TDgpt 还支持从动态模型目录加载通过 JSON 配置文件描述的预训练模型(见 service_registry.py 的
register_service_from_file),适合集成训练好的机器学习模型。 - 自定义目录:用户自定义预测算法也可放入
algo/custom/fc/目录(对应源码 tools/tdgpt/taosanalytics/algo/custom/fc/),加载失败时不会导致 Anode 启动中断,适合迭代开发。
按照上述约定完成开发、部署与测试后,你的预测算法就具备了与内置算法完全一致的调用体验:应用侧仅修改 SQL 中algo参数即可切换预测引擎,实现分析能力的按需扩展与平滑升级。
【免费下载链接】TDengineHigh-performance, scalable time-series database designed for Industrial IoT (IIoT) scenarios项目地址: https://gitcode.com/GitHub_Trending/tde/TDengine
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考