Python RSI策略回测DataFrame Shape及KeyError问题求助
RSI指标SPY交易回测代码报错求助
我运行基于RSI指标的SPY交易回测Python代码时,遇到了KeyError: 'Buy_Signal'及后续的ValueError: Expected a 1D array, got an array with shape (501, 501)错误,尝试过ChatGPT给出的方案但无法解决,恳请协助排查问题。
报错信息
--------------------------------------------------------------------------- KeyError Traceback (most recent call last) /usr/local/lib/python3.10/dist-packages/pandas/core/indexes/base.py in get_loc(self, key) 3804 try: -> 3805 return self._engine.get_loc(casted_key) 3806 except KeyError as err: index.pyx in pandas._libs.index.IndexEngine.get_loc() index.pyx in pandas._libs.index.IndexEngine.get_loc() pandas/_libs/hashtable_class_helper.pxi in pandas._libs.hashtable.PyObjectHashTable.get_item() pandas/_libs/hashtable_class_helper.pxi in pandas._libs.hashtable.PyObjectHashTable.get_item() KeyError: 'Buy_Signal' The above exception was the direct cause of the following exception: KeyError Traceback (most recent call last) 10 frames KeyError: 'Buy_Signal' During handling of the above exception, another exception occurred: ValueError Traceback (most recent call last) /usr/local/lib/python3.10/dist-packages/pandas/core/internals/managers.py in insert(self, loc, item, value, refs) 1368 value = value.T 1369 if len(value) > 1: -> 1370 raise ValueError( 1371 f"Expected a 1D array, got an array with shape {value.T.shape}" 1372 ) ValueError: Expected a 1D array, got an array with shape (501, 501)
原代码
import pandas as pd import numpy as np import yfinance as yf import matplotlib.pyplot as plt # 获取SPY历史数据 def fetch_data(symbol, start_date, end_date): data = yf.download(symbol, start=start_date, end=end_date) return data # 计算RSI指标 def calculate_rsi(data, window=14): delta = data['Close'].diff() gain = (delta.where(delta > 0, 0)).rolling(window=window).mean() loss = (-delta.where(delta < 0, 0)).rolling(window=window).mean() rs = gain / loss rsi = 100 - (100 / (1 + rs)) data['RSI'] = rsi return data # 根据RSI阈值生成买卖信号 def create_signals(data, buy_threshold=30, sell_threshold=70): data['Buy_Signal'] = np.where(data['RSI'] < buy_threshold, data['Close'], np.nan) data['Sell_Signal'] = np.where(data['RSI'] > sell_threshold, data['Close'], np.nan) return data class RSIStrategy: def __init__(self, initial_investment): self.initial_investment = initial_investment self.portfolio_value = initial_investment self.shares_owned = 0 def execute_trade(self, price, signal_type): if signal_type == 'buy': if self.shares_owned == 0: # 仅空仓时买入 self.shares_owned = self.portfolio_value / price self.portfolio_value = 0 # 全部资金买入股票 return "Bought" elif signal_type == 'sell': if self.shares_owned > 0: # 仅持仓时卖出 self.portfolio_value = self.shares_owned * price self.shares_owned = 0 # 全部股票卖出 return "Sold" return "No action" def simulate_trading(data, initial_investment): strategy = RSIStrategy(initial_investment) trades = [] for index, row in data.iterrows(): if pd.notna(row['Buy_Signal']): action = strategy.execute_trade(row['Close'], 'buy') trades.append((index, action, row['Close'])) elif pd.notna(row['Sell_Signal']): action = strategy.execute_trade(row['Close'], 'sell') trades.append((index, action, row['Close'])) # 计算最终组合价值 final_value = strategy.portfolio_value + (strategy.shares_owned * data['Close'].iloc[-1]) return final_value, trades # 定义参数并运行回测 start_date = '2022-01-01' end_date = '2023-12-31' symbol = 'SPY' initial_investment = 10000 # 获取并预处理数据 spy_data = fetch_data(symbol, start_date, end_date) spy_data = calculate_rsi(spy_data) spy_data = create_signals(spy_data) # 模拟交易 final_value, trades = simulate_trading(spy_data, initial_investment) # 输出结果 print(f"Final Portfolio Value: ${final_value:.2f}") print("Trades executed:") for trade in trades: print(trade) # 绘制结果图 plt.figure(figsize=(14, 7)) plt.plot(spy_data['Close'], label='SPY Close Price', alpha=0.5) plt.scatter(spy_data.index, spy_data['Buy_Signal'], marker='^', color='g', label='Buy Signal', alpha=1) plt.scatter(spy_data.index, spy_data['Sell_Signal'], marker='v', color='r', label='Sell Signal', alpha=1) plt.title('SPY Trading Strategy with RSI Signals') plt.xlabel('Date') plt.ylabel('Price') plt.legend() plt.show()
已尝试的方案
ChatGPT提示检查布尔值的位运算,但self.shares_owned是单个变量而非DataFrame列,该方案无法解决问题。
内容的提问来源于stack exchange,提问作者user28102875
相关产品推荐
相关产品推荐

