mirror of
https://github.com/silencesdg/mt5_python_ea_suite.git
synced 2026-08-02 21:57:43 +00:00
add ml
This commit is contained in:
+142
-168
@@ -1,16 +1,16 @@
|
||||
import sys
|
||||
import os
|
||||
import pandas as pd
|
||||
import json
|
||||
from datetime import datetime
|
||||
from tqdm import tqdm
|
||||
|
||||
# 添加项目根目录到路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__name__)))
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
from core.data_providers import BacktestDataProvider
|
||||
from core.risk import RiskController
|
||||
from execution.dynamic_weights import DynamicWeightManager # Still needed for strategy list
|
||||
from config import SYMBOL, TIMEFRAME, BACKTEST_COUNT, BACKTEST_START_DATE, BACKTEST_END_DATE, USE_DATE_RANGE, BACKTEST_CONFIG, INITIAL_CAPITAL, SIGNAL_THRESHOLDS, DEFAULT_WEIGHTS
|
||||
from config import SYMBOL, TIMEFRAME, BACKTEST_COUNT, BACKTEST_START_DATE, BACKTEST_END_DATE, USE_DATE_RANGE, BACKTEST_CONFIG, INITIAL_CAPITAL
|
||||
from logger import logger
|
||||
|
||||
# Import all strategy classes
|
||||
@@ -25,9 +25,8 @@ from strategies.turtle import TurtleStrategy
|
||||
from strategies.daily_breakout import DailyBreakoutStrategy
|
||||
from strategies.wave_theory import WaveTheoryStrategy
|
||||
|
||||
|
||||
def run_full_backtest():
|
||||
"""完整的、独立的、已重构的回测流程"""
|
||||
"""完整的、独立的、已重构的回测流程 (支持市场状态)"""
|
||||
# 1. 加载数据
|
||||
from core.utils import get_rates, initialize, shutdown
|
||||
initialize()
|
||||
@@ -41,189 +40,164 @@ def run_full_backtest():
|
||||
logger.error("未能获取历史数据,回测终止。")
|
||||
return
|
||||
|
||||
logger.info(f"实际获取数据量: {len(rates)} 条 (请求: {BACKTEST_COUNT} 条)")
|
||||
|
||||
# 如果实际获取的数据量超过配置,则截取配置的数据量
|
||||
if len(rates) > BACKTEST_COUNT:
|
||||
rates = rates[-BACKTEST_COUNT:] # 取最新的数据
|
||||
logger.info(f"截取数据量为: {len(rates)} 条")
|
||||
|
||||
df = pd.DataFrame(rates)
|
||||
df.set_index(pd.to_datetime(df['time'], unit='s'), inplace=True)
|
||||
|
||||
# 2. 初始化组件 (data_provider, risk_controller)
|
||||
data_provider = BacktestDataProvider(df, initial_equity=INITIAL_CAPITAL)
|
||||
risk_controller = RiskController(data_provider, trade_direction=BACKTEST_CONFIG['trade_direction'])
|
||||
# weight_manager is not directly used for signal generation in this new architecture,
|
||||
# but its strategy_blueprints can be used to instantiate strategies.
|
||||
|
||||
# 3. 策略实例化和信号预生成 (NEW STAGE 1)
|
||||
logger.info("开始预生成所有策略信号...")
|
||||
|
||||
# Define strategy blueprints directly here or get from DynamicWeightManager if it's refactored
|
||||
# For now, let's define them directly as DynamicWeightManager is for real-time dynamic weights.
|
||||
strategy_blueprints = {
|
||||
'ma_cross': (MACrossStrategy, {}),
|
||||
'rsi': (RSIStrategy, {}),
|
||||
'bollinger': (BollingerStrategy, {}),
|
||||
'mean_reversion': (MeanReversionStrategy, {}),
|
||||
'momentum_breakout': (MomentumBreakoutStrategy, {}),
|
||||
'macd': (MACDStrategy, {}),
|
||||
'kdj': (KDJStrategy, {}),
|
||||
'turtle': (TurtleStrategy, {}),
|
||||
'daily_breakout': (DailyBreakoutStrategy, {}),
|
||||
'wave_theory': (WaveTheoryStrategy, {})
|
||||
}
|
||||
# 2. 加载市场状态参数
|
||||
try:
|
||||
with open('regime_optimal_params.json', 'r') as f:
|
||||
regime_params = json.load(f)
|
||||
logger.info("成功加载市场状态参数文件。")
|
||||
except FileNotFoundError:
|
||||
logger.error("错误: regime_optimal_params.json 未找到。请先运行 regime_optimizer.py。")
|
||||
return
|
||||
|
||||
# 3. 市场状态分类
|
||||
logger.info("开始为整个数据集分类市场状态...")
|
||||
regime_detector = WaveTheoryStrategy(None, SYMBOL, TIMEFRAME)
|
||||
df_with_indicators = regime_detector._calculate_indicators(df.copy())
|
||||
|
||||
regimes = []
|
||||
adx_values = df_with_indicators['adx'].values
|
||||
ema_short_values = df_with_indicators['ema_short'].values
|
||||
ema_medium_values = df_with_indicators['ema_medium'].values
|
||||
ema_long_values = df_with_indicators['ema_long'].values
|
||||
adx_threshold = regime_detector.adx_threshold
|
||||
|
||||
for i in range(len(df_with_indicators)):
|
||||
if pd.isna(adx_values[i]) or pd.isna(ema_short_values[i]):
|
||||
regimes.append("Ranging")
|
||||
elif adx_values[i] < adx_threshold:
|
||||
regimes.append("Ranging")
|
||||
elif ema_short_values[i] > ema_medium_values[i] > ema_long_values[i]:
|
||||
regimes.append("Uptrend")
|
||||
else:
|
||||
regimes.append("Downtrend")
|
||||
df['regime'] = regimes
|
||||
logger.info("市场状态分类完成。")
|
||||
logger.info(df['regime'].value_counts())
|
||||
|
||||
# 4. 分状态生成信号
|
||||
logger.info("开始分状态生成所有策略信号...")
|
||||
all_signals_df = pd.DataFrame(index=df.index)
|
||||
|
||||
# Map strategy class names to config keys for weights
|
||||
strategy_name_to_config_key = {
|
||||
"MACrossStrategy": "ma_cross",
|
||||
"RSIStrategy": "rsi",
|
||||
"BollingerStrategy": "bollinger",
|
||||
"MeanReversionStrategy": "mean_reversion",
|
||||
"MomentumBreakoutStrategy": "momentum_breakout",
|
||||
"MACDStrategy": "macd",
|
||||
"KDJStrategy": "kdj",
|
||||
"TurtleStrategy": "turtle",
|
||||
"DailyBreakoutStrategy": "daily_breakout",
|
||||
"WaveTheoryStrategy": "wave_theory"
|
||||
strategy_classes = {
|
||||
"MACrossStrategy": MACrossStrategy,
|
||||
"RSIStrategy": RSIStrategy,
|
||||
"BollingerStrategy": BollingerStrategy,
|
||||
"MeanReversionStrategy": MeanReversionStrategy,
|
||||
"MomentumBreakoutStrategy": MomentumBreakoutStrategy,
|
||||
"MACDStrategy": MACDStrategy,
|
||||
"KDJStrategy": KDJStrategy,
|
||||
"TurtleStrategy": TurtleStrategy,
|
||||
"DailyBreakoutStrategy": DailyBreakoutStrategy,
|
||||
"WaveTheoryStrategy": WaveTheoryStrategy
|
||||
}
|
||||
|
||||
strategies_for_backtest = []
|
||||
for name, (strategy_class, params) in strategy_blueprints.items():
|
||||
# Pass None for data_provider, symbol, timeframe as run_backtest only needs df
|
||||
# But strategy __init__ expects them. So, pass dummy values.
|
||||
strategies_for_backtest.append(strategy_class(None, SYMBOL, TIMEFRAME, **params))
|
||||
|
||||
for strategy in tqdm(strategies_for_backtest, desc="Generating Strategy Signals"):
|
||||
if hasattr(strategy, 'run_backtest') and callable(getattr(strategy, 'run_backtest')):
|
||||
try:
|
||||
strategy_signals = strategy.run_backtest(df.copy()) # Pass a copy to avoid modifying original df
|
||||
if strategy_signals is not None:
|
||||
all_signals_df[strategy.name] = strategy_signals
|
||||
logger.info(f"策略 {strategy.name} 信号生成成功")
|
||||
else:
|
||||
logger.warning(f"策略 {strategy.name} 返回空信号")
|
||||
except Exception as e:
|
||||
logger.error(f"策略 {strategy.name} 执行失败: {str(e)}")
|
||||
# 为失败的策略创建全0信号序列
|
||||
all_signals_df[strategy.name] = pd.Series(0, index=df.index)
|
||||
continue
|
||||
else:
|
||||
logger.warning(f"策略 {strategy.name} 没有实现 'run_backtest' 方法,将被跳过。")
|
||||
|
||||
# 组合信号
|
||||
logger.info("开始组合策略信号...")
|
||||
combined_weighted_signals = pd.Series(0, index=df.index)
|
||||
if not all_signals_df.empty:
|
||||
weighted_signals_list = []
|
||||
for col_name in all_signals_df.columns:
|
||||
# Extract strategy class name from column name (e.g., 'MACrossStrategy')
|
||||
strategy_class_name = col_name
|
||||
config_key = strategy_name_to_config_key.get(strategy_class_name)
|
||||
|
||||
if config_key and config_key in DEFAULT_WEIGHTS:
|
||||
weight = DEFAULT_WEIGHTS[config_key]
|
||||
weighted_signals_list.append(all_signals_df[col_name] * weight)
|
||||
else:
|
||||
logger.warning(f"未找到策略 {strategy_class_name} 的默认权重,使用权重1.0。")
|
||||
weighted_signals_list.append(all_signals_df[col_name] * 1.0)
|
||||
|
||||
if weighted_signals_list:
|
||||
combined_weighted_signals = pd.concat(weighted_signals_list, axis=1).sum(axis=1)
|
||||
else:
|
||||
logger.warning("没有生成任何加权信号。")
|
||||
|
||||
# 应用阈值得到最终信号
|
||||
final_signals = pd.Series(0, index=df.index)
|
||||
buy_threshold = SIGNAL_THRESHOLDS.get("buy_threshold", 1.0)
|
||||
sell_threshold = SIGNAL_THRESHOLDS.get("sell_threshold", -1.0)
|
||||
|
||||
final_signals[combined_weighted_signals > buy_threshold] = 1
|
||||
final_signals[combined_weighted_signals < sell_threshold] = -1
|
||||
|
||||
logger.info("信号预生成和组合完成。")
|
||||
|
||||
# 4. 回测主循环 (NEW STAGE 2)
|
||||
logger.info(f"回测主循环开始... 数据量: {len(df)} 条")
|
||||
|
||||
# Reset data_provider's internal index to 0 for the loop
|
||||
data_provider.current_index = 0
|
||||
|
||||
for i in tqdm(range(len(df)), desc=f"Backtesting ({len(df)} bars)"):
|
||||
try:
|
||||
# Get current price from the data_provider (which uses its internal index)
|
||||
current_price = data_provider.get_current_price(SYMBOL)
|
||||
if not current_price:
|
||||
logger.warning(f"无法获取当前价格在索引 {i},跳过。")
|
||||
continue
|
||||
|
||||
# Get the pre-generated signal for the current bar
|
||||
current_signal = final_signals.iloc[i]
|
||||
|
||||
direction = None
|
||||
if current_signal == 1:
|
||||
direction = "buy"
|
||||
elif current_signal == -1:
|
||||
direction = "sell"
|
||||
|
||||
if direction:
|
||||
risk_controller.process_trading_signal(direction, current_price, abs(current_signal))
|
||||
|
||||
# Monitor and update positions
|
||||
risk_controller.monitor_positions(current_price)
|
||||
|
||||
# Advance data provider to the next time step
|
||||
data_provider.tick()
|
||||
except Exception as e:
|
||||
logger.error(f"回测循环中在索引 {i} 发生错误: {str(e)}")
|
||||
# 继续下一个bar,不中断整个回测
|
||||
for regime in ['Uptrend', 'Downtrend', 'Ranging']:
|
||||
logger.info(f"--- 为 {regime} 状态生成信号 ---")
|
||||
regime_df = df[df['regime'] == regime]
|
||||
if regime_df.empty:
|
||||
logger.info(f"{regime} 状态没有数据,跳过。")
|
||||
continue
|
||||
|
||||
# 5. 结束和报告 (Keep as is)
|
||||
params = regime_params.get(regime, {}).get('best_parameters', {})
|
||||
if not params:
|
||||
logger.warning(f"未找到 {regime} 的参数,将无法为该状态生成信号。")
|
||||
continue
|
||||
|
||||
for strat_name, strat_class in strategy_classes.items():
|
||||
instance = strat_class(None, SYMBOL, TIMEFRAME)
|
||||
strat_params = {}
|
||||
prefix_map = {
|
||||
'MACrossStrategy': 'ma_cross_',
|
||||
'RSIStrategy': 'rsi_',
|
||||
'BollingerStrategy': 'bollinger_',
|
||||
'MACDStrategy': 'macd_',
|
||||
'MeanReversionStrategy': 'mean_reversion_',
|
||||
'MomentumBreakoutStrategy': 'momentum_breakout_',
|
||||
'KDJStrategy': 'kdj_',
|
||||
'TurtleStrategy': 'turtle_',
|
||||
'DailyBreakoutStrategy': 'daily_breakout_',
|
||||
'WaveTheoryStrategy': 'wave_'
|
||||
}
|
||||
prefix = prefix_map[strat_name]
|
||||
for p_name, p_val in params.items():
|
||||
if p_name.startswith(prefix):
|
||||
param_name = p_name.replace(prefix, '')
|
||||
if strat_name == 'WaveTheoryStrategy' and p_name == 'wave_period':
|
||||
param_name = 'wave_period'
|
||||
strat_params[param_name] = p_val
|
||||
|
||||
instance.set_params(strat_params)
|
||||
|
||||
try:
|
||||
signals = instance.run_backtest(regime_df.copy())
|
||||
if signals is not None:
|
||||
all_signals_df.loc[regime_df.index, strat_name] = signals
|
||||
except Exception as e:
|
||||
logger.error(f"策略 {strat_name} 在 {regime} 状态下执行失败: {e}")
|
||||
|
||||
all_signals_df.fillna(0, inplace=True)
|
||||
|
||||
# 5. 组合信号
|
||||
logger.info("开始组合策略信号...")
|
||||
final_signals = pd.Series(0.0, index=df.index)
|
||||
for regime in ['Uptrend', 'Downtrend', 'Ranging']:
|
||||
regime_df_indices = df[df['regime'] == regime].index
|
||||
if regime_df_indices.empty: continue
|
||||
|
||||
params = regime_params.get(regime, {}).get('best_parameters', {})
|
||||
weights = {k.replace('weight_', ''): v for k, v in params.items() if k.startswith('weight_')}
|
||||
buy_threshold = params.get('buy_threshold', 1.5)
|
||||
sell_threshold = params.get('sell_threshold', -1.5)
|
||||
|
||||
regime_signals = all_signals_df.loc[regime_df_indices]
|
||||
weighted_sum = pd.Series(0.0, index=regime_signals.index)
|
||||
for strat_name, weight in weights.items():
|
||||
if strat_name in regime_signals.columns:
|
||||
weighted_sum += regime_signals[strat_name] * weight
|
||||
|
||||
# BUG FIX: Use indices to avoid alignment errors
|
||||
buy_indices = weighted_sum[weighted_sum > buy_threshold].index
|
||||
sell_indices = weighted_sum[weighted_sum < sell_threshold].index
|
||||
|
||||
final_signals.loc[buy_indices] = 1
|
||||
final_signals.loc[sell_indices] = -1
|
||||
|
||||
# 6. 回测主循环
|
||||
logger.info("回测主循环开始...")
|
||||
data_provider = BacktestDataProvider(df, initial_equity=INITIAL_CAPITAL)
|
||||
risk_controller = RiskController(data_provider, trade_direction=BACKTEST_CONFIG['trade_direction'])
|
||||
|
||||
for i in tqdm(range(len(df)), desc="Backtesting"):
|
||||
current_price = data_provider.get_current_price(SYMBOL)
|
||||
if not current_price: continue
|
||||
|
||||
current_signal = final_signals.iloc[i]
|
||||
direction = "buy" if current_signal == 1 else "sell" if current_signal == -1 else None
|
||||
|
||||
if direction:
|
||||
risk_controller.process_trading_signal(direction, current_price, abs(current_signal))
|
||||
|
||||
risk_controller.monitor_positions(current_price)
|
||||
data_provider.tick()
|
||||
|
||||
# 7. 结束和报告
|
||||
logger.info("回测完成,生成性能报告...")
|
||||
summary = risk_controller.position_manager.get_trade_summary()
|
||||
|
||||
# 打印报告
|
||||
logger.info("=" * 80)
|
||||
logger.info("回测性能报告")
|
||||
logger.info(f"总交易次数: {summary.get('total_trades', 0)}")
|
||||
logger.info(f"胜率: {summary.get('win_rate', 0):.2f}%")
|
||||
logger.info(f"总盈亏: ${summary.get('total_profit_loss', 0):.2f}")
|
||||
logger.info("=" * 80)
|
||||
|
||||
# 保存交易记录
|
||||
try:
|
||||
risk_controller.position_manager.save_trade_history("backtest_trades")
|
||||
logger.info("交易记录保存成功")
|
||||
except Exception as e:
|
||||
logger.error(f"保存交易记录失败: {str(e)}")
|
||||
# 尝试手动保存
|
||||
try:
|
||||
import json
|
||||
trades = risk_controller.position_manager.closed_trades
|
||||
if trades:
|
||||
with open("backtest_trades_manual.json", 'w', encoding='utf-8') as f:
|
||||
json.dump(trades, f, ensure_ascii=False, indent=2, default=str)
|
||||
logger.info("手动保存交易记录到 backtest_trades_manual.json")
|
||||
except Exception as e2:
|
||||
logger.error(f"手动保存交易记录也失败: {str(e2)}")
|
||||
|
||||
risk_controller.position_manager.save_trade_history("backtest_trades")
|
||||
|
||||
def main():
|
||||
print("=" * 60)
|
||||
print("MetaTrader 5 智能交易系统 - 回测")
|
||||
print("MetaTrader 5 智能交易系统 - 回测 (市场状态模式)")
|
||||
print("=" * 60)
|
||||
|
||||
try:
|
||||
run_full_backtest()
|
||||
print("\n回测完成!")
|
||||
except Exception as e:
|
||||
import traceback
|
||||
print(f"\n回测出错: {e}")
|
||||
traceback.print_exc()
|
||||
run_full_backtest()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
Reference in New Issue
Block a user