如何在股票预测LSTM代码中使用自定义CSV替代雅虎财经数据源
适配CSV自定义数据的修改方案
一、需要调整的核心代码部分
- 删掉不必要的依赖:原代码中用于爬取雅虎财经数据的
pandas_datareader库可以移除,对应的导入语句删除即可。 - 替换训练数据读取逻辑:原
web.DataReader读取在线数据的代码,替换为pd.read_csv读取本地CSV文件。因为你提到当前CSV无表头、后续可添加,两种场景的写法如下:无表头场景:假设你的CSV第1列是日期、第5列是收盘价(和雅虎财经返回的列顺序一致:日期、开盘、最高、最低、收盘、调整后收盘、成交量),读取时手动指定列名即可
有表头场景:后续你在Excel添加了Date、Close这类表头后,简化读取逻辑即可 - 替换测试数据读取逻辑:原测试数据也是在线拉取的,同样替换为本地CSV读取,你可以选择把全量数据存在一个CSV里直接按时间切片,或者分开存训练、测试两个CSV文件。
- 时间筛选逻辑调整:如果你的CSV里包含了超出训练/测试时间范围的数据,可以用pandas的时间筛选功能切出对应区间,和原逻辑保持一致。
二、修改后的完整代码示例
import numpy as np import matplotlib.pyplot as plt import pandas as pd import datetime as dt from sklearn.preprocessing import MinMaxScaler from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense, Dropout, LSTM # ---------------------- 修改1:配置CSV路径和参数 ---------------------- CSV_PATH = "你的股票数据文件路径.csv" prediction_days = 60 # 保持和原逻辑一致的训练、测试时间范围 train_start = dt.datetime(2012,1,1) train_end = dt.datetime(2021,1,1) test_start = dt.datetime(2020,1,1) test_end = dt.datetime.now() # ---------------------- 修改2:读取本地CSV数据 ---------------------- # 无表头时用当前写法 df = pd.read_csv(CSV_PATH, # 手动指定列名,和你CSV实际列顺序对应即可,如果只需要收盘价可以只保留日期、收盘价两列 names=["Date", "Open", "High", "Low", "Close", "Adj Close", "Volume"], parse_dates=["Date"], # 自动解析日期列为时间格式 index_col="Date") # 将日期设为索引方便后续按时间筛选 # 后续你给CSV加了表头后,替换为下面的写法即可 # df = pd.read_csv(CSV_PATH, parse_dates=["Date"], index_col="Date") # 切出训练区间数据 data = df.loc[train_start:train_end] scaler = MinMaxScaler(feature_range=(0,1)) scaled_data = scaler.fit_transform(data['Close'].values.reshape(-1, 1)) x_train = [] y_train = [] for x in range(prediction_days, len(scaled_data)): x_train.append(scaled_data[x-prediction_days:x, 0]) y_train.append(scaled_data[x, 0]) x_train, y_train = np.array(x_train), np.array(y_train) x_train = np.reshape(x_train, (x_train.shape[0], x_train.shape[1], 1)) # 模型构建部分不需要修改 model = Sequential() model.add(LSTM(units=50, return_sequences=True, input_shape=(x_train.shape[1], 1))) model.add(Dropout(0.2)) model.add(LSTM(units=50, return_sequences=True)) model.add(Dropout(0.2)) model.add(LSTM(units=50)) model.add(Dropout(0.2)) model.add(Dense(units=1)) # 预测下一日收盘价 model.compile(optimizer='adam', loss='mean_squared_error') model.fit(x_train, y_train, epochs=25, batch_size=32) # ---------------------- 修改3:切出测试区间数据 ---------------------- test_dataset = df.loc[test_start:test_end] actual_prices = test_dataset['Close'].values total_dataset = pd.concat((data['Close'], test_dataset['Close']), axis=0) model_inputs = total_dataset[len(total_dataset)-len(test_dataset)-prediction_days:].values model_inputs = model_inputs.reshape(-1,1) model_inputs = scaler.transform(model_inputs) # 测试集预测逻辑不需要修改 x_test = [] for x in range(prediction_days, len(model_inputs)): x_test.append(model_inputs[x-prediction_days:x, 0]) x_test = np.array(x_test) x_test = np.reshape(x_test,(x_test.shape[0], x_test.shape[1],1)) predicted_prices = model.predict(x_test) predicted_prices = scaler.inverse_transform(predicted_prices) # 绘图部分不需要修改 plt.plot(actual_prices, color="green", label="实际价格") plt.plot(predicted_prices, color="blue", label="预测价格") plt.title("GER40股价") plt.xlabel('时间') plt.ylabel('GER40价格') plt.legend() plt.show() # 下一日价格预测逻辑不需要修改 real_dataset = [model_inputs[len(model_inputs)+1-prediction_days:len(model_inputs+1), 0]] real_dataset = np.array(real_dataset) real_dataset = np.reshape(real_dataset, (real_dataset.shape[0], real_dataset.shape[1], 1)) prediction = model.predict(real_dataset) prediction = scaler.inverse_transform(prediction) print(f"下一日收盘价预测: {prediction}")
三、注意事项
- 如果CSV日期格式解析失败,可以在
pd.read_csv中新增date_parser参数指定你本地CSV的日期格式,比如date_parser=lambda x: pd.to_datetime(x, format="%Y/%m/%d") - 如果你的CSV只存了收盘价数据,调整
read_csv的names参数为["Date", "Close"]即可,和实际列顺序对应就行
内容的提问来源于stack exchange,提问作者Terry
相关产品推荐
相关产品推荐

