一个已经拟合好的statsmodels模型,能算出的东西远不止大多数代码从里面取出的那串数字。点预测只是它能完成的最小任务。下面这三个技巧,本质上都是向结果对象索取它已经算好的东西,而不是自己动手重建一遍。

同一份数据集、同一个模型、三个经常被人重新实现的方法。三个技巧都跑在同一份月度序列和同一个拟合模型上,唯一的区别是调用了结果对象上的哪个方法。以下内容均基于statsmodels 0.15.0验证。

打开网易新闻 查看精彩图片

先安装statsmodels:

pip install statsmodels

技巧一:要区间,不只要数字

res.forecast(12)给你十二个数字。res.get_forecast(12)给你的是一个PredictionResults对象,它恰好也携带了模型已经估计出的不确定性。点预测用predicted_mean,边界用conf_int,两者一起拿用summary_frame()。

区间并不是额外的工作量,它们是同一次计算的结果,只是更短的那个方法把它们丢掉了:

import statsmodels.api as sm
from statsmodels.tsa.arima.model import ARIMA
co2 = sm.datasets.co2.load_pandas().data["co2"]
co2 = co2.resample("MS").mean().ffill()
train, recent = co2[:-12], co2[-12:]
res = ARIMA(train, order=(1, 1, 1), seasonal_order=(1, 1, 1, 12)).fit()
print(res.get_forecast(12).summary_frame().head())

输出:

co2 mean mean_se mean_ci_lower mean_ci_upper
2001-01-01 370.523929 0.322722 369.891406 371.156452
2001-02-01 371.253673 0.388214 370.492787 372.014559
2001-03-01 372.200726 0.429518 371.358887 373.042566
2001-04-01 373.468351 0.463501 372.559905 374.376797
2001-05-01 373.856957 0.494157 372.888427 374.825487

发布一个不带区间的预测是一种选择。如果你想要的是历史数据上的拟合值而不是未来的路径,get_prediction就是同一个思路应用在可以包含样本内区间的范围上。

技巧二:不重新拟合就加入新数据

十二个月的新观测到了。条件反射是把它们拼接到训练数据上再调用一次.fit(),这会从头重新估计每一个参数。append做的事情更便宜:它在合并后的数据上重建结果对象,并且在refit=False时保留你已经估计出的参数:

updated = res.append(recent, refit=False)
print(updated.get_forecast(6).summary_frame().head())

输出:

co2 mean mean_se mean_ci_lower mean_ci_upper
2002-01-01 371.969954 0.322722 371.337432 372.602477
2002-02-01 372.750021 0.388214 371.989135 373.510907
2002-03-01 373.654908 0.429518 372.813068 374.496748
2002-04-01 374.834276 0.463501 373.925830 375.742722
2002-05-01 375.328719 0.494157 374.360189 376.297248

默认就是refit=False,复用你已经有的估计值。当积累的新数据足够多、你希望重新计算它们时,传入refit=True。

这类方法有三个:

  • append:在原始数据和新数据上重新跑一遍滤波
  • extend:只对新观测做滤波,历史很长时更快
  • apply:用于另一份数据集,而不是当前数据的延续

技巧三:让STL处理季节性

手动版本的流程分三步:

  1. 分解序列
  2. 预测季节性调整后的部分
  3. 再把季节成分加回到预测上

符号错误和索引错位就发生在第三步。STLForecast把整个循环封装成一个对象。文档把它描述为:先用STL减去估计出的季节性,再用时间序列模型(例如ARIMA)预测去季节化后的数据,然后进行预测。

from statsmodels.tsa.forecasting.stl import STLForecast
stlf = STLForecast(train, ARIMA, model_kwargs={"order": (1, 1, 1), "trend": "t"})
print(stlf.fit().forecast(12).head())

输出:

2001-01-01 370.529117
2001-02-01 370.963627
2001-03-01 371.921080
2001-04-01 373.137720
2001-05-01 373.144192
Freq: MS, dtype: float64

注意传进去的是什么:ARIMA类本身,而不是一个已拟合的实例,它的参数通过model_kwargs单独传入。这是这个API里唯一真正让人意外的地方,传入ARIMA(...)是大多数人在这里犯的第一个错误。

收尾

这里的每一个技巧,都是你已经构建好的对象上已经存在的方法。手工实现的替代方案更长、更慢,而且更容易出错,所以在这种情况下偏离内置方法确实不值得。它通常之所以被写出来,是因为没人看过.fit()返回了什么。读一读结果对象,然后别再重写它了。