You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 15:02:01