运行StatsForecast遇TypeError:预期dtype object却获numpy.dtype[float32]
问题
严格遵循《使用ARIMA和ETS进行预测》教程,在Anaconda的Jupyter Notebook中执行代码,运行到Y_hat_df = sf.forecast(horizon)时抛出错误:TypeError: expected dtype object, got 'numpy.dtype[float32]'。
完整代码
import numpy as np import pandas as pd from IPython.display import display, Markdown import matplotlib.pyplot as plt from statsforecast import StatsForecast from statsforecast.models import AutoARIMA, ETS, Naive #Imports the models you will use from statsforecast.utils import AirPassengersDF Y_df = AirPassengersDF Y_df.head()
输出:
unique_id ds y 0 1.0 1949-01-31 112.0 1 1.0 1949-02-28 118.0 2 1.0 1949-03-31 132.0 3 1.0 1949-04-30 129.0 4 1.0 1949-05-31 121.0
Y_train_df = Y_df[Y_df.ds<='1959-12-31'] # 132 monthly observations for train Y_test_df = Y_df[Y_df.ds>'1959-12-31'] # 12 monthly observations for test
season_length = 12 # Monthly data horizon = len(Y_test_df) # Predict the lenght of the test df # Include the models you imported models = [ AutoARIMA(season_length=season_length), ETS(season_length=season_length), Naive() ] # Instansiate the StatsForecast class as sf sf = StatsForecast( df=Y_train_df, models=models, freq='M', n_jobs=-1 ) # Forecast for the defined horizon Y_hat_df = sf.forecast(horizon) Y_hat_df.head()
报错堆栈
TypeError Traceback (most recent call last) <ipython-input-10-a9ee1bd8ce20> in <module> 18 19 # Forecast for the defined horizon ---> 20 Y_hat_df = sf.forecast(horizon) 21 22 Y_hat_df.head() ~\spyder\lib\site-packages\statsforecast\core.py in forecast(self, h, df, X_df, level, fitted, sort_df) 668 X, level = self._parse_X_level(h=h, X=X_df, level=level) 669 if self.n_jobs == 1: ---> 670 res_fcsts = self.ga.forecast( 671 models=self.models, 672 h=h, ~\spyder\lib\site-packages\statsforecast\core.py in forecast(self, models, h, fallback_model, fitted, X, level, verbose) 197 ) 198 else: ---> 199 raise error 200 cols_m = [ 201 key ~\spyder\lib\site-packages\statsforecast\core.py in forecast(self, models, h, fallback_model, fitted, X, level, verbose) 183 kwargs["level"] = level 184 try: ---> 185 res_i = model.forecast( 186 h=h, y=y_train, X=X_train, X_future=X_f, fitted=fitted, **kwargs 187 ) ~\spyder\lib\site-packages\statsforecast\models.py in forecast(self, y, h, X, X_future, level, fitted) 306 """ 307 with np.errstate(invalid="ignore"): ---> 308 mod = auto_arima_f( 309 x=y, 310 d=self.d, ~\spyder\lib\site-packages\statsforecast\arima.py in auto_arima_f(x, d, D, max_p, max_q, max_P, max_Q, max_order, max_d, max_D, start_p, start_q, start_P, start_Q, stationary, seasonal, ic, stepwise, nmodels, trace, approximation, method, truncate, xreg, test, test_kwargs, seasonal_test, seasonal_test_kwargs, allowdrift, allowmean, blambda, biasadj, parallel, num_cores, period) 1785 D = 0 1786 elif D is None: ---> 1787 D = nsdiffs( 1788 xx, period=m, test=seasonal_test, max_D=max_D, **seasonal_test_kwargs 1789 ) ~\spyder\lib\site-packages\statsforecast\arima.py in nsdiffs(x, test, alpha, period, max_D, **kwargs) 1608 while dodiff and D < max_D: 1609 D += 1 ---> 1610 x = diff(x, period, 1) 1611 if is_constant(x): 1612 return D ~\spyder\lib\site-packages\statsforecast\arima.py in diff(x, lag, differences) 583 def diff(x, lag, differences): 584 if x.ndim == 1: ---> 585 y = diff1d(x, lag, differences) 586 nan_mask = np.isnan(y) 587 elif x.ndim == 2: TypeError: expected dtype object, got 'numpy.dtype[float32]'
预期输出
unique_id ds AutoARIMA ETS Naive 1.0 1960-01-31 424.160156 406.651276 405.0 1.0 1960-02-29 407.081696 401.732910 405.0 1.0 1960-03-31 470.860535 456.289642 405.0 1.0 1960-04-30 460.913605 440.870514 405.0 1.0 1960-05-31 484.900879 440.333923 405.0
错误原因分析
报错根源是AirPassengersDF中的y列数据类型为float32,而statsforecast库底层ARIMA实现依赖的diff1d函数仅支持float64(numpy默认浮点类型),不兼容float32类型,导致类型匹配错误。
解决方案
将数据集的y列转换为float64类型即可解决问题,修改加载数据集的代码部分:
Y_df = AirPassengersDF.copy() Y_df['y'] = Y_df['y'].astype(np.float64)
完整修改后的代码片段:
import numpy as np import pandas as pd from IPython.display import display, Markdown import matplotlib.pyplot as plt from statsforecast import StatsForecast from statsforecast.models import AutoARIMA, ETS, Naive from statsforecast.utils import AirPassengersDF # 加载数据集并转换y列数据类型 Y_df = AirPassengersDF.copy() Y_df['y'] = Y_df['y'].astype(np.float64) Y_df.head()
之后按原流程执行后续代码即可得到预期输出。
内容的提问来源于stack exchange,提问作者user11222532
相关产品推荐
相关产品推荐

