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

如何优化运行超24小时的RSI背离策略参数寻优Python代码?

RSI背离策略参数寻优代码优化求助

我有一段用于寻找特定时段内RSI背离交易策略最优参数的Python代码,目前已运行24小时仍未输出结果,无法预估剩余耗时。我并非技术专家,不清楚如何修改代码,希望有人能帮忙优化这段代码,并愿意学习相关方法。

原代码

import pandas as pd
import numpy as np
import ta

def load_data(file_path, start_date, end_date):
    """
    Loads data for the specified symbol and date range from a CSV file
    """
    df = pd.read_csv(file_path)
    if 'Date' not in df.columns:
        df['Date'] = pd.to_datetime(df.index)
    df['Date'] = pd.to_datetime(df['Date'])
    df = df.set_index('Date')
    df = df[(df.index >= start_date) & (df.index <= end_date)]
    return df

def calc_rsi(df, n):
    """
    Calculates the relative strength index (RSI) for the given dataframe and window size
    """
    delta = df["Close"].diff()
    gain = delta.where(delta > 0, 0)
    loss = abs(delta.where(delta < 0, 0))
    avg_gain = gain.rolling(window=n).mean()
    avg_loss = loss.rolling(window=n).mean()
    rs = avg_gain / avg_loss
    rsi = 100 - (100 / (1 + rs))
    return rsi

def calc_pivot_point(df, pivot_point_type, pivot_point_n):
    """
    Calculates the pivot point for the given dataframe and pivot point type
    """
    if pivot_point_type == "Close":
        pivot_point = df["Close"].rolling(window=pivot_point_n).mean()
    elif pivot_point_type == "High/Low":
        pivot_point = (df["High"].rolling(window=pivot_point_n).mean() + df["Low"].rolling(window=pivot_point_n).mean()) / 2
    else:
        raise ValueError("Invalid pivot point type")
    return pivot_point

def calc_divergence(df, rsi, pivot_point, divergence_type, max_pivot_point, max_bars_to_check):
    """
    Calculates the divergence for the given dataframe and parameters
    """
    if divergence_type == "Regular":
        pivot_point_delta = pivot_point.diff()
        pivot_point_delta_sign = pivot_point_delta.where(pivot_point_delta > 0, -1)
        pivot_point_delta_sign[pivot_point_delta_sign > 0] = 1
        rsi_delta = rsi.diff()
        rsi_delta_sign = rsi_delta.where(rsi_delta > 0, -1)
        rsi_delta_sign[rsi_delta_sign > 0] = 1
        divergence = pivot_point_delta_sign * rsi_delta_sign
        divergence[divergence < 0] = -1
        divergence = divergence.rolling(window=max_pivot_point).sum()
        divergence = divergence.rolling(window=max_bars_to_check).sum()
        divergence = divergence.where(divergence > 0, 0)
        divergence[divergence < 0] = -1
    else:
        raise ValueError("Invalid divergence type")
    return divergence

def backtest(df, rsi_period, pivot_point_type, pivot_point_n, divergence_type, max_pivot_point, max_bars_to_check, trailing_stop, starting_capital):
    """
    Backtests the strategy for the given dataframe and parameters
    """
    rsi = calc_rsi(df, rsi_period)
    pivot_point = calc_pivot_point(df, pivot_point_type, pivot_point_n)
    divergence = calc_divergence(df, rsi, pivot_point, divergence_type, max_pivot_point, max_bars_to_check)
    positions = pd.DataFrame(index=df.index, columns=["Position", "Stop Loss"])
    positions["Position"] = 0.0
    positions["Stop Loss"] = 0.0
    capital = starting_capital
    for i, row in enumerate(df.iterrows()):
        date = row[0]
        close = row[1]["Close"]
        rsi_val = rsi.loc[date]
        pivot_val = pivot_point.loc[date]
        divergence_val = divergence.loc[date]
        if divergence_val > 0 and positions.loc[date]["Position"] == 0:
            positions.at[date, "Position"] = capital / close
            positions.at[date, "Stop Loss"] = close * (1 - trailing_stop)
        elif divergence_val < 0 and positions.loc[date]["Position"] > 0:
            capital = positions.loc[date]["Position"] * close
            positions.at[date, "Position"] = 0.0
            positions.at[date, "Stop Loss"] = 0.0
        elif close < positions.loc[date]["Stop Loss"] and positions.loc[date]["Position"] > 0:
            capital = positions.loc[date]["Position"] * close
            positions.at[date, "Position"] = 0.0
            positions.at[date, "Stop Loss"] = 0.0
    return capital

def find_best_iteration(df, start_rsi_period, end_rsi_period, pivot_point_types, start_pivot_point_n, end_pivot_point_n, divergence_types, start_max_pivot_point, end_max_pivot_point, start_max_bars_to_check, end_max_bars_to_check, start_trailing_stop, end_trailing_stop, starting_capital):
    """
    Finds the best iteration for the given parameters
    """
    best_result = 0.0
    best_params = None
    for rsi_period in range(start_rsi_period, end_rsi_period + 1):
        for pivot_point_type in pivot_point_types:
            for pivot_point_n in range(start_pivot_point_n, end_pivot_point_n + 1):
                for divergence_type in divergence_types:
                     for max_pivot_point in range(start_max_pivot_point, end_max_pivot_point + 1):
                        for max_bars_to_check in range(start_max_bars_to_check, end_max_bars_to_check + 1):
                            for trailing_stop in np.arange(start_trailing_stop, end_trailing_stop + 0.01, 0.01):
                                result = backtest(df, rsi_period, pivot_point_type, pivot_point_n, divergence_type, max_pivot_point, max_bars_to_check, trailing_stop, starting_capital)
                                if result > best_result:
                                    best_result = result
                                    best_params = (rsi_period, pivot_point_type, pivot_point_n, divergence_type, max_pivot_point, max_bars_to_check, trailing_stop)
    return best_result, best_params

# Define the parameters
file_path = 'C:\\Users\\The Death\\Downloads\\Binance_BTCUSDT_spot.csv'
start_date = "2020-03-16"
end_date = "2021-04-12"
df = load_data(file_path, start_date, end_date)

def load_data(start_date, end_date):
    # Your code to load the data for the specified date range
    # ...
    return df

# Define the parameters for the backtesting
start_rsi_period = 1
end_rsi_period = 30
pivot_point_types = ["Close", "High/Low"]
start_pivot_point_n = 1
end_pivot_point_n = 50
divergence_types = ["Regular"]
start_max_pivot_point = 1
end_max_pivot_point = 20
start_max_bars_to_check = 30
end_max_bars_to_check = 200
start_trailing_stop = 0.01
end_trailing_stop = 0.5
starting_capital = 10000

# Run the backtesting
df = load_data(start_date, end_date)
best_result, best_params = find_best_iteration(df, start_rsi_period, end_rsi_period, pivot_point_types, start_pivot_point_n, end_pivot_point_n, divergence_types, start_max_pivot_point, end_max_pivot_point, start_max_bars_to_check, end_max_bars_to_check, start_trailing_stop, end_trailing_stop, starting_capital)


# Print the results
print("Best result: ", best_result)
print("Best parameters: ", best_params)

优化方案

1. 先修正代码错误

代码里重复定义了load_data函数,后面的空函数覆盖了前面读取CSV的逻辑,直接删掉后面的空load_data函数,保证数据正常加载。

2. 大幅减少冗余计算

把和止损参数无关的指标计算(RSI、枢轴点、背离信号)移到外层循环,缓存结果,避免每次遍历止损都重复计算:

def find_best_iteration(df, start_rsi_period, end_rsi_period, pivot_point_types, start_pivot_point_n, end_pivot_point_n, divergence_types, start_max_pivot_point, end_max_pivot_point, start_max_bars_to_check, end_max_bars_to_check, start_trailing_stop, end_trailing_stop, starting_capital):
    best_result = 0.0
    best_params = None
    # 先遍历所有非止损参数,缓存指标
    for rsi_period in range(start_rsi_period, end_rsi_period + 1):
        rsi = calc_rsi(df, rsi_period)
        for pivot_point_type in pivot_point_types:
            for pivot_point_n in range(start_pivot_point_n, end_pivot_point_n + 1):
                pivot_point = calc_pivot_point(df, pivot_point_type, pivot_point_n)
                for divergence_type in divergence_types:
                    for max_pivot_point in range(start_max_pivot_point, end_max_pivot_point + 1):
                        for max_bars_to_check in range(start_max_bars_to_check, end_max_bars_to_check + 1):
                            divergence = calc_divergence(df, rsi, pivot_point, divergence_type, max_pivot_point, max_bars_to_check)
                            # 现在只遍历止损参数,用缓存好的指标回测
                            for trailing_stop in np.arange(start_trailing_stop, end_trailing_stop + 0.01, 0.01):
                                result = backtest_with_cached(df, rsi, pivot_point, divergence, trailing_stop, starting_capital)
                                if result > best_result:
                                    best_result = result
                                    best_params = (rsi_period, pivot_point_type, pivot_point_n, divergence_type, max_pivot_point, max_bars_to_check, trailing_stop)
    return best_result, best_params

# 新增用缓存指标的回测函数
def backtest_with_cached(df, rsi, pivot_point, divergence, trailing_stop, starting_capital):
    positions = pd.DataFrame(index=df.index, columns=["Position", "Stop Loss"])
    positions["Position"] = 0.0
    positions["Stop Loss"] = 0.0
    capital = starting_capital
    # 后续逻辑和原backtest一致,只是不用再计算指标
    for i, row in enumerate(df.iterrows()):
        date = row[0]
        close = row[1]["Close"]
        divergence_val = divergence.loc[date]
        if divergence_val > 0 and positions.loc[date]["Position"] == 0:
            positions.at[date, "Position"] = capital / close
            positions.at[date, "Stop Loss"] = close * (1 - trailing_stop)
        elif divergence_val < 0 and positions.loc[date]["Position"] > 0:
            capital = positions.loc[date]["Position"] * close
            positions.at[date, "Position"] = 0.0
            positions.at[date, "Stop Loss"] = 0.0
        elif close < positions.loc[date]["Stop Loss"] and positions.loc[date]["Position"] > 0:
            capital = positions.loc[date]["Position"] * close
            positions.at[date, "Position"] = 0.0
            positions.at[date, "Stop Loss"] = 0.0
    return capital

3. 优化回测循环(关键提速)

替换df.iterrows()为向量化操作,彻底删除逐行循环:

def backtest_with_cached(df, rsi, pivot_point, divergence, trailing_stop, starting_capital):
    # 初始化临时列
    temp_df = df.copy()
    temp_df['position'] = 0.0
    temp_df['stop_loss'] = 0.0
    temp_df['capital'] = starting_capital

    # 生成买入信号:背离>0且当前无仓位
    buy_signal = (divergence > 0) & (temp_df['position'].shift(1, fill_value=0) == 0)
    # 生成卖出信号:背离<0且当前有仓位,或者价格跌破止损
    sell_signal = ((divergence < 0) | (temp_df['Close'] < temp_df['stop_loss'].shift(1))) & (temp_df['position'].shift(1) > 0)

    # 执行买入操作
    temp_df.loc[buy_signal, 'position'] = temp_df['capital'].loc[buy_signal] / temp_df['Close'].loc[buy_signal]
    temp_df.loc[buy_signal, 'stop_loss'] = temp_df['Close'].loc[buy_signal] * (1 - trailing_stop)
    
    # 延续仓位和止损状态
    temp_df['position'] = temp_df['position'].fillna(method='ffill')
    temp_df['stop_loss'] = temp_df['stop_loss'].fillna(method='ffill')
    
    # 执行卖出操作
    temp_df.loc[sell_signal, 'capital'] = temp_df['position'].loc[sell_signal] * temp_df['Close'].loc[sell_signal]
    temp_df.loc[sell_signal, 'position'] = 0
    temp_df.loc[sell_signal, 'stop_loss'] = 0

    # 处理最终未平仓仓位
    final_position = temp_df['position'].iloc[-1]
    if final_position > 0:
        final_capital = final_position * temp_df['Close'].iloc[-1]
    else:
        final_capital = temp_df['capital'].dropna().iloc[-1] if not temp_df['capital'].dropna().empty else starting_capital
    return final_capital

4. 缩小参数搜索范围

原参数范围导致总循环次数超过5000万次,完全不现实,先缩小到合理范围:

start_rsi_period = 7
end_rsi_period = 14  # RSI常用周期7-14
start_pivot_point_n = 5
end_pivot_point_n = 20  # 枢轴点周期5-20
start_max_bars_to_check = 50
end_max_bars_to_check = 100  # 背离检查范围50-100
start_trailing_stop = 0.05
end_trailing_stop = 0.2
trailing_stop_step = 0.05  # 止损步长改成0.05,减少到4个值

5. 并行计算加速

用joblib实现并行遍历,利用多核CPU:

from joblib import Parallel, delayed

def find_best_iteration(df, start_rsi_period, end_rsi_period, pivot_point_types, start_pivot_point_n, end_pivot_point_n, divergence_types, start_max_pivot_point, end_max_pivot_point, start_max_bars_to_check, end_max_bars_to_check, start_trailing_stop, end_trailing_stop, starting_capital):
    # 生成所有参数组合及对应的缓存指标
    param_list = []
    for rsi_period in range(start_rsi_period, end_rsi_period + 1):
        rsi = calc_rsi(df, rsi_period)
        for pivot_point_type in pivot_point_types:
            for pivot_point_n in range(start_pivot_point_n, end_pivot_point_n + 1):
                pivot_point = calc_pivot_point(df, pivot_point_type, pivot_point_n)
                for divergence_type in divergence_types:
                    for max_pivot_point in range(start_max_pivot_point, end_max_pivot_point + 1):
                        for max_bars_to_check in range(start_max_bars_to_check, end_max_bars_to_check + 1):
                            divergence = calc_divergence(df, rsi, pivot_point, divergence_type, max_pivot_point, max_bars_to_check)
                            for trailing_stop in np.arange(start_trailing_stop, end_trailing_stop + 0.05, 0.05):
                                # 保存参数元组和缓存指标
                                param_tuple = (rsi_period, pivot_point_type, pivot_point_n, divergence_type, max_pivot_point, max_bars_to_check, trailing_stop)
                                param_list.append((df, rsi, pivot_point, divergence, trailing_stop, starting_capital, param_tuple))
    
    # 并行计算所有组合
    results = Parallel(n_jobs=-1)(delayed(lambda x: (backtest_with_cached(*x[:6]), x[6]))(params) for params in param_list)
    
    # 筛选最优结果
    best_result, best_params = max(results, key=lambda x: x[0])
    return best_result, best_params

学习建议

  • 优先掌握Pandas向量化操作:避免逐行循环是提速核心,shift()、布尔索引、fillna()等方法能替代大部分循环逻辑。
  • 理解参数寻优的复杂度:参数组合是指数级增长的,必须基于交易常识缩小范围,或使用网格搜索、遗传算法等高效寻优方法。
  • 入门并行计算:用joblib或multiprocessing可以快速利用多核CPU,适合这类CPU密集型的参数遍历任务。

内容的提问来源于stack exchange,提问作者Danky Kang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 13:35:19