已重塑数组仍遇ValueError:数组维度为4,估算器要求≤2求助
解决ValueError:Found array with dim 4. Estimator expected <=2的问题
咱们来拆解下你遇到的问题,根源主要出在日期数组的重塑逻辑和预测时的输入维度嵌套上,一步步来修正:
问题1:日期数组的reshape方向搞反了
看你第25行的代码:
dates = dates.reshape(dates.shape[1], -1)
你之前通过dates.append([int(date.day)])生成的dates是(n_samples, 1)的二维数组(比如20条数据就是(20,1)),但你用dates.shape[1](也就是1)作为第一个维度,reshape后变成了(1, 20)的行向量——而Scikit-learn的模型要求输入是列向量(n_samples, n_features),也就是每个样本占一行,特征占一列,这个维度错误是触发问题的核心。
问题2:预测时的输入嵌套层数太多
第36行你调用函数传的是[[28]],函数内部又写了svr_rbf.predict([[x]]),这就导致最终传入的是[[[28]]],变成了三维数组,叠加之前的日期维度问题,就出现了dim 4的报错。
修正后的完整代码及说明
1. 修正日期数组的reshape
把第25行改成正确的reshape方式,确保是(n_samples, 1)的形状:
# 替换原来错误的reshape代码 dates = dates.reshape(-1, 1)
reshape(-1,1)的意思是让Python自动计算样本数量,固定特征数为1,完美符合模型的输入要求。
2. 修正预测时的输入维度
调用predict_prices时直接传单个值28,不用嵌套多层列表:
# 替换原来的调用代码 predicted_price = predict_prices(dates, prices, 28)
函数内部的预测部分保持svr_rbf.predict([[x]])即可,这里是把单个值转换成模型需要的(1,1)形状数组。
完整修正后的代码
from datetime import datetime from iexfinance.stocks import Stock import pandas as pd import numpy as np from sklearn.svm import SVR import matplotlib.pyplot as plt start = datetime(2020, 1, 1) end = datetime(2020, 1, 29) def get_price_vol(symbol): get_info= get_historical_data(symbol, start, end, token='xyz', close_only=True, output_format='pandas' ) return get_info aapl_df = get_price_vol('aapl').reset_index() df = aapl_df[['date','close']].iloc[:-1] df_dates = df.loc[:,'date'] df_close = df.loc[:,'close'] dates = [] prices = [] for date in df_dates: dates.append([int(date.day)] ) for close_price in df_close: prices.append(float(close_price)) dates = np.array(dates) # 修正reshape方向 dates = dates.reshape(-1, 1) prices = np.array(prices) def predict_prices(dates, prices, x): svr_rbf = SVR(kernel='rbf', C=1e3, gamma=0.1) svr_rbf.fit(dates, prices) plt.scatter(dates, prices, color = 'black', label='Data') plt.plot(dates, svr_rbf.predict(dates), color = 'red', label='RBF model') plt.xlabel('Date') plt.ylabel('Price') plt.legend() # 加上图例让图表更清晰 plt.show() return svr_rbf.predict([[x]])[0] # 修正传入的参数维度 predicted_price = predict_prices(dates, prices, 28) print(predicted_price)
额外小提醒
- 你代码里第4行
from pandas import pandas是多余的,已经有import pandas as pd了,可以删掉; - 画图时加上
plt.legend()能让图例正常显示,方便观察模型拟合效果; - 注意确认
get_historical_data的调用是否符合当前iexfinance版本的API要求,如果这部分报错可以查下官方文档调整。
内容的提问来源于stack exchange,提问作者Jonathan
相关产品推荐
相关产品推荐

