feat: 适应度门槛95+swing_point策略+多项改进

- 新增适应度门槛: min_backtest_fitness=95, 适应度<95暂停开仓
- 新增 SwingPointRetest 策略替代 Turtle
- 新增 monday_reset.py 周重置脚本
- exit_rules: 拖尾止损相对回撤模式
- market_state: 趋势检测优化
- position: 一票制并发锁+合约规格缓存
- optimize: Optuna 替代 DEAP 遗传算法
- realtime_trader: 适应度门槛+同向递增
- weights: 动态权重管理
- cron_optimize: PYTHONPATH 修复
- .gitignore: 排除生成文件
This commit is contained in:
silencesdg
2026-05-21 20:26:55 +08:00
parent 218efef538
commit e1691c3c41
14 changed files with 699 additions and 449 deletions
+17 -10
View File
@@ -1,10 +1,17 @@
/GEMINI.md
/logs
/myenv
__pycache__/
*.json
*.csv
.claude/
.claude/**
.claude/settings.local.json
.claude/
1|/GEMINI.md
2|/logs
3|/myenv
4|__pycache__/
5|*.json
6|*.csv
7|.claude/
8|.claude/**
9|.claude/settings.local.json
10|.claude/
config_backups/
optimization_report_*.txt
optimization_results_*.json
.ea_pid
.ea.lock
.hermes/
realtime_trades_*
+114 -94
View File
@@ -1,4 +1,10 @@
# MT5 代理服务配置
import importlib as _importlib, sys as _sys
def reload():
"""★ 热重载配置 — 运行时调用,无需重启 EA"""
_importlib.reload(_sys.modules[__name__])
SERVER_HOST = "0.0.0.0"
SERVER_PORT = 5555
@@ -29,7 +35,7 @@ USE_DATE_RANGE = False # 设置为False可强制使用数据量模式
# 兼容性配置 (如果日期配置不可用,则使用数据量)
BACKTEST_COUNT = 30000 # 回测数据量
OPTIMIZER_COUNT = 50000 # 优化器数据量
OPTIMIZER_COUNT = 6000 # 优化器数据量(约4天M1,更快响应行情)
RISK_CONFIG_CONST = {
'enable_time_based_exit': False # 关闭超时平仓,让止盈/止损/跟踪止损接管
@@ -62,8 +68,8 @@ BACKTEST_CONFIG = {
REALTIME_CONFIG = {
"update_interval": 5, # 更新间隔(秒)
"daily_reset_time": "00:00", # 每日重置时间
"max_long_positions": 1, # 同方向只持一单,避免重复开仓
"max_short_positions": 1, # 同方向只持一
"max_long_positions": 10, # 同方向最多10单
"max_short_positions": 10, # 同方向最多10
"min_trade_interval": 0, # 最小交易间隔(分钟),0表示无限制
"enable_auto_trading": True, # 是否启用自动交易
"dry_run": False, # 是否为模拟运行(不实际下单)
@@ -104,8 +110,8 @@ DATA_CONFIG = {
# 遗传算法优化器配置
GENETIC_OPTIMIZER_CONFIG = {
# 算法参数
"population_size": 50, # 种群大小
"generations": 10, # 进化代数
"population_size": 20, # 种群大小(精简化,6000bar快速收敛)
"generations": 6, # 进化代数
"crossover_probability": 0.7, # 交叉概率
"mutation_probability": 0.3, # 变异概率
@@ -135,57 +141,68 @@ GENETIC_OPTIMIZER_CONFIG = {
# 信号阈值配置(优化器结果 2026-05-10,5万根M1数据)
SIGNAL_THRESHOLDS = {
"buy_threshold": 1.0315663601353515,
"sell_threshold": -2.312867704316165
"buy_threshold": 2.087211936631902,
"sell_threshold": -0.6594468309213586
}
# 风险管理参数(优化器结果 2026-05-105万根M1数据
# 风险管理参数(手动设定,优化器不宜优化 — 保证金% 基础值,杠杆自适应缩放
RISK_CONFIG = {
# ★ 以下百分比均为保证金%开仓成本%),非金价涨跌%
# 杠杆统一 2000x0.01手保证金≈$2.33
# -50%保证金 = -$1.17 ≈ 1.2点(很紧,但2000x下正常
"risk_leverage": 0, # 0=自动读MT5杠杆(2000x)
"stop_loss_pct": -0.50, # -50% 保证金
"profit_retracement_pct": 0.35, # 35% 保证金(拖尾回撤容忍
"min_profit_for_trailing": 0.70, # +70% 保证金后激活拖尾
"take_profit_pct": 1.00, # +100% 保证金
# ★ 以下百分比均为保证金%100x基准),实盘按 (MT5杠杆/100) 缩放
# 2000x下:-0.70 → -14×杠杆 → -1400%保证金 = -$32 ≈ 32点
"risk_leverage": 100, # 100x基准计算保证金%(不改实际杠杆
"stop_loss_pct": -0.10, # -10% 保证金 ←手动设定,不进优化器
"profit_retracement_pct": 0.03, # ★ 3% 回撤容忍(相对!3%×峰值利润) ←手动设定
"retracement_mode": "relative", # ★ 回撤模式从绝对值改为相对(利润的3%
"min_profit_for_trailing": 0.10, # +10% 后激活拖尾 ←手动设定
"take_profit_pct": 0.15, # +15% 止盈 ←手动设定
"max_daily_loss": -0.3,
"max_holding_minutes": 133,
"min_profit_for_time_exit": 0.010,
"max_holding_minutes": 0, # 0=不启用时间平仓
"min_profit_for_time_exit": 0.05,
"cooldown_bars": 30,
# ★ 同向加仓门槛递增:第N单需要的信号 = 基础阈值 × α^(N-1)
"entry_escalation_alpha": 1.2, # 递增系数(1.2=每单信号强20%)
# ★ 波动率自适应拖尾:激活点数 = 基础点数 × (当前ATR / 长周期ATR)
"vol_adaptive_trailing": True, # 启动波动率自适应拖尾
"trailing_atr_period": 14, # ATR短周期
"trailing_atr_baseline": 100, # ATR长周期(基准)
# ★ 硬止损倍率:MT5 服务器端 SL/TP = 软止损 × 倍率(兜底,仅 EA 挂掉时触发)
"hard_sl_multiplier": 1.5, # 硬 SL = 软 SL × 1.5-1.5%→-2.25%账户)
"hard_tp_multiplier": 1.3, # 硬 TP = 软 TP × 1.3+3.0%→+3.9%账户)
"hard_sl_multiplier": 1.5, # 硬 SL = 软 SL × 1.5
"hard_tp_multiplier": 1.3, # 硬 TP = 软 TP × 1.3
# ★ 回测适应度门槛:低于此值不开新单(适应度≈回测总盈亏$)
"min_backtest_fitness": 95, # 适应度<95 → 不开仓(低于95说明市场难做)
}
# ★ 最近优化适应度(daily_optimize.py 写入,EA 热加载读取)
LAST_OPTIMIZATION_FITNESS = 58.63
# 市场状态分析参数
MARKET_STATE_CONFIG = {
"trend_period": 83,
"retracement_tolerance": 0.15567164907853615,
"volume_period": 29,
"volume_ma_period": 23
"trend_period": 60,
"retracement_tolerance": 0.3857733104065197,
"volume_period": 42,
"volume_ma_period": 30
}
# 策略参数配置(优化器结果 2026-05-10,5万根M1数据)
STRATEGY_CONFIG = {
"ma_cross": {
"short_window": 13,
"long_window": 17
"short_window": 7,
"long_window": 18
},
"rsi": {
"period": 30,
"overbought": 68,
"oversold": 32
"period": 19,
"overbought": 75,
"oversold": 31
},
"bollinger": {
"period": 20,
"std_dev": 2.4663422649047333
"std_dev": 2.2757986281351896
},
"macd": {
"fast_ema": 14,
"slow_ema": 24,
"signal_period": 9
"fast_ema": 19,
"slow_ema": 23,
"signal_period": 13
},
"mean_reversion": {
"period": 29,
@@ -193,105 +210,108 @@ STRATEGY_CONFIG = {
},
"momentum_breakout": {
"period": 15,
"momentum_period": 13
"momentum_period": 19
},
"kdj": {
"period": 21
},
"turtle": {
"period": 42
"swing_point": {
"left_bars": 4,
"right_bars": 3,
"tolerance_pct": 0.0006056955205083707,
"num_swings": 2
},
"daily_breakout": {
"bars_count": 2008
"bars_count": 1392
},
"wave_theory": {
"ema_short": 3,
"ema_short": 8,
"ema_medium": 11,
"ema_long": 46,
"wave_period": 11,
"range_period": 35,
"adx_period": 20,
"ema_long": 39,
"wave_period": 19,
"range_period": 15,
"adx_period": 24,
"momentum_period": 10,
"range_threshold": 0.01,
"adx_threshold": 21
"range_threshold": 0.0023249568571626195,
"adx_threshold": 28
}
}
# 市场趋势判断权重配置
TREND_INDICATOR_WEIGHTS = {
"price_breakout": 0.1504199442999585,
"volume_confirmation": 0.12782205952949638,
"price_breakout": 0.2506453627492913,
"volume_confirmation": 0.2851951176719032,
"momentum oscillator": 0.3677,
"moving_average": 0.4092273363154768
"moving_average": 0.353451489159348
}
# 趋势判断阈值
TREND_THRESHOLDS = {
"strong_trend": 0.7940886082643032,
"weak_trend": 0.4565953163045441,
"volume_spike": 2.7329673335105396,
"oversold": 32,
"overbought": 68
"strong_trend": 0.4040165929726079,
"weak_trend": 0.2963326526031991,
"volume_spike": 1.5234871986051206,
"oversold": 31,
"overbought": 75
}
# 动态权重配置(优化器结果 2026-05-10,5万根M1数据)
DEFAULT_WEIGHTS = {
"ma_cross": 0.32048324544087786,
"rsi": 1.906009222972592,
"bollinger": 0.9232444456662969,
"mean_reversion": 1.1649306516234597,
"momentum_breakout": 0.8769099086178171,
"macd": 0.748930985323352,
"kdj": 1.4647470623067596,
"turtle": 1.637931136008917,
"daily_breakout": 1.478598649389487,
"wave_theory": 0.818283180489803
"ma_cross": 0.4771175228437661,
"rsi": 0.478479356771334,
"bollinger": 1.3391582231825303,
"mean_reversion": 1.9958459833789346,
"momentum_breakout": 0.7567594572894596,
"macd": 1.170771336653618,
"kdj": 0.013296663814979182,
"swing_point": 0.5194535227966377,
"daily_breakout": 0.6506588091889913,
"wave_theory": 1.6765424080384428
}
# 市场状态策略权重配置
MARKET_STATE_WEIGHTS = {
"uptrend": {
"ma_cross": 0.32048324544087786,
"momentum_breakout": 0.8769099086178171,
"turtle": 1.637931136008917,
"macd": 0.748930985323352,
"daily_breakout": 1.478598649389487,
"rsi": 1.906009222972592,
"bollinger": 0.9232444456662969,
"kdj": 1.4647470623067596,
"mean_reversion": 1.1649306516234597,
"wave_theory": 0.818283180489803
"ma_cross": 0.4771175228437661,
"momentum_breakout": 0.7567594572894596,
"swing_point": 0.5194535227966377,
"macd": 1.170771336653618,
"daily_breakout": 0.6506588091889913,
"rsi": 0.478479356771334,
"bollinger": 1.3391582231825303,
"kdj": 0.013296663814979182,
"mean_reversion": 1.9958459833789346,
"wave_theory": 1.6765424080384428
},
"downtrend": {
"ma_cross": 0.32048324544087786,
"momentum_breakout": 0.8769099086178171,
"turtle": 1.637931136008917,
"macd": 0.748930985323352,
"daily_breakout": 1.478598649389487,
"rsi": 1.906009222972592,
"bollinger": 0.9232444456662969,
"kdj": 1.4647470623067596,
"mean_reversion": 1.1649306516234597,
"wave_theory": 0.818283180489803
"ma_cross": 0.4771175228437661,
"momentum_breakout": 0.7567594572894596,
"swing_point": 0.5194535227966377,
"macd": 1.170771336653618,
"daily_breakout": 0.6506588091889913,
"rsi": 0.478479356771334,
"bollinger": 1.3391582231825303,
"kdj": 0.013296663814979182,
"mean_reversion": 1.9958459833789346,
"wave_theory": 1.6765424080384428
},
"ranging": {
"rsi": 1.906009222972592,
"bollinger": 0.9232444456662969,
"mean_reversion": 1.1649306516234597,
"kdj": 1.4647470623067596,
"wave_theory": 0.818283180489803,
"ma_cross": 0.32048324544087786,
"macd": 0.748930985323352,
"turtle": 1.637931136008917,
"momentum_breakout": 0.8769099086178171,
"daily_breakout": 1.478598649389487
"rsi": 0.478479356771334,
"bollinger": 1.3391582231825303,
"mean_reversion": 1.9958459833789346,
"kdj": 0.013296663814979182,
"wave_theory": 1.6765424080384428,
"ma_cross": 0.4771175228437661,
"macd": 1.170771336653618,
"swing_point": 0.5194535227966377,
"momentum_breakout": 0.7567594572894596,
"daily_breakout": 0.6506588091889913
},
"none": DEFAULT_WEIGHTS
}
# 市场趋势置信度阈值配置
CONFIDENCE_THRESHOLDS = {
"high_confidence": 0.6813641209473742,
"medium_confidence": 0.6336441706563001
"high_confidence": 0.5050416445534804,
"medium_confidence": 0.6102467437388168
}
+10 -3
View File
@@ -54,15 +54,22 @@ class TrailingStopRule(BaseExitRule):
def check(self, ctx: ExitContext) -> tuple[str, str]:
min_profit = self.config.get("min_profit_for_trailing", 0.01)
retracement_pct = self.config.get("profit_retracement_pct", 0.10)
retracement_mode = self.config.get("retracement_mode", "absolute")
if ctx.peak_profit_pct <= min_profit:
return "none", ""
# ★ 回撤从峰值绝对值扣除(账户%,非相对%):峰值+2.0%回撤1.0%→止损在+1.0%
stop_level = ctx.peak_profit_pct - retracement_pct
if retracement_mode == "relative":
# ★ 相对回撤:止损 = 峰值 × (1-回撤%)
# 例: 峰值+50% 回撤30% → 止损+35%(利润从50%回落到35%时平仓)
stop_level = ctx.peak_profit_pct * (1 - retracement_pct)
else:
# ★ 绝对值扣除(旧模式):峰值+2.0%回撤1.0%→止损在+1.0%
stop_level = ctx.peak_profit_pct - retracement_pct
if ctx.current_profit_pct <= stop_level:
mode_tag = "相对" if retracement_mode == "relative" else "绝对"
return "close", (
f"追踪止损触发 "
f"追踪止损触发({mode_tag}回撤) "
f"(峰值 {ctx.peak_profit_pct:.2%} 回落至 {ctx.current_profit_pct:.2%})"
)
return "none", ""
+4 -3
View File
@@ -1,6 +1,7 @@
import pandas as pd
import numpy as np
from logger import logger
import config
from config import (
MARKET_STATE_CONFIG, SYMBOL, DEFAULT_WEIGHTS,
TREND_INDICATOR_WEIGHTS, TREND_THRESHOLDS,
@@ -288,7 +289,7 @@ class MarketStateAnalyzer:
base_weights = dict(individual_weights)
else:
# 正常模式:从配置获取市场状态对应权重
base_weights = dict(MARKET_STATE_WEIGHTS.get(market_state, DEFAULT_WEIGHTS))
base_weights = dict(config.MARKET_STATE_WEIGHTS.get(market_state, config.DEFAULT_WEIGHTS))
high_conf = self.confidence_thresholds.get("high_confidence", 0.7)
medium_conf = self.confidence_thresholds.get("medium_confidence", 0.4)
@@ -297,9 +298,9 @@ class MarketStateAnalyzer:
return {k: v * confidence for k, v in base_weights.items()}
elif confidence > medium_conf:
return {
k: (v * confidence + DEFAULT_WEIGHTS.get(k, 1.0) * (1 - confidence))
k: (v * confidence + config.DEFAULT_WEIGHTS.get(k, 1.0) * (1 - confidence))
for k, v in base_weights.items()
}
else:
# 低置信度:individual_weights 优先(优化器模式),否则回退到 DEFAULT_WEIGHTS
return dict(individual_weights) if individual_weights is not None else dict(DEFAULT_WEIGHTS)
return dict(individual_weights) if individual_weights is not None else dict(config.DEFAULT_WEIGHTS)
+73 -8
View File
@@ -3,10 +3,8 @@ import numpy as np
import json
import os
from logger import logger
from config import (
RISK_CONFIG, SYMBOL, INITIAL_CAPITAL, CAPITAL_ALLOCATION,
RISK_CONFIG_CONST, SIMULATION_CONFIG
)
import config
from config import RISK_CONFIG, SYMBOL, INITIAL_CAPITAL, CAPITAL_ALLOCATION, RISK_CONFIG_CONST, SIMULATION_CONFIG
from core.risk.exit_rules import ExitRuleEngine, ExitContext
@@ -61,6 +59,7 @@ class PositionManager:
self.min_profit_for_trailing = risk.get("min_profit_for_trailing", 0.01) * self._lev_ratio
self.take_profit_pct = risk.get("take_profit_pct", 0.20) * self._lev_ratio
self.min_profit_for_time_exit = risk.get("min_profit_for_time_exit", 0.001) * self._lev_ratio
self.retracement_mode = risk.get("retracement_mode", "absolute") # relative/absolute
# 资金管理
self.initial_capital = INITIAL_CAPITAL
@@ -89,13 +88,24 @@ class PositionManager:
"max_holding_minutes": self.max_holding_minutes,
"min_profit_for_time_exit": self.min_profit_for_time_exit,
"max_daily_loss": self.max_daily_loss,
"retracement_mode": self.retracement_mode,
}
self.exit_engine = ExitRuleEngine(exit_config)
self._exit_config = exit_config # 保存引用,供波动率自适应更新
# ★ 波动率自适应拖尾 — 保存基准值
self._base_min_profit_for_trailing = self.min_profit_for_trailing
self._base_profit_retracement_pct = self.profit_retracement_pct
self._vol_adaptive = risk.get("vol_adaptive_trailing", False)
self._trailing_atr_period = risk.get("trailing_atr_period", 14)
self._trailing_atr_baseline = risk.get("trailing_atr_baseline", 100)
self._last_atr_update = None # 节流:最多1分钟更新一次
if self._vol_adaptive:
logger.info(f"📐 波动率自适应拖尾已启用 (ATR{self._trailing_atr_period}/ATR{self._trailing_atr_baseline})")
# 峰值数据持久化(优化器中禁用文件I/O避免多进程竞争)
self.peak_data_file = "position_peaks.json"
if self._persist_peaks:
self._load_peak_data()
self._saved_peaks = self._load_peak_data() if self._persist_peaks else {}
# 对冲管理器
from config import HEDGE_CONFIG
@@ -106,6 +116,54 @@ class PositionManager:
self._pending_long = 0
self._pending_short = 0
# ── ATR 计算与波动率自适应 ──
def _compute_atr(self, period: int) -> float | None:
"""计算指定周期的 ATRAverage True Range"""
try:
import numpy as np
rates = self.data_provider.get_historical_data(self.symbol, 1, period + 1)
if rates is None or len(rates) < period + 1:
return None
highs = np.array([r[2] for r in rates[-period-1:]]) # high
lows = np.array([r[3] for r in rates[-period-1:]]) # low
closes = np.array([r[4] for r in rates[-period-1:]]) # close
tr = np.maximum(
highs[1:] - lows[1:],
np.maximum(
np.abs(highs[1:] - closes[:-1]),
np.abs(lows[1:] - closes[:-1])
)
)
return float(np.mean(tr))
except Exception as e:
logger.debug(f"ATR 计算失败: {e}")
return None
def _update_volatility_trailing(self):
"""波动率自适应:按 ATR 比率调整拖尾激活和回撤参数"""
if not self._vol_adaptive:
return
# 节流:最多每分钟更新一次
from datetime import datetime
now = datetime.now()
if self._last_atr_update is not None:
if (now - self._last_atr_update).total_seconds() < 60:
return
short_atr = self._compute_atr(self._trailing_atr_period)
long_atr = self._compute_atr(self._trailing_atr_baseline)
if short_atr is None or long_atr is None or long_atr <= 0:
return
vol_ratio = short_atr / long_atr
# 限制极端值:0.5 ~ 2.0
vol_ratio = max(0.5, min(2.0, vol_ratio))
# 更新
self.min_profit_for_trailing = self._base_min_profit_for_trailing * vol_ratio
self.profit_retracement_pct = self._base_profit_retracement_pct * vol_ratio
self._exit_config["min_profit_for_trailing"] = self.min_profit_for_trailing
self._exit_config["profit_retracement_pct"] = self.profit_retracement_pct
self._last_atr_update = now
# ── 仓位计算 ──
def _calculate_position_size(self, capital_to_allocate, current_price):
@@ -277,6 +335,9 @@ class PositionManager:
if not self.positions:
return
# ★ 波动率自适应拖尾 — 每分钟更新一次
self._update_volatility_trailing()
current_time = current_price.get('time', pd.Timestamp.now())
positions_to_remove = []
@@ -326,7 +387,7 @@ class PositionManager:
self.cleanup_peak_data()
# ── 对冲评估 ──
if self.positions and self.hedge_manager:
if self.positions and hasattr(self, 'hedge_manager') and self.hedge_manager:
hedge_actions = self.hedge_manager.evaluate(weighted_signal, current_price)
for action, target, reason in hedge_actions:
if action == "hedge":
@@ -441,6 +502,8 @@ class PositionManager:
if live_positions is None:
return
saved_peaks = self._load_peak_data()
# 合并构造器中加载的峰值(优先已持久化)
saved_peaks = {**saved_peaks, **getattr(self, '_saved_peaks', {})}
existing_peaks = {pos['ticket']: pos.get('peak_profit_pct', 0.0) for pos in self.positions}
merged_peaks = {**existing_peaks, **saved_peaks}
self.positions.clear()
@@ -496,7 +559,9 @@ class PositionManager:
try:
if os.path.exists(self.peak_data_file):
with open(self.peak_data_file, 'r', encoding='utf-8') as f:
return json.load(f)
raw = json.load(f)
# ★ JSON 键是字符串,MT5 ticket 是整数 → 统一转 int
return {int(k): v for k, v in raw.items()}
except Exception as e:
logger.error(f"加载峰值数据失败: {e}")
return {}
+132 -282
View File
@@ -1,21 +1,21 @@
"""遗传算法优化器(重写版
"""贝叶斯优化器(Optuna TPE
高效替代 DEAP 遗传算法,50次试验 ≈ 2分钟收敛
关键改进:
1. 权重基因通过 individual_weights → MarketStateAnalyzer.get_strategy_weights() 真正生效
2. 使用 StrategyRegistry 单例消,除重复的策略列表
3. 使用 MultiTimeframeDataStore 支持正确的多周期市场状态分析
4. evaluate_fitness 逐 bar 调用 get_market_state(bar_index) 获取每根bar的动态权重
1. TPE sampler 建模适应度曲面,采样效率高 3-5 倍
2. 权重基因通过 individual_weights → MarketStateAnalyzer.get_strategy_weights() 真正生效
3. evaluate_fitness 逐 bar 调用 get_market_state(bar_index) 获取每根bar的动态权重
"""
import random
import numpy as np
from deap import base, creator, tools, algorithms
import multiprocessing
import pandas as pd
import sys
import os
import logging
from datetime import datetime
from tqdm import tqdm
import optuna
SEED = 42
random.seed(SEED)
@@ -31,7 +31,7 @@ from core.signal.combiner import SignalCombiner
from config import (
SYMBOL, TIMEFRAME, OPTIMIZER_COUNT, OPTIMIZER_START_DATE, OPTIMIZER_END_DATE,
USE_DATE_RANGE, INITIAL_CAPITAL, SIGNAL_THRESHOLDS, DEFAULT_WEIGHTS,
RISK_CONFIG, GENETIC_OPTIMIZER_CONFIG,
RISK_CONFIG,
MARKET_STATE_CONFIG, TREND_INDICATOR_WEIGHTS, TREND_THRESHOLDS, CONFIDENCE_THRESHOLDS,
DATA_PROVIDER_MODE, REMOTE_SERVER_HOST, REMOTE_SERVER_PORT
)
@@ -45,8 +45,6 @@ import strategies # noqa: F401
# ── 全局缓存 ──
_multi_tf: MultiTimeframeDataStore | None = None
_cached_signals: pd.DataFrame | None = None
_registry: StrategyRegistry | None = None
# MT5 结构化数组 dtype(用于远程 API JSON → numpy 转换)
_MT5_RATES_DTYPE = np.dtype([
@@ -61,15 +59,7 @@ _MT5_RATES_DTYPE = np.dtype([
])
def init_worker():
"""多进程worker初始化:抑制日志噪音"""
import logging
logging.getLogger().setLevel(logging.WARNING)
for name in ["StrategyLogger", "PositionManager", "RiskController", "DataProvider"]:
logging.getLogger(name).setLevel(logging.WARNING)
# ── 参数定义(与旧版一致,保持兼容)──
# ── 参数定义 ──
PARAMETER_DEFINITIONS = [
# MACrossStrategy
{'name': 'ma_cross_short_window', 'type': 'int', 'min': 3, 'max': 15, 'strategy': 'MACrossStrategy'},
@@ -93,8 +83,11 @@ PARAMETER_DEFINITIONS = [
{'name': 'momentum_breakout_momentum_period', 'type': 'int', 'min': 5, 'max': 30, 'strategy': 'MomentumBreakoutStrategy'},
# KDJStrategy
{'name': 'kdj_period', 'type': 'int', 'min': 5, 'max': 21, 'strategy': 'KDJStrategy'},
# TurtleStrategy
{'name': 'turtle_period', 'type': 'int', 'min': 10, 'max': 50, 'strategy': 'TurtleStrategy'},
# SwingPointRetestStrategy — 前高前低回踩
{'name': 'swing_point_left_bars', 'type': 'int', 'min': 2, 'max': 5, 'strategy': 'SwingPointRetestStrategy'},
{'name': 'swing_point_right_bars', 'type': 'int', 'min': 2, 'max': 5, 'strategy': 'SwingPointRetestStrategy'},
{'name': 'swing_point_tolerance_pct', 'type': 'float', 'min': 0.0003, 'max': 0.002, 'strategy': 'SwingPointRetestStrategy'},
{'name': 'swing_point_num_swings', 'type': 'int', 'min': 1, 'max': 3, 'strategy': 'SwingPointRetestStrategy'},
# WaveTheoryStrategy
{'name': 'wave_ema_short', 'type': 'int', 'min': 3, 'max': 10, 'strategy': 'WaveTheoryStrategy'},
{'name': 'wave_ema_medium', 'type': 'int', 'min': 8, 'max': 20, 'strategy': 'WaveTheoryStrategy'},
@@ -110,14 +103,7 @@ PARAMETER_DEFINITIONS = [
# Signal Thresholds
{'name': 'buy_threshold', 'type': 'float', 'min': 0.5, 'max': 3.0, 'strategy': 'signal'},
{'name': 'sell_threshold', 'type': 'float', 'min': -3.0, 'max': -0.5, 'strategy': 'signal'},
# Risk Management(范围适配保证金%基准)
{'name': 'stop_loss_pct', 'type': 'float', 'min': -1.00, 'max': -0.10, 'strategy': 'risk'},
{'name': 'profit_retracement_pct', 'type': 'float', 'min': 0.10, 'max': 0.80, 'strategy': 'risk'},
{'name': 'min_profit_for_trailing', 'type': 'float', 'min': 0.30, 'max': 2.00, 'strategy': 'risk'},
{'name': 'take_profit_pct', 'type': 'float', 'min': 0.50, 'max': 3.00, 'strategy': 'risk'},
{'name': 'max_holding_minutes', 'type': 'int', 'min': 30, 'max': 180, 'strategy': 'risk'},
{'name': 'min_profit_for_time_exit', 'type': 'float', 'min': 0.02, 'max': 0.20, 'strategy': 'risk'},
# ★ Strategy Weights — 这些基因现在真正影响适应度
# ★ Strategy Weights
{'name': 'weight_MACrossStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_RSIStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_BollingerStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
@@ -125,7 +111,7 @@ PARAMETER_DEFINITIONS = [
{'name': 'weight_MomentumBreakoutStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_MACDStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_KDJStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_TurtleStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_SwingPointRetestStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_DailyBreakoutStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
{'name': 'weight_WaveTheoryStrategy', 'type': 'float', 'min': 0.0, 'max': 2.0, 'strategy': 'weight'},
# Market State Analysis parameters
@@ -149,7 +135,6 @@ PARAMETER_DEFINITIONS = [
{'name': 'confidence_medium', 'type': 'float', 'min': 0.3, 'max': 0.7, 'strategy': 'confidence'},
]
# class_name → config_key 映射(从 StrategyRegistry 获取)
_class_name_to_config_key = {
"MACrossStrategy": "ma_cross",
"RSIStrategy": "rsi",
@@ -158,27 +143,27 @@ _class_name_to_config_key = {
"MomentumBreakoutStrategy": "momentum_breakout",
"MACDStrategy": "macd",
"KDJStrategy": "kdj",
"TurtleStrategy": "turtle",
"SwingPointRetestStrategy": "swing_point",
"DailyBreakoutStrategy": "daily_breakout",
"WaveTheoryStrategy": "wave_theory",
}
def _parse_individual(individual):
"""解析个体基因为命名参数字典,并裁剪到定义范围"""
def _trial_to_params(trial: optuna.Trial) -> dict:
"""从 Optuna trial 提取参数 → 兼容旧 _parse_individual 格式"""
parsed = {}
for idx, param_def in enumerate(PARAMETER_DEFINITIONS):
val = individual[idx]
for param_def in PARAMETER_DEFINITIONS:
name = param_def['name']
if param_def['type'] == 'int':
val = int(round(val))
# 裁剪到 [min, max] 防止变异越界
val = max(param_def['min'], min(param_def['max'], val))
parsed[param_def['name']] = val
val = trial.suggest_int(name, param_def['min'], param_def['max'])
else:
val = trial.suggest_float(name, param_def['min'], param_def['max'])
parsed[name] = val
return parsed
def _extract_strategy_params(parsed: dict) -> dict:
"""从解析后的基因提取策略参数字典"""
"""从解析后的参数提取策略参数字典"""
strategy_params = {}
for param_def in PARAMETER_DEFINITIONS:
strategy_name = param_def.get('strategy')
@@ -188,57 +173,37 @@ def _extract_strategy_params(parsed: dict) -> dict:
if strategy_name not in strategy_params:
strategy_params[strategy_name] = {}
# 参数名映射(保持与旧版兼容)
param_name = param_def['name']
if param_name == 'ma_cross_short_window':
param_name = 'short_window'
elif param_name == 'ma_cross_long_window':
param_name = 'long_window'
elif param_name == 'rsi_period':
param_name = 'period'
elif param_name == 'rsi_overbought':
param_name = 'overbought'
elif param_name == 'rsi_oversold':
param_name = 'oversold'
elif param_name == 'bollinger_period':
param_name = 'period'
elif param_name == 'bollinger_std_dev':
param_name = 'std_dev'
elif param_name == 'macd_fast_ema':
param_name = 'fast_ema'
elif param_name == 'macd_slow_ema':
param_name = 'slow_ema'
elif param_name == 'macd_signal_period':
param_name = 'signal_period'
elif param_name == 'mean_reversion_period':
param_name = 'period'
elif param_name == 'mean_reversion_std_dev':
param_name = 'std_dev'
elif param_name == 'momentum_breakout_period':
param_name = 'period'
elif param_name == 'momentum_breakout_momentum_period':
param_name = 'momentum_period'
elif param_name == 'kdj_period':
param_name = 'period'
elif param_name == 'turtle_period':
param_name = 'period'
if param_name == 'ma_cross_short_window': param_name = 'short_window'
elif param_name == 'ma_cross_long_window': param_name = 'long_window'
elif param_name == 'rsi_period': param_name = 'period'
elif param_name == 'rsi_overbought': param_name = 'overbought'
elif param_name == 'rsi_oversold': param_name = 'oversold'
elif param_name == 'bollinger_period': param_name = 'period'
elif param_name == 'bollinger_std_dev': param_name = 'std_dev'
elif param_name == 'macd_fast_ema': param_name = 'fast_ema'
elif param_name == 'macd_slow_ema': param_name = 'slow_ema'
elif param_name == 'macd_signal_period': param_name = 'signal_period'
elif param_name == 'mean_reversion_period': param_name = 'period'
elif param_name == 'mean_reversion_std_dev': param_name = 'std_dev'
elif param_name == 'momentum_breakout_period': param_name = 'period'
elif param_name == 'momentum_breakout_momentum_period': param_name = 'momentum_period'
elif param_name == 'kdj_period': param_name = 'period'
elif param_name == 'swing_point_left_bars': param_name = 'left_bars'
elif param_name == 'swing_point_right_bars': param_name = 'right_bars'
elif param_name == 'swing_point_tolerance_pct': param_name = 'tolerance_pct'
elif param_name == 'swing_point_num_swings': param_name = 'num_swings'
elif param_name.startswith("wave_"):
param_name = param_name[5:]
if param_name == "period":
param_name = "wave_period"
elif param_name == 'daily_breakout_bars_count':
param_name = 'bars_count'
if param_name == "period": param_name = "wave_period"
elif param_name == 'daily_breakout_bars_count': param_name = 'bars_count'
strategy_params[strategy_name][param_name] = parsed[param_def['name']]
return strategy_params
def _extract_weight_genes(parsed: dict) -> dict:
"""提取权重基因 → {config_key: weight}
这些权重会通过 individual_weights 参数传递给
MarketStateAnalyzer.get_strategy_weights(),使权重基因真正生效。
"""
"""提取权重基因 → {config_key: weight}"""
weight_genes = {}
for param_def in PARAMETER_DEFINITIONS:
if param_def.get('strategy') != 'weight':
@@ -258,7 +223,6 @@ def _precompute_h1_states(multi_tf: MultiTimeframeDataStore, analyzer: MarketSta
m1_index = multi_tf.main_df.index
h1_index = h1_df.index
# 计算每个H1 bar的状态
h1_states = []
for i in range(len(h1_df)):
lookback = analyzer.trend_period + 10
@@ -267,11 +231,9 @@ def _precompute_h1_states(multi_tf: MultiTimeframeDataStore, analyzer: MarketSta
state, conf = analyzer._calculate_state_from_df(df_slice)
h1_states.append((state, conf))
# 映射到M1时间轴
m1_states = []
h1_pos = 0
first_h1_time = h1_index[0]
for m1_time in m1_index:
if m1_time < first_h1_time:
m1_states.append(("none", 0.0))
@@ -282,44 +244,34 @@ def _precompute_h1_states(multi_tf: MultiTimeframeDataStore, analyzer: MarketSta
m1_states.append(h1_states[h1_pos])
else:
m1_states.append(("none", 0.0))
return m1_states
def evaluate_fitness(individual, multi_tf: MultiTimeframeDataStore,
_signals_df_unused=None) -> tuple:
"""★ 重写的适应度函数 — 所有权重基因、策略参数和风控参数真正生效
def evaluate_fitness(trial: optuna.Trial, multi_tf: MultiTimeframeDataStore) -> float:
"""★ Optuna 适应度函数 — 所有权重/策略/趋势参数真正生效"""
parsed = _trial_to_params(trial)
三个关键基因全部生效:
1. 策略参数 → 重新实例化策略并 run_backtest → 影响信号序列
2. 权重基因 → per-bar get_strategy_weights(individual_weights=...) → 影响信号组合
3. 风控参数 → PositionManager 在回测中真正使用
"""
parsed = _parse_individual(individual)
# 1. 提取策略参数
# 1. 策略参数
strategy_params = _extract_strategy_params(parsed)
# 2. ★ 提取权重基因(核心修复)
# 2. 权重基因
weight_genes = _extract_weight_genes(parsed)
# 3. 提取风控参数
# 3. 风控参数(从 config 读,不进优化器)
risk_params = {
'stop_loss_pct': parsed.get('stop_loss_pct', RISK_CONFIG.get('stop_loss_pct', -0.01)),
'profit_retracement_pct': parsed.get('profit_retracement_pct', RISK_CONFIG.get('profit_retracement_pct', 0.10)),
'min_profit_for_trailing': parsed.get('min_profit_for_trailing', RISK_CONFIG.get('min_profit_for_trailing', 0.01)),
'take_profit_pct': parsed.get('take_profit_pct', RISK_CONFIG.get('take_profit_pct', 0.20)),
'max_holding_minutes': parsed.get('max_holding_minutes', RISK_CONFIG.get('max_holding_minutes', 60)),
'min_profit_for_time_exit': parsed.get('min_profit_for_time_exit', RISK_CONFIG.get('min_profit_for_time_exit', 0.005)),
'max_daily_loss': parsed.get('max_daily_loss', RISK_CONFIG.get('max_daily_loss', -0.30)),
'stop_loss_pct': RISK_CONFIG.get('stop_loss_pct', -0.50),
'profit_retracement_pct': RISK_CONFIG.get('profit_retracement_pct', 0.35),
'min_profit_for_trailing': RISK_CONFIG.get('min_profit_for_trailing', 0.70),
'take_profit_pct': RISK_CONFIG.get('take_profit_pct', 1.00),
'max_holding_minutes': RISK_CONFIG.get('max_holding_minutes', 0),
'min_profit_for_time_exit': RISK_CONFIG.get('min_profit_for_time_exit', 0.005),
'max_daily_loss': RISK_CONFIG.get('max_daily_loss', -0.30),
}
# 4. 提取市场状态和趋势参数
# 4. 市场状态和趋势参数
market_params = {
'trend_period': parsed.get('market_trend_period', MARKET_STATE_CONFIG.get('trend_period', 50)),
'retracement_tolerance': parsed.get('market_retracement_tolerance', MARKET_STATE_CONFIG.get('retracement_tolerance', 0.30)),
'volume_period': parsed.get('market_volume_period', MARKET_STATE_CONFIG.get('volume_period', 20)),
'volume_ma_period': parsed.get('market_volume_ma_period', MARKET_STATE_CONFIG.get('volume_ma_period', 10)),
'trend_period': parsed.get('market_trend_period', 50),
'retracement_tolerance': parsed.get('market_retracement_tolerance', 0.30),
'volume_period': parsed.get('market_volume_period', 20),
'volume_ma_period': parsed.get('market_volume_ma_period', 10),
'hourly_data_count': 100,
}
trend_weights = {
@@ -340,7 +292,7 @@ def evaluate_fitness(individual, multi_tf: MultiTimeframeDataStore,
'medium_confidence': parsed.get('confidence_medium', 0.4),
}
# 5. 构建 MarketStateAnalyzer 并预计算H1状态
# 5. MarketStateAnalyzer
analyzer = MarketStateAnalyzer(
timeframe=PERIOD_H1,
market_state_params=market_params,
@@ -350,8 +302,7 @@ def evaluate_fitness(individual, multi_tf: MultiTimeframeDataStore,
)
analyzer._precomputed_states = _precompute_h1_states(multi_tf, analyzer)
# 5.5. ★ 使用个体策略参数重新实例化策略并重新计算信号
# 将 class_name → {params} 映射为 config_key → {params}
# 5.5 用个体策略参数重新计算信号
strategy_params_by_key = {
_class_name_to_config_key.get(cn, cn.lower()): params
for cn, params in strategy_params.items()
@@ -375,29 +326,19 @@ def evaluate_fitness(individual, multi_tf: MultiTimeframeDataStore,
sell_th = parsed.get('sell_threshold', -1.5)
total_bars = multi_tf.length
# 7. 逐 bar 回测(权重基因真正影响信号组合)
# 7. 逐 bar 回测
for bar_index in range(total_bars):
current_price = data_provider.get_current_price(SYMBOL)
if not current_price:
data_provider.tick()
continue
# 获取当前 bar 的市场状态
state, conf = analyzer.get_market_state(bar_index)
weights = analyzer.get_strategy_weights(state, conf, individual_weights=weight_genes)
# ★ 权重基因在这里生效
weights = analyzer.get_strategy_weights(
state, conf, individual_weights=weight_genes
)
# 组合当前 bar 的信号(使用个体参数重新计算的信号)
bar_signals = {
col: signals_df[col].iloc[bar_index]
for col in signals_df.columns
}
bar_signals = {col: signals_df[col].iloc[bar_index] for col in signals_df.columns}
final_signal = SignalCombiner.combine_at_bar(bar_signals, weights, buy_th, sell_th)
# 执行交易
if final_signal == 1:
pm.open_position("buy", current_price, 1.0, dry_run=True)
elif final_signal == -1:
@@ -407,20 +348,16 @@ def evaluate_fitness(individual, multi_tf: MultiTimeframeDataStore,
data_provider.tick()
# 8. 适应度 = 总盈亏
total_pnl = pm.total_equity - INITIAL_CAPITAL
return (total_pnl,)
return pm.total_equity - INITIAL_CAPITAL
def _json_to_mt5_rates(rates_data):
"""远程API JSON rates 转换为 MT5 兼容的 numpy 结构化数组"""
"""远程API JSON → MT5 numpy 结构化数组"""
if not rates_data:
return None
records = []
for r in rates_data:
records.append((
r['time'], r['open'], r['high'], r['low'], r['close'],
r.get('tick_volume', 0), r.get('spread', 0), r.get('real_volume', 0),
))
records = [(r['time'], r['open'], r['high'], r['low'], r['close'],
r.get('tick_volume', 0), r.get('spread', 0), r.get('real_volume', 0))
for r in rates_data]
return np.array(records, dtype=_MT5_RATES_DTYPE)
@@ -430,152 +367,96 @@ def load_historical_data():
provider = RemoteDataProvider(host=REMOTE_SERVER_HOST, port=REMOTE_SERVER_PORT)
if not provider.initialize():
raise RuntimeError("远程MT5 API初始化失败,请检查 Windows MT5 是否运行")
rates_json = provider.get_historical_data(SYMBOL, TIMEFRAME, OPTIMIZER_COUNT)
provider.shutdown()
if not rates_json:
raise RuntimeError("远程获取历史数据失败")
rates = _json_to_mt5_rates(rates_json)
else:
initialize()
rates = (
get_rates(SYMBOL, TIMEFRAME, OPTIMIZER_COUNT, OPTIMIZER_START_DATE, OPTIMIZER_END_DATE)
if USE_DATE_RANGE else
get_rates(SYMBOL, TIMEFRAME, OPTIMIZER_COUNT)
)
rates = (get_rates(SYMBOL, TIMEFRAME, OPTIMIZER_COUNT, OPTIMIZER_START_DATE, OPTIMIZER_END_DATE)
if USE_DATE_RANGE else get_rates(SYMBOL, TIMEFRAME, OPTIMIZER_COUNT))
shutdown()
if rates is None or len(rates) == 0:
raise RuntimeError("获取历史数据失败")
if len(rates) > OPTIMIZER_COUNT:
rates = rates[-OPTIMIZER_COUNT:]
store = MultiTimeframeDataStore()
store.load_m1_data(rates)
# 预生成H1数据
store.ensure_timeframe(PERIOD_H1)
return store
def precompute_all_signals(multi_tf: MultiTimeframeDataStore,
strategy_params: dict = None) -> pd.DataFrame:
"""使用默认参数预计算所有策略的回测信号"""
registry = StrategyRegistry()
strategies = registry.instantiate_all(SYMBOL, TIMEFRAME, strategy_params)
signals_df = pd.DataFrame(index=multi_tf.main_df.index)
def save_optimization_results(best_params: dict, best_fitness: float, study: optuna.Study):
"""保存优化结果(兼容旧格式)"""
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
import json
for config_key, strategy in strategies.items():
try:
sig = strategy.run_backtest(multi_tf.main_df)
signals_df[config_key] = sig if sig is not None else pd.Series(0, index=signals_df.index)
except Exception:
signals_df[config_key] = pd.Series(0, index=signals_df.index)
results = {
'best_fitness': best_fitness,
'best_params': best_params,
'n_trials': len(study.trials),
'optimizer': 'optuna_tpe',
}
json_path = f"optimization_results_{timestamp}.json"
with open(json_path, 'w', encoding='utf-8') as f:
json.dump(results, f, ensure_ascii=False, indent=2, default=str)
logger.info(f"结果已保存: {json_path}")
return signals_df
txt_path = f"optimization_report_{timestamp}.txt"
with open(txt_path, 'w', encoding='utf-8') as f:
f.write(f"优化完成时间: {datetime.now()}\n")
f.write(f"最佳适应度: {best_fitness:.2f}\n")
f.write(f"试验次数: {len(study.trials)}\n\n")
f.write("最佳参数:\n")
for key, value in best_params.items():
f.write(f" {key}: {value}\n")
logger.info(f"报告已保存: {txt_path}")
def run_optimizer():
"""运行遗传算法优化(所有基因真正生效"""
"""★ Optuna TPE 贝叶斯优化(替代 DEAP 遗传算法"""
try:
# 1. 加载数据(一次性)
# 1. 加载数据
logger.info("加载历史数据...")
multi_tf = load_historical_data()
logger.info(f"数据加载完成: {multi_tf.length} 条M1数据")
# 2. DEAP 设置(策略信号在 evaluate_fitness 内部逐个体重新计算)
creator.create("FitnessMax", base.Fitness, weights=(1.0,))
creator.create("Individual", list, fitness=creator.FitnessMax)
# 2. Optuna study
sampler = optuna.samplers.TPESampler(seed=SEED, n_startup_trials=10)
study = optuna.create_study(direction="maximize", sampler=sampler)
toolbox = base.Toolbox()
for i, param_def in enumerate(PARAMETER_DEFINITIONS):
if param_def['type'] == 'int':
toolbox.register(f"attr_param_{i}", random.randint, param_def['min'], param_def['max'])
else:
toolbox.register(f"attr_param_{i}", random.uniform, param_def['min'], param_def['max'])
toolbox.register("individual", tools.initIterate, creator.Individual,
lambda: [toolbox.__getattribute__(f"attr_param_{j}")()
for j in range(len(PARAMETER_DEFINITIONS))])
toolbox.register("population", tools.initRepeat, list, toolbox.individual)
toolbox.register("evaluate", evaluate_fitness,
multi_tf=multi_tf)
toolbox.register("mate", tools.cxTwoPoint)
toolbox.register("mutate", tools.mutGaussian,
mu=GENETIC_OPTIMIZER_CONFIG["mutation_mu"],
sigma=GENETIC_OPTIMIZER_CONFIG["mutation_sigma"],
indpb=GENETIC_OPTIMIZER_CONFIG["mutation_indpb"])
toolbox.register("select", tools.selTournament,
tournsize=GENETIC_OPTIMIZER_CONFIG["tournament_size"])
# 多进程
if GENETIC_OPTIMIZER_CONFIG["enable_multiprocessing"]:
pool = multiprocessing.Pool(
processes=GENETIC_OPTIMIZER_CONFIG["processes"],
initializer=init_worker
# 3. 优化(4进程并行,静默回测日志避免 I/O 争用)
optuna.logging.set_verbosity(optuna.logging.WARNING)
import logging as _logging
_old_level = logger.level
_old_handler_levels = [h.level for h in logger.handlers]
logger.setLevel(_logging.WARNING)
for h in logger.handlers:
h.setLevel(_logging.WARNING)
logger.info(f"开始贝叶斯优化: 50 次试验, TPE sampler (回测日志已静默)")
try:
study.optimize(
lambda trial: evaluate_fitness(trial, multi_tf),
n_trials=50,
n_jobs=4,
show_progress_bar=False,
)
toolbox.register("map", pool.map)
else:
toolbox.register("map", map)
finally:
logger.setLevel(_old_level)
for h, lvl in zip(logger.handlers, _old_handler_levels):
h.setLevel(lvl)
# 4. 统计
stats = tools.Statistics(lambda ind: ind.fitness.values)
stats.register("avg", np.mean)
stats.register("std", np.std)
stats.register("min", np.min)
stats.register("max", np.max)
# 4. 结果
best_fitness = study.best_value
best_params = study.best_params
logger.info(f"优化完成: 最佳适应度={best_fitness:.2f}, 试验次数={len(study.trials)}")
# 5. 运行
population = toolbox.population(n=GENETIC_OPTIMIZER_CONFIG["population_size"])
ngen = GENETIC_OPTIMIZER_CONFIG["generations"]
cxpb = GENETIC_OPTIMIZER_CONFIG["crossover_probability"]
mutpb = GENETIC_OPTIMIZER_CONFIG["mutation_probability"]
save_optimization_results(best_params, best_fitness, study)
logger.info(f"开始进化: 种群 {len(population)}, {ngen}")
generation_info = []
for gen in range(ngen):
offspring = toolbox.select(population, len(population))
offspring = algorithms.varAnd(offspring, toolbox, cxpb, mutpb)
fits = toolbox.map(toolbox.evaluate, offspring)
for fit, ind in zip(fits, offspring):
ind.fitness.values = fit
population[:] = offspring
best = tools.selBest(population, k=1)[0]
best_fit = best.fitness.values[0]
avg_fit = np.mean([ind.fitness.values[0] for ind in population])
if GENETIC_OPTIMIZER_CONFIG.get("save_generation_info", True):
generation_info.append({
'generation': gen + 1,
'best_fitness': best_fit,
'avg_fitness': avg_fit,
})
if GENETIC_OPTIMIZER_CONFIG.get("verbose", True):
logger.info(f"代数 {gen + 1}/{ngen}: 最佳={best_fit:.2f}, 平均={avg_fit:.2f}")
# 6. 结果
if GENETIC_OPTIMIZER_CONFIG["enable_multiprocessing"]:
pool.close()
pool.join()
best_individual = tools.selBest(population, k=1)[0]
best_fitness = best_individual.fitness.values[0]
logger.info(f"优化完成: 最佳适应度={best_fitness:.2f}")
save_optimization_results(best_individual, best_fitness, generation_info)
return _parse_individual(best_individual), best_fitness
return best_params, best_fitness
except Exception as e:
import traceback
@@ -583,36 +464,5 @@ def run_optimizer():
raise
def save_optimization_results(best_individual, best_fitness, generation_info):
"""保存优化结果"""
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
parsed = _parse_individual(best_individual)
# JSON
import json
results = {
'best_fitness': best_fitness,
'best_params': parsed,
'generation_info': generation_info,
}
json_path = f"optimization_results_{timestamp}.json"
with open(json_path, 'w', encoding='utf-8') as f:
json.dump(results, f, ensure_ascii=False, indent=2, default=str)
logger.info(f"结果已保存: {json_path}")
# TXT报告
txt_path = f"optimization_report_{timestamp}.txt"
with open(txt_path, 'w', encoding='utf-8') as f:
f.write(f"优化完成时间: {datetime.now()}\n")
f.write(f"最佳适应度: {best_fitness:.2f}\n\n")
f.write("最佳参数:\n")
for key, value in parsed.items():
f.write(f" {key}: {value}\n")
f.write("\n代数统计:\n")
for gi in generation_info:
f.write(f" 代数 {gi['generation']}: 最佳={gi['best_fitness']:.2f}, 平均={gi['avg_fitness']:.2f}\n")
logger.info(f"报告已保存: {txt_path}")
if __name__ == "__main__":
run_optimizer()
+51 -18
View File
@@ -3,12 +3,12 @@ import signal
import sys
from datetime import datetime
from logger import logger
from config import SYMBOL, TIMEFRAME, REALTIME_CONFIG, SIGNAL_THRESHOLDS, RISK_CONFIG, RISK_CONFIG_CONST
import config
from core.risk import RiskController
from execution.weights import DynamicWeightManager
class RealtimeTrader:
"""实时交易器 (已重构为依赖注入)"""
"""实时交易器 (已重构为依赖注入) — ★ 每周期热加载配置"""
def __init__(self, data_provider, update_interval=60):
self.data_provider = data_provider
@@ -29,27 +29,28 @@ class RealtimeTrader:
signal.signal(signal.SIGINT, self._signal_handler)
signal.signal(signal.SIGTERM, self._signal_handler)
# ★ 杠杆自适应:打印有效值
# ★ 打印风控参数(保证金%基准)
try:
acct = self.data_provider.get_account_info()
lev = acct.leverage if hasattr(acct, 'leverage') else acct.get('leverage', 2000)
except Exception:
lev = 2000
ratio = lev / 100.0
rl = config.RISK_CONFIG.get('risk_leverage', 100)
# ★ 启动参数一览
logger.info("=" * 50)
logger.info(f"品种: {SYMBOL} | 周期: M{TIMEFRAME} | 间隔: {self.update_interval}s | 杠杆: {lev}x")
logger.info(f"风控: 止损={RISK_CONFIG['stop_loss_pct']*ratio:.1%} | "
f"止盈={RISK_CONFIG['take_profit_pct']*ratio:.1%} | "
f"拖尾激活={RISK_CONFIG['min_profit_for_trailing']*ratio:.1%} | "
f"拖尾回撤={RISK_CONFIG['profit_retracement_pct']*ratio:.1%}")
logger.info(f"信号: 买入阈值={SIGNAL_THRESHOLDS.get('buy_threshold',1.5)} | "
f"卖出阈值={SIGNAL_THRESHOLDS.get('sell_threshold',-1.5)}")
logger.info(f"仓位: 最多多={REALTIME_CONFIG['max_long_positions']} 最多空={REALTIME_CONFIG['max_short_positions']} | "
f"超时平仓={'' if RISK_CONFIG_CONST.get('enable_time_based_exit',True) else ''}")
logger.info(f"对冲: 信号对冲={'' if REALTIME_CONFIG.get('hedge_enabled',False) else ''} | "
f"锁仓={'' if REALTIME_CONFIG.get('lock_enabled',False) else ''}")
logger.info(f"品种: {config.SYMBOL} | 周期: M{config.TIMEFRAME} | 间隔: {self.update_interval}s | 杠杆: {lev}x")
logger.info(f"风控: 止损={config.RISK_CONFIG['stop_loss_pct']:.0%} | "
f"止盈={config.RISK_CONFIG['take_profit_pct']:.0%} | "
f"拖尾激活={config.RISK_CONFIG['min_profit_for_trailing']:.0%} | "
f"拖尾回撤={config.RISK_CONFIG['profit_retracement_pct']:.0%}"
f" (基准={rl}x)")
logger.info(f"信号: 买入阈值={config.SIGNAL_THRESHOLDS.get('buy_threshold',1.5)} | "
f"卖出阈值={config.SIGNAL_THRESHOLDS.get('sell_threshold',-1.5)}")
logger.info(f"仓位: 最多多={config.REALTIME_CONFIG['max_long_positions']} 最多空={config.REALTIME_CONFIG['max_short_positions']} | "
f"超时平仓={'' if config.RISK_CONFIG_CONST.get('enable_time_based_exit',True) else ''}")
logger.info(f"对冲: 信号对冲={'' if config.REALTIME_CONFIG.get('hedge_enabled',False) else ''} | "
f"锁仓={'' if config.REALTIME_CONFIG.get('lock_enabled',False) else ''}")
logger.info("=" * 50)
return True
@@ -60,9 +61,13 @@ class RealtimeTrader:
def _run_cycle(self):
try:
self._cycle_count += 1
# ★ 热加载配置:cron 优化器改完 config.py 后自动生效
config.reload()
self.risk_controller.sync_state()
current_price = self.data_provider.get_current_price(SYMBOL)
current_price = self.data_provider.get_current_price(config.SYMBOL)
if not current_price:
return
@@ -75,8 +80,8 @@ class RealtimeTrader:
weights.append(weight)
weighted_signal_sum = sum(s * w for s, w in zip(signals, weights))
buy_threshold = SIGNAL_THRESHOLDS.get('buy_threshold', 1.5)
sell_threshold = SIGNAL_THRESHOLDS.get('sell_threshold', -1.5)
buy_threshold = config.SIGNAL_THRESHOLDS.get('buy_threshold', 1.5)
sell_threshold = config.SIGNAL_THRESHOLDS.get('sell_threshold', -1.5)
direction = None
if weighted_signal_sum > buy_threshold:
@@ -84,6 +89,33 @@ class RealtimeTrader:
elif weighted_signal_sum < sell_threshold:
direction = "sell"
# ★ 同向门槛递增:已有N单同向时,第N+1单需要更强信号
if direction:
# ★ 适应度门槛:回测亏钱就不开新单
last_fitness = getattr(config, 'LAST_OPTIMIZATION_FITNESS', 0)
min_fitness = config.RISK_CONFIG.get('min_backtest_fitness', 250)
if last_fitness < min_fitness:
if self._cycle_count % 30 == 0:
logger.warning(f"⚠️ 回测适应度{last_fitness:.0f}<{min_fitness},暂停开仓(监控持仓中)")
direction = None
if direction:
alpha = config.RISK_CONFIG.get('entry_escalation_alpha', 1.2)
pm = self.risk_controller.position_manager
n_long = sum(1 for p in pm.positions if p['position_type'] == 'long')
n_short = sum(1 for p in pm.positions if p['position_type'] == 'short')
if direction == 'buy':
adjusted = buy_threshold * (alpha ** n_long)
if weighted_signal_sum <= adjusted:
logger.debug(f"BUY信号{weighted_signal_sum:.2f}<调整阈值{adjusted:.2f}(已有{n_long}多单), 忽略")
direction = None
elif direction == 'sell':
adjusted = sell_threshold * (alpha ** n_short)
if weighted_signal_sum >= adjusted:
logger.debug(f"SELL信号{weighted_signal_sum:.2f}>调整阈值{adjusted:.2f}(已有{n_short}空单), 忽略")
direction = None
# 只在信号触发时打印决策依据
if direction:
logger.info(f"⚡ 信号触发 | 加权={weighted_signal_sum:.2f} | "
@@ -124,6 +156,7 @@ class RealtimeTrader:
self.running = False
try:
if self.risk_controller:
self.risk_controller.position_manager.cleanup_peak_data()
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
self.risk_controller.save_trade_history(f"realtime_trades_{timestamp}")
summary = self.risk_controller.position_manager.get_trade_summary()
+4 -3
View File
@@ -10,7 +10,7 @@ import strategies # noqa: F401 — 触发策略注册
from logger import logger
from core.risk.market_state import MarketStateAnalyzer
from core.signal.registry import StrategyRegistry
from config import SYMBOL, TIMEFRAME
import config
class DynamicWeightManager:
@@ -27,7 +27,7 @@ class DynamicWeightManager:
def _ensure_strategies(self):
if not self._strategies_initialized:
self.registry.instantiate_all(SYMBOL, TIMEFRAME,
self.registry.instantiate_all(config.SYMBOL, config.TIMEFRAME,
data_provider=self.data_provider)
self._strategies_initialized = True
@@ -46,7 +46,8 @@ class DynamicWeightManager:
return result
def get_current_weights(self) -> dict:
"""获取当前实时权重"""
"""获取当前实时权重 — ★ 热加载配置"""
config.reload()
market_state, confidence = self.analyzer.get_market_state()
return self.analyzer.get_strategy_weights(market_state, confidence)
+11
View File
@@ -4,6 +4,7 @@
import sys
import os
import fcntl
from datetime import datetime
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
@@ -16,8 +17,18 @@ from config import (
)
from logger import setup_logger
LOCK_FILE = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), '.ea.lock')
def main():
# ★ 进程锁:防止多实例同时运行
lock_fd = os.open(LOCK_FILE, os.O_CREAT | os.O_RDWR, 0o644)
try:
fcntl.flock(lock_fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError:
print("❌ 已有 EA 实例在运行,退出")
os.close(lock_fd)
return
log_level = REALTIME_CONFIG.get('logging_level', 'INFO')
setup_logger(log_level)
+4 -3
View File
@@ -1,5 +1,6 @@
#!/bin/bash
# cron 环境修复:设置正确的 PYTHONPATH 后运行优化器
export PYTHONPATH="/home/songkl/.hermes/profiles/bot1/home/.local/lib/python3.13/site-packages:$PYTHONPATH"
# 每日自动优化 + 重启 EA
# 设置 PYTHONPATH 确保 cron 环境下能找到 user-site packages(如 deap, moocore 等)
export PYTHONPATH="$HOME/.local/lib/python$(python3 -c 'import sys; print(f"{sys.version_info.major}.{sys.version_info.minor}")')/site-packages:$PYTHONPATH"
cd /home/songkl/mt5_python_ea_suite
exec /usr/bin/python3 scripts/daily_optimize.py "$@"
exec python3 scripts/daily_optimize.py "$@"
+39 -23
View File
@@ -1,5 +1,5 @@
#!/usr/bin/env python3
"""每日自动优化 — 遗传算法跑参数 → 写入config → 重启实盘EA
"""每日自动优化 — 遗传算法跑参数 → 写入config → EA热加载生效(无需重启)
通过 cron 调用: python3 scripts/daily_optimize.py
"""
@@ -22,12 +22,12 @@ RESTART_SIGNAL = PROJECT_DIR / ".restart_signal"
def run_optimizer():
"""运行遗传算法优化,返回 best_params dict"""
"""运行遗传算法优化,返回 (best_params dict, fitness)"""
from execution.optimize import run_optimizer as _run
logger.info("🧬 开始遗传算法优化...")
best_params, fitness = _run()
logger.info(f"✅ 优化完成 适应度={fitness:.2f}")
return best_params
return best_params, fitness
def backup_config():
@@ -39,18 +39,13 @@ def backup_config():
logger.info(f"📦 已备份配置: {dst}")
def update_config(best_params: dict):
def update_config(best_params: dict, fitness: float = 0.0):
"""将优化结果写回 config.py"""
content = CONFIG_PATH.read_text(encoding="utf-8")
# ═══ RISK_CONFIG ═══
# ═══ RISK_CONFIG — 手动设定,不进优化器 ═══
# optimizer.py PARAM_SPACE 已移除风控基因,此 map 清空)
risk_map = {
"stop_loss_pct": "stop_loss_pct",
"profit_retracement_pct": "profit_retracement_pct",
"min_profit_for_trailing": "min_profit_for_trailing",
"take_profit_pct": "take_profit_pct",
"max_holding_minutes": "max_holding_minutes",
"min_profit_for_time_exit": "min_profit_for_time_exit",
}
for opt_key, cfg_key in risk_map.items():
if opt_key in best_params:
@@ -115,8 +110,11 @@ def update_config(best_params: dict):
"momentum_breakout_momentum_period": ("momentum_breakout", "momentum_period"),
# KDJStrategy
"kdj_period": ("kdj", "period"),
# TurtleStrategy
"turtle_period": ("turtle", "period"),
# SwingPointRetestStrategy
"swing_point_left_bars": ("swing_point", "left_bars"),
"swing_point_right_bars": ("swing_point", "right_bars"),
"swing_point_tolerance_pct": ("swing_point", "tolerance_pct"),
"swing_point_num_swings": ("swing_point", "num_swings"),
# DailyBreakoutStrategy
"daily_breakout_bars_count": ("daily_breakout", "bars_count"),
# WaveTheoryStrategy
@@ -151,7 +149,7 @@ def update_config(best_params: dict):
"weight_MomentumBreakoutStrategy": "momentum_breakout",
"weight_MACDStrategy": "macd",
"weight_KDJStrategy": "kdj",
"weight_TurtleStrategy": "turtle",
"weight_SwingPointRetestStrategy": "swing_point",
"weight_DailyBreakoutStrategy": "daily_breakout",
"weight_WaveTheoryStrategy": "wave_theory",
}
@@ -213,6 +211,13 @@ def update_config(best_params: dict):
content
)
# ═══ LAST_OPTIMIZATION_FITNESS ═══
content = re.sub(
r'LAST_OPTIMIZATION_FITNESS\s*=\s*[\d.\-e]+',
f'LAST_OPTIMIZATION_FITNESS = {fitness:.2f}',
content
)
CONFIG_PATH.write_text(content, encoding="utf-8")
logger.info("✏️ 配置已更新")
@@ -222,11 +227,24 @@ def restart_ea():
import subprocess
logger.info("🔄 重启 EA...")
# 杀旧进程
# 杀旧进程,等锁释放再启动新的
subprocess.run(["pkill", "-f", "python.*run/realtime.py"], capture_output=True)
import time; time.sleep(2)
subprocess.run(["pkill", "-9", "-f", "python.*run/realtime.py"], capture_output=True)
time.sleep(1)
# 确认旧进程已死 + 锁已释放(轮询最多等 5 秒)
lock_file = PROJECT_DIR / ".ea.lock"
import fcntl as _fcntl
for _ in range(10):
time.sleep(0.5)
try:
fd = os.open(str(lock_file), os.O_RDONLY)
_fcntl.flock(fd, _fcntl.LOCK_EX | _fcntl.LOCK_NB)
os.close(fd) # 立即释放,只是测试
break
except (BlockingIOError, OSError):
pass
else:
logger.warning("⚠️ 旧进程锁未释放,强制启动(旧进程可能僵死)")
# 清空旧交易记录
for f in PROJECT_DIR.glob("realtime_trades_*"):
@@ -259,20 +277,18 @@ def main():
# 2. 运行优化
try:
best_params = run_optimizer()
best_params, fitness = run_optimizer()
except Exception as e:
logger.error(f"优化失败: {e}")
import traceback
traceback.print_exc()
return 1
# 3. 写入 config.py
update_config(best_params)
# 3. 写入 config.py(含适应度)
update_config(best_params, fitness)
# 4. 触发重启
restart_ea()
logger.info("✅ 每日优化流程完成")
# 4. ★ EA 通过 config.reload() 自动热加载,无需重启
logger.info("✅ 每日优化流程完成 (EA 将在下个周期自动读取新配置)")
return 0
+61
View File
@@ -0,0 +1,61 @@
#!/usr/bin/env python3
"""开市后全平持仓 + 清峰值 + 重启 EA"""
import requests, json, time, sys, os, subprocess
API = "http://192.168.1.5:5555/api"
SYMBOL = "XAUUSDz"
PROJECT_DIR = "/home/songkl/mt5_python_ea_suite"
def log(msg):
print(f"[{time.strftime('%H:%M:%S')}] {msg}")
# 1. 等市场开市(重试最多30次,间隔10秒)
log("等待市场开市...")
for i in range(30):
try:
r = requests.get(f"{API}/account", timeout=5)
if r.status_code == 200:
acct = r.json()
if acct.get("trade_allowed"):
log("✅ 市场已开市")
break
except:
pass
time.sleep(10)
else:
log("❌ 等待超时,市场未开市")
sys.exit(1)
# 2. 获取所有持仓并平仓
log("获取持仓...")
r = requests.get(f"{API}/positions/{SYMBOL}")
positions = r.json().get("positions", [])
log(f"{len(positions)}")
total_pnl = 0
for p in positions:
resp = requests.post(f"{API}/close",
json={"ticket": p["ticket"], "symbol": SYMBOL, "volume": p["volume"]})
ok = resp.json().get("success", False)
total_pnl += p["profit"]
log(f"{'' if ok else ''} {p['ticket']} ${p['profit']:+.2f}")
log(f"总盈亏: ${total_pnl:+.2f}")
# 3. 清峰值数据
peak_file = os.path.join(PROJECT_DIR, "position_peaks.json")
if os.path.exists(peak_file):
os.remove(peak_file)
log("✅ 峰值数据已清除")
# 4. 杀掉旧 EA,启动新 EA
subprocess.run(["pkill", "-9", "-f", "python.*run/realtime.py"], capture_output=True)
time.sleep(2)
log_file = os.path.join(PROJECT_DIR, "logs", "strategy.log")
with open(log_file, "a") as f:
proc = subprocess.Popen(
[sys.executable, "run/realtime.py"],
cwd=PROJECT_DIR,
stdout=f, stderr=subprocess.STDOUT,
)
log(f"✅ EA 已重启 PID={proc.pid}")
+2 -2
View File
@@ -5,7 +5,7 @@ from strategies.rsi import RSIStrategy
from strategies.bollinger import BollingerStrategy
from strategies.macd import MACDStrategy
from strategies.kdj import KDJStrategy
from strategies.turtle import TurtleStrategy
from strategies.swing_point import SwingPointRetestStrategy
from strategies.mean_reversion import MeanReversionStrategy
from strategies.momentum_breakout import MomentumBreakoutStrategy
from strategies.daily_breakout import DailyBreakoutStrategy
@@ -18,7 +18,7 @@ _registry.register('rsi', RSIStrategy)
_registry.register('bollinger', BollingerStrategy)
_registry.register('macd', MACDStrategy)
_registry.register('kdj', KDJStrategy)
_registry.register('turtle', TurtleStrategy)
_registry.register('swing_point', SwingPointRetestStrategy)
_registry.register('mean_reversion', MeanReversionStrategy)
_registry.register('momentum_breakout', MomentumBreakoutStrategy)
_registry.register('daily_breakout', DailyBreakoutStrategy)
+177
View File
@@ -0,0 +1,177 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
SwingPointRetestStrategy — 摆动点回踩策略
基于前高前低的支撑阻力位:
1. 找局部摆动高/低点(比左右各N根K线更高/更低)
2. 价格接近前高 → SELL(阻力反弹),接近前低 → BUY(支撑反弹)
3. ★ 回踩确认:价格突破后回踩原支撑/阻力位 → 更高胜率的反转信号
黄金M1参数:左右各3根K线,容差 0.05%~0.1%(约2~5点)
"""
import pandas as pd
import numpy as np
from .base_strategy import BaseStrategy
from config import STRATEGY_CONFIG
class SwingPointRetestStrategy(BaseStrategy):
"""摆动点回踩策略 — 前高前低 + 回踩确认"""
def __init__(self, data_provider, symbol, timeframe,
left_bars=None, right_bars=None,
tolerance_pct=None, num_swings=None):
super().__init__(data_provider, symbol, timeframe)
config = STRATEGY_CONFIG.get('swing_point', {})
self.left_bars = left_bars if left_bars is not None else config.get('left_bars', 3)
self.right_bars = right_bars if right_bars is not None else config.get('right_bars', 3)
self.tolerance_pct = tolerance_pct if tolerance_pct is not None else config.get('tolerance_pct', 0.0008)
self.num_swings = num_swings if num_swings is not None else config.get('num_swings', 2)
self.lookback = max(self.left_bars + self.right_bars + 10, 50)
def _find_swing_points(self, df):
"""找摆动高低点"""
highs = df['high'].values
lows = df['low'].values
n = len(df)
L, R = self.left_bars, self.right_bars
swing_highs = [] # (index, price)
swing_lows = []
for i in range(L, n - R):
# swing_high: 比左边L根和右边R根都高
left_max = np.max(highs[i - L:i])
right_max = np.max(highs[i + 1:i + 1 + R])
if highs[i] > left_max and highs[i] > right_max:
swing_highs.append((i, highs[i]))
# swing_low: 比左边L根和右边R根都低
left_min = np.min(lows[i - L:i])
right_min = np.min(lows[i + 1:i + 1 + R])
if lows[i] < left_min and lows[i] < right_min:
swing_lows.append((i, lows[i]))
return swing_highs, swing_lows
def _calculate_indicators(self, df):
"""计算摆动点并标记到 DataFrame"""
df = df.copy()
swing_highs, swing_lows = self._find_swing_points(df)
return df, swing_highs, swing_lows
def _signal_from_swings(self, current_price, swings, is_high):
"""
判断当前价格是否接近摆动点
is_high=True: 接近前高 → 阻力 → SELL 信号
is_high=False: 接近前低 → 支撑 → BUY 信号
"""
if not swings:
return 0
tolerance = current_price * self.tolerance_pct
best_signal = 0
for idx, swing_price in swings[-self.num_swings:]:
distance_pct = abs(current_price - swing_price) / swing_price
if distance_pct <= self.tolerance_pct:
# 价格在摆动点容差范围内
# 信号强度 = 1 - (距离/容差),越近信号越强
strength = 1.0 - (distance_pct / self.tolerance_pct)
# ★ 回踩确认:价格曾突破过该摆动点
if is_high and current_price <= swing_price:
# 价格在阻力位下方 → 正常卖点
signal = -strength
elif not is_high and current_price >= swing_price:
# 价格在支撑位上方 → 正常买点
signal = strength
else:
# 价格在错误一侧,不给信号
continue
if abs(signal) > abs(best_signal):
best_signal = signal
# 归一化到 [-1, 1]
return max(-1.0, min(1.0, best_signal))
def generate_signal(self):
"""生成交易信号"""
rates = self.data_provider.get_historical_data(
self.symbol, self.timeframe, self.lookback
)
if rates is None or len(rates) < self.lookback:
return 0
df = pd.DataFrame(rates)
_, swing_highs, swing_lows = self._calculate_indicators(df)
current_price = df['close'].iloc[-1]
# 接近前高 → SELL
sell_signal = self._signal_from_swings(current_price, swing_highs, is_high=True)
# 接近前低 → BUY
buy_signal = self._signal_from_swings(current_price, swing_lows, is_high=False)
# 合并信号(sell为负,buy为正)
total = buy_signal + sell_signal # sell_signal 已经是负数
return max(-1.0, min(1.0, total))
def run_backtest(self, df):
"""回测模式 — 向量化计算信号"""
df = df.copy()
n = len(df)
L, R = self.left_bars, self.right_bars
tolerance = self.tolerance_pct
signals = pd.Series(0.0, index=df.index)
highs = df['high'].values
lows = df['low'].values
closes = df['close'].values
# 预计算摆动点
swing_high_mask = np.zeros(n, dtype=bool)
swing_low_mask = np.zeros(n, dtype=bool)
for i in range(L, n - R):
if highs[i] > np.max(highs[i - L:i]) and highs[i] > np.max(highs[i + 1:i + 1 + R]):
swing_high_mask[i] = True
if lows[i] < np.min(lows[i - L:i]) and lows[i] < np.min(lows[i + 1:i + 1 + R]):
swing_low_mask[i] = True
# 生成信号
for i in range(self.lookback, n):
current = closes[i]
# 找最近的摆动点
prev_highs = np.where(swing_high_mask[:i])[0]
prev_lows = np.where(swing_low_mask[:i])[0]
signal = 0.0
# 检查前高(阻力位)
for sh_idx in prev_highs[-self.num_swings:]:
sh_price = highs[sh_idx]
dist_pct = abs(current - sh_price) / sh_price
if dist_pct <= tolerance and current <= sh_price:
strength = 1.0 - (dist_pct / tolerance)
signal -= strength
break
# 检查前低(支撑位)
for sl_idx in prev_lows[-self.num_swings:]:
sl_price = lows[sl_idx]
dist_pct = abs(current - sl_price) / sl_price
if dist_pct <= tolerance and current >= sl_price:
strength = 1.0 - (dist_pct / tolerance)
signal += strength
break
signals.iloc[i] = max(-1.0, min(1.0, signal))
return signals