如何优化运行超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
相关产品推荐
相关产品推荐

