mirror of
https://github.com/xavierchuan/FX-ML-Trading-Engine.git
synced 2026-08-25 23:08:04 +00:00
Add files via upload
This commit is contained in:
@@ -0,0 +1,37 @@
|
||||
# strategies/__init__.py
|
||||
from typing import Dict, Callable, Any
|
||||
|
||||
# 策略注册表:name -> class
|
||||
_REGISTRY: Dict[str, Callable[..., Any]] = {}
|
||||
|
||||
def register(name: str):
|
||||
"""用作装饰器:@register('sma_atr')"""
|
||||
def deco(cls):
|
||||
_REGISTRY[name] = cls
|
||||
return cls
|
||||
return deco
|
||||
|
||||
def _lazy_import_all():
|
||||
"""
|
||||
懒加载:首次 load_strategy 时再导入具体策略文件,
|
||||
这样不会因为循环依赖或路径问题导致注册表是空的。
|
||||
"""
|
||||
# 在这里逐个导入具体策略模块;导入发生时模块内的 @register 会把类放进 _REGISTRY
|
||||
from . import sma_atr # noqa: F401
|
||||
from . import regime_sma # noqa: F401
|
||||
from . import band_mean_revert # noqa: F401
|
||||
from . import bollinger_mean_revert # noqa: F401
|
||||
from . import ma_crossover # noqa: F401
|
||||
from . import momentum # noqa: F401
|
||||
from . import xgb_signal # noqa: F401
|
||||
|
||||
def load_strategy(name: str, **kwargs):
|
||||
# 第一次用时尝试懒加载,填充注册表
|
||||
if not _REGISTRY:
|
||||
_lazy_import_all()
|
||||
if name not in _REGISTRY:
|
||||
# 再尝试一次(防止用户后来才添加文件)
|
||||
_lazy_import_all()
|
||||
if name not in _REGISTRY:
|
||||
raise ValueError(f"Unknown strategy: {name}. Available: {list(_REGISTRY.keys())}")
|
||||
return _REGISTRY[name](**kwargs)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,63 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from . import register
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
@register("band_mean_revert")
|
||||
class BandMeanRevert(Strategy):
|
||||
"""
|
||||
简单区间/均值回归策略:
|
||||
- 使用慢均线 +/- ATR*mult 作为区间带;
|
||||
- 当价格跌破下带且 RSI 低于阈值时做多;
|
||||
- 当价格突破上带且 RSI 高于阈值时做空(可选)。
|
||||
- 价格回到均值或 RSI 归中时离场。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
band_atr_mult: float = 1.5,
|
||||
rsi_long: float = 35.0,
|
||||
rsi_short: float = 65.0,
|
||||
exit_rsi_mid: float = 50.0,
|
||||
allow_short: bool = True,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
band_atr_mult=band_atr_mult,
|
||||
rsi_long=rsi_long,
|
||||
rsi_short=rsi_short,
|
||||
exit_rsi_mid=exit_rsi_mid,
|
||||
allow_short=allow_short,
|
||||
)
|
||||
self.band_atr_mult = float(band_atr_mult)
|
||||
self.rsi_long = float(rsi_long)
|
||||
self.rsi_short = float(rsi_short)
|
||||
self.exit_rsi_mid = float(exit_rsi_mid)
|
||||
self.allow_short = bool(allow_short)
|
||||
|
||||
def on_bar(self, state: dict) -> dict:
|
||||
close = state.get("close")
|
||||
sma_slow = state.get("sma_slow")
|
||||
atr = state.get("curr_atr")
|
||||
rsi = state.get("rsi")
|
||||
position = state.get("position", 0)
|
||||
|
||||
if close is None or sma_slow is None or atr is None or rsi is None:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
upper = sma_slow + self.band_atr_mult * atr
|
||||
lower = sma_slow - self.band_atr_mult * atr
|
||||
|
||||
if position == 0:
|
||||
if close <= lower and rsi <= self.rsi_long:
|
||||
return {"action": "ENTER_LONG"}
|
||||
if self.allow_short and close >= upper and rsi >= self.rsi_short:
|
||||
return {"action": "ENTER_SHORT"}
|
||||
elif position > 0:
|
||||
if close >= sma_slow or rsi >= self.exit_rsi_mid:
|
||||
return {"action": "EXIT_LONG"}
|
||||
elif position < 0:
|
||||
if close <= sma_slow or rsi <= self.exit_rsi_mid:
|
||||
return {"action": "EXIT_SHORT"}
|
||||
|
||||
return {"action": "HOLD"}
|
||||
@@ -0,0 +1,16 @@
|
||||
# strategies/base.py
|
||||
import os
|
||||
import sys
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from typing import Dict, Any
|
||||
|
||||
class Strategy:
|
||||
"""
|
||||
只负责“发信号”,不做撮合和资金结算。
|
||||
on_bar 输入 state,输出 {"action": "ENTER_LONG"/"EXIT_LONG"/"HOLD"}。
|
||||
"""
|
||||
def __init__(self, **params: Any) -> None:
|
||||
self.params = params
|
||||
|
||||
def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {"action": "HOLD"}
|
||||
@@ -0,0 +1,78 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from . import register
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
@register("bollinger_mean_revert")
|
||||
class BollingerMeanRevert(Strategy):
|
||||
"""
|
||||
Simple Bollinger-band based mean reversion strategy skeleton.
|
||||
- Enters long when price falls `enter_z` standard deviations below the mean.
|
||||
- Enters short symmetrically (if allow_short).
|
||||
- Exits when price reverts back within `exit_z` standard deviations.
|
||||
This is meant to be combined with other sleeves in portfolio tests.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window: int = 50,
|
||||
num_std: float = 2.0,
|
||||
enter_z: float = 1.0,
|
||||
exit_z: float = 0.2,
|
||||
allow_short: bool = True,
|
||||
cooldown: int = 0,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
window=window,
|
||||
num_std=num_std,
|
||||
enter_z=enter_z,
|
||||
exit_z=exit_z,
|
||||
allow_short=allow_short,
|
||||
cooldown=cooldown,
|
||||
)
|
||||
self.window = int(window)
|
||||
self.num_std = float(num_std)
|
||||
self.enter_z = float(enter_z)
|
||||
self.exit_z = float(exit_z)
|
||||
self.allow_short = bool(allow_short)
|
||||
self.cooldown = int(cooldown or 0)
|
||||
self._next_entry_bar = 0
|
||||
|
||||
def on_bar(self, state: dict) -> dict:
|
||||
closes = state.get("close_history")
|
||||
bar_idx = state.get("bar_idx", 0)
|
||||
position = state.get("position", 0)
|
||||
if closes is None or len(closes) < self.window:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
window_data = np.array(closes[-self.window:], dtype=float)
|
||||
mean = window_data.mean()
|
||||
std = window_data.std(ddof=0)
|
||||
if std == 0:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
price = float(window_data[-1])
|
||||
z_score = (price - mean) / std
|
||||
|
||||
# Enforce cooldown between fresh entries
|
||||
if bar_idx < self._next_entry_bar and position == 0:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
if position == 0:
|
||||
if z_score <= -self.enter_z:
|
||||
self._next_entry_bar = bar_idx + self.cooldown
|
||||
return {"action": "ENTER_LONG"}
|
||||
if self.allow_short and z_score >= self.enter_z:
|
||||
self._next_entry_bar = bar_idx + self.cooldown
|
||||
return {"action": "ENTER_SHORT"}
|
||||
elif position > 0:
|
||||
if z_score >= -self.exit_z:
|
||||
return {"action": "EXIT_LONG"}
|
||||
else: # position < 0
|
||||
if z_score <= self.exit_z:
|
||||
return {"action": "EXIT_SHORT"}
|
||||
|
||||
return {"action": "HOLD"}
|
||||
@@ -0,0 +1,33 @@
|
||||
# FX_BACKTEST/strategies/ma_cross.py
|
||||
import os
|
||||
import sys
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from collections import deque
|
||||
from core.events import TickEvent, SignalEvent
|
||||
|
||||
class MACross:
|
||||
def __init__(self, q, symbol: str, short: int = 20, long: int = 50, size: float = 10000.0):
|
||||
assert short < long
|
||||
self.q = q
|
||||
self.symbol = symbol
|
||||
self.short_n, self.long_n = short, long
|
||||
self.short_win, self.long_win = deque(maxlen=short), deque(maxlen=long)
|
||||
self.pos = 0.0 # 当前方向(>0 多 / <0 空 / =0 空仓)
|
||||
self.size = size
|
||||
|
||||
def on_event(self, ev):
|
||||
if isinstance(ev, TickEvent) and ev.symbol == self.symbol:
|
||||
mid = (ev.bid + ev.ask) / 2.0
|
||||
self.short_win.append(mid)
|
||||
self.long_win.append(mid)
|
||||
if len(self.long_win) < self.long_n:
|
||||
return
|
||||
sma_s = sum(self.short_win) / len(self.short_win)
|
||||
sma_l = sum(self.long_win) / len(self.long_win)
|
||||
# 交叉信号
|
||||
if self.pos <= 0 and sma_s > sma_l: # 金叉 -> 做多
|
||||
self.q.put(SignalEvent(ev.ts, self.symbol, "LONG", self.size))
|
||||
self.pos = 1
|
||||
elif self.pos >= 0 and sma_s < sma_l: # 死叉 -> 做空
|
||||
self.q.put(SignalEvent(ev.ts, self.symbol, "SHORT", self.size))
|
||||
self.pos = -1
|
||||
@@ -0,0 +1,94 @@
|
||||
"""Simple moving-average crossover strategy registered for StrategyEngine combos."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Dict, Any
|
||||
|
||||
from . import register
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
@register("ma_crossover")
|
||||
class MovingAverageCrossover(Strategy):
|
||||
"""
|
||||
Emits ENTER/EXIT signals when a fast SMA crosses a slow SMA.
|
||||
|
||||
State inputs expected from StrategyEngine:
|
||||
- sma_fast / sma_slow
|
||||
- bar_idx (int)
|
||||
- position_units (float)
|
||||
- default_qty (float)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size_mult: float = 1.0,
|
||||
cooldown_bars: int = 0,
|
||||
exit_buffer_pct: float = 0.0,
|
||||
allow_short: bool = True,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
size_mult=size_mult,
|
||||
cooldown_bars=cooldown_bars,
|
||||
exit_buffer_pct=exit_buffer_pct,
|
||||
allow_short=allow_short,
|
||||
)
|
||||
self.size_mult = float(size_mult)
|
||||
self.cooldown_bars = int(max(0, cooldown_bars))
|
||||
self.exit_buffer_pct = float(max(0.0, exit_buffer_pct))
|
||||
self.allow_short = bool(allow_short)
|
||||
|
||||
self._prev_fast: float | None = None
|
||||
self._prev_slow: float | None = None
|
||||
self._block_until: int = 0
|
||||
|
||||
def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]:
|
||||
fast = state.get("sma_fast")
|
||||
slow = state.get("sma_slow")
|
||||
bar_idx = int(state.get("bar_idx", 0) or 0)
|
||||
position = float(state.get("position_units", 0.0) or 0.0)
|
||||
|
||||
if fast is None or slow is None:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
prev_fast = self._prev_fast
|
||||
prev_slow = self._prev_slow
|
||||
self._prev_fast = fast
|
||||
self._prev_slow = slow
|
||||
|
||||
if prev_fast is None or prev_slow is None:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
if bar_idx < self._block_until:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
buffer = self.exit_buffer_pct
|
||||
size = self._position_size(state)
|
||||
|
||||
crossed_up = prev_fast <= prev_slow and fast > slow
|
||||
crossed_down = prev_fast >= prev_slow and fast < slow
|
||||
|
||||
if crossed_up:
|
||||
self._block_until = bar_idx + self.cooldown_bars
|
||||
return {"action": "ENTER_LONG", "size": size}
|
||||
|
||||
if crossed_down and self.allow_short:
|
||||
self._block_until = bar_idx + self.cooldown_bars
|
||||
return {"action": "ENTER_SHORT", "size": size}
|
||||
|
||||
slow_with_buffer = slow * (1.0 + buffer)
|
||||
slow_lower = slow * (1.0 - buffer)
|
||||
|
||||
if position > 0 and fast < slow_lower:
|
||||
return {"action": "EXIT_LONG"}
|
||||
|
||||
if position < 0 and (fast > slow_with_buffer or not self.allow_short):
|
||||
return {"action": "EXIT_SHORT"}
|
||||
|
||||
return {"action": "HOLD"}
|
||||
|
||||
def _position_size(self, state: Dict[str, Any]) -> float | None:
|
||||
default_qty = state.get("default_qty")
|
||||
if default_qty is None:
|
||||
return None
|
||||
return float(default_qty) * self.size_mult
|
||||
@@ -0,0 +1,17 @@
|
||||
import pandas as pd
|
||||
|
||||
def mean_reversion_strategy(df: pd.DataFrame, short_window=20, long_window=50):
|
||||
"""
|
||||
简单的均值回归策略(示例):
|
||||
- 价格高于长期均线则卖出
|
||||
- 价格低于长期均线则买入
|
||||
"""
|
||||
df = df.copy()
|
||||
df["short_ma"] = df["close"].rolling(window=short_window).mean()
|
||||
df["long_ma"] = df["close"].rolling(window=long_window).mean()
|
||||
|
||||
df["signal"] = 0
|
||||
df.loc[df["short_ma"] > df["long_ma"], "signal"] = 1 # 买入
|
||||
df.loc[df["short_ma"] < df["long_ma"], "signal"] = -1 # 卖出
|
||||
|
||||
return df
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Basic momentum breakout strategy usable inside StrategyEngine combos."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Dict, Any, Sequence
|
||||
|
||||
from . import register
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
@register("momentum_breakout")
|
||||
class MomentumBreakout(Strategy):
|
||||
"""
|
||||
Uses rate-of-change over a configurable lookback to enter in the dominant direction.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
lookback : int
|
||||
Number of bars between comparisons.
|
||||
enter_threshold : float
|
||||
Minimum absolute return (%) to trigger an entry (e.g. 0.002 = 20 bps).
|
||||
exit_threshold : float
|
||||
Momentum magnitude below which existing positions are flattened.
|
||||
size_mult : float
|
||||
Multiplier applied to StrategyEngine default_qty when sizing trades.
|
||||
allow_short : bool
|
||||
Whether to take short trades when momentum turns negative.
|
||||
cooldown_bars : int
|
||||
Minimum bars between successive entries.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
lookback: int = 24,
|
||||
enter_threshold: float = 0.0015,
|
||||
exit_threshold: float = 0.0005,
|
||||
size_mult: float = 1.0,
|
||||
allow_short: bool = True,
|
||||
cooldown_bars: int = 0,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
lookback=lookback,
|
||||
enter_threshold=enter_threshold,
|
||||
exit_threshold=exit_threshold,
|
||||
size_mult=size_mult,
|
||||
allow_short=allow_short,
|
||||
cooldown_bars=cooldown_bars,
|
||||
)
|
||||
self.lookback = max(1, int(lookback))
|
||||
self.enter_threshold = float(enter_threshold)
|
||||
self.exit_threshold = float(exit_threshold)
|
||||
self.size_mult = float(size_mult)
|
||||
self.allow_short = bool(allow_short)
|
||||
self.cooldown_bars = int(max(0, cooldown_bars))
|
||||
|
||||
self._block_until: int = 0
|
||||
|
||||
def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]:
|
||||
closes = state.get("close_history")
|
||||
bar_idx = int(state.get("bar_idx", 0) or 0)
|
||||
position = float(state.get("position_units", 0.0) or 0.0)
|
||||
|
||||
if not self._has_enough_history(closes):
|
||||
return {"action": "HOLD"}
|
||||
|
||||
roc = self._rate_of_change(closes)
|
||||
|
||||
if position > 0 and roc < self.exit_threshold:
|
||||
return {"action": "EXIT_LONG"}
|
||||
if position < 0 and roc > -self.exit_threshold:
|
||||
return {"action": "EXIT_SHORT"}
|
||||
|
||||
if bar_idx < self._block_until:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
size = self._position_size(state)
|
||||
if roc >= self.enter_threshold:
|
||||
self._block_until = bar_idx + self.cooldown_bars
|
||||
return {"action": "ENTER_LONG", "size": size}
|
||||
if roc <= -self.enter_threshold and self.allow_short:
|
||||
self._block_until = bar_idx + self.cooldown_bars
|
||||
return {"action": "ENTER_SHORT", "size": size}
|
||||
|
||||
return {"action": "HOLD"}
|
||||
|
||||
def _has_enough_history(self, closes: Any) -> bool:
|
||||
if closes is None:
|
||||
return False
|
||||
if isinstance(closes, Sequence):
|
||||
return len(closes) > self.lookback
|
||||
return False
|
||||
|
||||
def _rate_of_change(self, closes: Sequence[float]) -> float:
|
||||
recent = float(closes[-1])
|
||||
past = float(closes[-(self.lookback + 1)])
|
||||
if past == 0:
|
||||
return 0.0
|
||||
return (recent - past) / past
|
||||
|
||||
def _position_size(self, state: Dict[str, Any]) -> float | None:
|
||||
default_qty = state.get("default_qty")
|
||||
if default_qty is None:
|
||||
return None
|
||||
return float(default_qty) * self.size_mult
|
||||
@@ -0,0 +1,246 @@
|
||||
# strategies/regime_sma.py
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from . import register
|
||||
from .base import Strategy
|
||||
from .sma_atr import SmaAtr
|
||||
|
||||
|
||||
@register("regime_sma")
|
||||
class RegimeSMAStrategy(Strategy):
|
||||
"""
|
||||
Wraps the SMA+ATR strategy with a simple regime filter.
|
||||
|
||||
- 在趋势 regime(由 StrategyEngine 提供)下,沿用 SmaAtr 信号。
|
||||
- 在震荡 regime 下,可选择保持空仓或使用简单的 RSI 均值回归。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
trend_params: dict | None = None,
|
||||
range_mode: str = "flat",
|
||||
range_rsi_high: float = 75.0,
|
||||
range_rsi_low: float = 25.0,
|
||||
trend_min_bars: int = 0,
|
||||
atr_percentile_min: float | None = None,
|
||||
size_tiers: list | None = None,
|
||||
base_size_mult: float = 1.0,
|
||||
htf_alignment: bool = False,
|
||||
htf_rsi_range: tuple[float, float] | list | None = None,
|
||||
risk_rules: list | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.range_mode = (range_mode or "flat").lower()
|
||||
self.range_rsi_high = float(range_rsi_high)
|
||||
self.range_rsi_low = float(range_rsi_low)
|
||||
self.range_exit_mid = (self.range_rsi_high + self.range_rsi_low) / 2.0
|
||||
self.trend_min_bars = int(trend_min_bars or 0)
|
||||
self.atr_percentile_min = atr_percentile_min if atr_percentile_min is None else float(atr_percentile_min)
|
||||
self.base_size_mult = float(base_size_mult or 1.0)
|
||||
self.size_tiers = self._normalize_tiers(size_tiers)
|
||||
self.htf_alignment = bool(htf_alignment)
|
||||
if htf_rsi_range:
|
||||
lo, hi = htf_rsi_range
|
||||
self.htf_rsi_min = float(lo)
|
||||
self.htf_rsi_max = float(hi)
|
||||
else:
|
||||
self.htf_rsi_min = 30.0
|
||||
self.htf_rsi_max = 70.0
|
||||
self.risk_rules = risk_rules or []
|
||||
|
||||
trend_params = trend_params or {}
|
||||
self.trend_strategy = SmaAtr(**trend_params)
|
||||
|
||||
def on_bar(self, state: dict):
|
||||
cooldown = self._check_risk_rules(state)
|
||||
if cooldown:
|
||||
return {"action": "HOLD", "cooldown_bars": cooldown}
|
||||
|
||||
regime_label = state.get("regime_label", "unknown")
|
||||
trend_streak = int(state.get("regime_trend_bars", 0) or 0)
|
||||
atr_percentile = state.get("atr_percentile")
|
||||
base_signal = self.trend_strategy.on_bar(state) or {}
|
||||
action = base_signal.get("action", "HOLD")
|
||||
|
||||
if regime_label == "trend":
|
||||
if self.trend_min_bars and trend_streak < self.trend_min_bars:
|
||||
regime_label = "range"
|
||||
elif self.atr_percentile_min is not None:
|
||||
if atr_percentile is None or atr_percentile < self.atr_percentile_min:
|
||||
regime_label = "range"
|
||||
if regime_label == "trend" and action.startswith("ENTER"):
|
||||
if not self._passes_htf_filter(action, state):
|
||||
regime_label = "range"
|
||||
else:
|
||||
sized = dict(base_signal)
|
||||
sized["size"] = self._position_size(action, state)
|
||||
return sized
|
||||
elif regime_label == "trend":
|
||||
return base_signal
|
||||
|
||||
# 在趋势 regime,直接沿用趋势策略的信号
|
||||
if regime_label == "trend":
|
||||
return base_signal
|
||||
|
||||
# 非趋势 regime:允许趋势策略发出的平仓指令生效,但屏蔽入场
|
||||
if action.startswith("EXIT"):
|
||||
return base_signal
|
||||
|
||||
# 根据 range_mode 决定行为
|
||||
range_action = self._range_signal(state)
|
||||
if range_action:
|
||||
return {"action": range_action}
|
||||
|
||||
# 默认保持空仓
|
||||
if state.get("position"):
|
||||
# 持仓状态下交给风控(risk exit)或趋势策略的平仓指令处理
|
||||
return {"action": "HOLD"}
|
||||
return {"action": "HOLD"}
|
||||
|
||||
def _range_signal(self, state: dict) -> str | None:
|
||||
if self.range_mode != "mean_revert":
|
||||
return None
|
||||
rsi = state.get("rsi")
|
||||
if rsi is None:
|
||||
return None
|
||||
|
||||
position = state.get("position", 0)
|
||||
if position == 0:
|
||||
if rsi >= self.range_rsi_high:
|
||||
return "ENTER_SHORT"
|
||||
if rsi <= self.range_rsi_low:
|
||||
return "ENTER_LONG"
|
||||
elif position > 0 and rsi >= self.range_exit_mid:
|
||||
return "EXIT_LONG"
|
||||
elif position < 0 and rsi <= self.range_exit_mid:
|
||||
return "EXIT_SHORT"
|
||||
return None
|
||||
|
||||
def _passes_htf_filter(self, action: str, state: dict) -> bool:
|
||||
if not self.htf_alignment:
|
||||
return True
|
||||
htf_ema = state.get("htf_ema")
|
||||
if htf_ema is None:
|
||||
return False
|
||||
close = state.get("close")
|
||||
if close is None:
|
||||
return False
|
||||
if action == "ENTER_LONG" and close < htf_ema:
|
||||
return False
|
||||
if action == "ENTER_SHORT" and close > htf_ema:
|
||||
return False
|
||||
htf_rsi = state.get("htf_rsi")
|
||||
if htf_rsi is not None:
|
||||
if htf_rsi < self.htf_rsi_min or htf_rsi > self.htf_rsi_max:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _normalize_tiers(self, tiers: list | None) -> list:
|
||||
if not tiers:
|
||||
return [{"size_mult": self.base_size_mult}]
|
||||
normalized = []
|
||||
for tier in tiers:
|
||||
if not isinstance(tier, dict):
|
||||
continue
|
||||
entry = tier.copy()
|
||||
entry["size_mult"] = float(entry.get("size_mult", 1.0))
|
||||
normalized.append(entry)
|
||||
if not normalized:
|
||||
normalized.append({"size_mult": self.base_size_mult})
|
||||
normalized.sort(key=lambda t: t.get("size_mult", 0), reverse=True)
|
||||
return normalized
|
||||
|
||||
def _position_size(self, action: str, state: dict) -> float:
|
||||
base_qty = float(state.get("default_qty") or 0.0)
|
||||
if base_qty <= 0:
|
||||
return 0.0
|
||||
atr_pct = state.get("atr_percentile")
|
||||
trend_strength = state.get("trend_strength")
|
||||
streak = int(state.get("regime_trend_bars", 0) or 0)
|
||||
for tier in self.size_tiers:
|
||||
if self._tier_matches(tier, atr_pct, trend_strength, streak, action):
|
||||
return base_qty * tier.get("size_mult", 1.0)
|
||||
return base_qty * self.base_size_mult
|
||||
|
||||
def _tier_matches(self, tier: dict, atr_pct, trend_strength, streak: int, action: str) -> bool:
|
||||
min_atr = tier.get("min_atr_pct")
|
||||
max_atr = tier.get("max_atr_pct")
|
||||
min_strength = tier.get("min_trend_strength")
|
||||
min_streak = tier.get("min_trend_bars")
|
||||
allow_short = tier.get("allow_short")
|
||||
allow_long = tier.get("allow_long")
|
||||
if min_atr is not None:
|
||||
if atr_pct is None or atr_pct < float(min_atr):
|
||||
return False
|
||||
if max_atr is not None and atr_pct is not None:
|
||||
if atr_pct > float(max_atr):
|
||||
return False
|
||||
if min_strength is not None:
|
||||
if trend_strength is None or trend_strength < float(min_strength):
|
||||
return False
|
||||
if min_streak is not None and streak < int(min_streak):
|
||||
return False
|
||||
if action == "ENTER_LONG" and allow_long is False:
|
||||
return False
|
||||
if action == "ENTER_SHORT" and allow_short is False:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _check_risk_rules(self, state: dict) -> int:
|
||||
if not self.risk_rules:
|
||||
return 0
|
||||
atr_pct = state.get("atr_percentile")
|
||||
ts = self._to_datetime(state.get("ts"))
|
||||
for rule in self.risk_rules:
|
||||
rtype = rule.get("type")
|
||||
if rtype == "atr_percentile":
|
||||
min_v = rule.get("min")
|
||||
max_v = rule.get("max")
|
||||
triggered = False
|
||||
if min_v is not None:
|
||||
if atr_pct is None or atr_pct < float(min_v):
|
||||
triggered = True
|
||||
if max_v is not None and atr_pct is not None and atr_pct > float(max_v):
|
||||
triggered = True
|
||||
if triggered:
|
||||
return int(rule.get("cooldown_bars", 0) or 0)
|
||||
elif rtype == "calendar":
|
||||
if ts is None:
|
||||
continue
|
||||
dates = rule.get("dates") or []
|
||||
date_str = ts.strftime("%Y-%m-%d")
|
||||
if date_str in dates:
|
||||
return int(rule.get("cooldown_bars", 0) or 0)
|
||||
windows = rule.get("windows") or []
|
||||
for window in windows:
|
||||
start = self._to_datetime(window.get("start"))
|
||||
end = self._to_datetime(window.get("end"))
|
||||
if start and end and start <= ts <= end:
|
||||
return int(rule.get("cooldown_bars", 0) or 0)
|
||||
elif rtype == "time_window":
|
||||
start = self._to_datetime(rule.get("start"))
|
||||
end = self._to_datetime(rule.get("end"))
|
||||
if ts and start and end and start <= ts <= end:
|
||||
return int(rule.get("cooldown_bars", 0) or 0)
|
||||
return 0
|
||||
|
||||
def _to_datetime(self, value):
|
||||
if value is None:
|
||||
return None
|
||||
if hasattr(value, "to_pydatetime"):
|
||||
try:
|
||||
return value.to_pydatetime()
|
||||
except Exception:
|
||||
pass
|
||||
if isinstance(value, datetime):
|
||||
return value
|
||||
text = str(value)
|
||||
if not text:
|
||||
return None
|
||||
try:
|
||||
return datetime.fromisoformat(text.replace("Z", "+00:00"))
|
||||
except Exception:
|
||||
return None
|
||||
@@ -0,0 +1,110 @@
|
||||
# strategies/sma_atr.py
|
||||
from .base import Strategy
|
||||
from . import register
|
||||
|
||||
@register("sma_atr")
|
||||
class SmaAtr(Strategy):
|
||||
"""
|
||||
参数(与 runner 对齐):
|
||||
- long_only_above_slow: bool
|
||||
- allow_short: bool
|
||||
- short_only_below_slow: bool
|
||||
- slope_lookback: int
|
||||
- cooldown: int
|
||||
- fast_win: int
|
||||
- slow_win: int
|
||||
- atr_sl / atr_tp / atr_window(仅用于记录,不在策略层计算)
|
||||
"""
|
||||
def on_bar(self, state):
|
||||
c = state["close"]
|
||||
position = state["position"]
|
||||
rsi = state.get("rsi")
|
||||
sma_fast = state.get("sma_fast")
|
||||
sma_slow = state.get("sma_slow")
|
||||
bar_idx = state["bar_idx"]
|
||||
next_entry_bar_idx_long = state.get("next_entry_bar_idx_long", 0)
|
||||
next_entry_bar_idx_short = state.get("next_entry_bar_idx_short", 0)
|
||||
sma_fast_hist = state["sma_fast_hist"] # deque,最近值在右侧
|
||||
|
||||
fast_win = self.params.get("fast_win")
|
||||
slow_win = self.params.get("slow_win")
|
||||
long_only_above_slow = self.params.get("long_only_above_slow", False)
|
||||
slope_lookback = self.params.get("slope_lookback", 0)
|
||||
cooldown = self.params.get("cooldown", 0)
|
||||
allow_short = self.params.get("allow_short", True)
|
||||
short_only_below_slow = self.params.get("short_only_below_slow", False)
|
||||
rsi_long_thresh = self.params.get("rsi_long_thresh")
|
||||
rsi_short_thresh = self.params.get("rsi_short_thresh")
|
||||
|
||||
# 均线就绪才判断
|
||||
if sma_fast is None or sma_slow is None:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
go_long = (position == 0) and (sma_fast > sma_slow)
|
||||
exit_long = (position == 1) and (sma_fast < sma_slow)
|
||||
|
||||
# 仅做多需在慢均线上方
|
||||
if long_only_above_slow and go_long:
|
||||
if not (c > sma_slow):
|
||||
go_long = False
|
||||
|
||||
# fast 斜率确认
|
||||
if slope_lookback and go_long:
|
||||
if len(sma_fast_hist) > slope_lookback:
|
||||
# 最近一个值与 L 根前比较
|
||||
if not (sma_fast_hist[-1] > sma_fast_hist[-1 - slope_lookback]):
|
||||
go_long = False
|
||||
else:
|
||||
go_long = False
|
||||
|
||||
# 冷却
|
||||
if cooldown and go_long:
|
||||
if bar_idx < next_entry_bar_idx_long:
|
||||
go_long = False
|
||||
|
||||
# RSI 过滤(如果配置了阈值且 RSI 可用)
|
||||
if go_long and rsi_long_thresh is not None:
|
||||
if rsi is None:
|
||||
go_long = False
|
||||
else:
|
||||
if not (rsi > float(rsi_long_thresh)):
|
||||
go_long = False
|
||||
|
||||
go_short = False
|
||||
exit_short = False
|
||||
if allow_short:
|
||||
go_short = (position == 0) and (sma_fast < sma_slow)
|
||||
exit_short = (position == -1) and (sma_fast > sma_slow)
|
||||
|
||||
if short_only_below_slow and go_short:
|
||||
if not (c < sma_slow):
|
||||
go_short = False
|
||||
|
||||
if slope_lookback and go_short:
|
||||
if len(sma_fast_hist) > slope_lookback:
|
||||
if not (sma_fast_hist[-1] < sma_fast_hist[-1 - slope_lookback]):
|
||||
go_short = False
|
||||
else:
|
||||
go_short = False
|
||||
|
||||
if cooldown and go_short:
|
||||
if bar_idx < next_entry_bar_idx_short:
|
||||
go_short = False
|
||||
|
||||
# RSI 过滤(空头侧)
|
||||
if go_short and rsi_short_thresh is not None:
|
||||
if rsi is None:
|
||||
go_short = False
|
||||
else:
|
||||
if not (rsi < float(rsi_short_thresh)):
|
||||
go_short = False
|
||||
|
||||
if exit_long:
|
||||
return {"action": "EXIT_LONG"}
|
||||
if exit_short:
|
||||
return {"action": "EXIT_SHORT"}
|
||||
if go_long:
|
||||
return {"action": "ENTER_LONG"}
|
||||
if go_short:
|
||||
return {"action": "ENTER_SHORT"}
|
||||
return {"action": "HOLD"}
|
||||
@@ -0,0 +1,274 @@
|
||||
"""XGBoost-based signal strategy (long-only v1).
|
||||
|
||||
Loads a trained Booster + feature list + thresholds and emits ENTER_LONG/EXIT_LONG
|
||||
decisions based on predicted probability compared to configured thresholds.
|
||||
|
||||
Registration name: "xgb_signal"
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
from . import register
|
||||
from .base import Strategy
|
||||
|
||||
|
||||
def _safe_get(d: Dict[str, Any], key: str, default=None):
|
||||
v = d.get(key, default)
|
||||
return v if v is not None else default
|
||||
|
||||
|
||||
def _rsi_from_series(arr: np.ndarray, period: int = 14) -> float | None:
|
||||
if arr.size < period + 1:
|
||||
return None
|
||||
diffs = np.diff(arr)
|
||||
gains = diffs[diffs > 0]
|
||||
losses = -diffs[diffs < 0]
|
||||
avg_gain = gains.mean() if gains.size > 0 else 0.0
|
||||
avg_loss = losses.mean() if losses.size > 0 else 0.0
|
||||
if avg_loss == 0.0 and avg_gain == 0.0:
|
||||
return 50.0
|
||||
if avg_loss == 0.0:
|
||||
return 100.0
|
||||
rs = avg_gain / avg_loss
|
||||
return 100.0 - (100.0 / (1.0 + rs))
|
||||
|
||||
|
||||
@register("xgb_signal")
|
||||
class XGBSignal(Strategy):
|
||||
def __init__(
|
||||
self,
|
||||
model_dir: Optional[str] = None,
|
||||
latest_ptr: str = "QuantResearch/artifacts/models/usdjpy_h1_xgb_latest.json",
|
||||
prob_long: Optional[float] = None,
|
||||
prob_exit: Optional[float] = None,
|
||||
size_mult: float = 1.0,
|
||||
cooldown_bars: int = 0,
|
||||
min_atr_pct: Optional[float] = None,
|
||||
low_atr_pct: Optional[float] = None,
|
||||
prob_long_low: Optional[float] = None,
|
||||
cooldown_low: Optional[int] = None,
|
||||
min_vol_24: Optional[float] = None,
|
||||
atr_relax_pct: Optional[float] = None,
|
||||
prob_long_relaxed: Optional[float] = None,
|
||||
cooldown_relaxed: Optional[int] = None,
|
||||
debug_log_hits: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
# Resolve model directory
|
||||
if not model_dir:
|
||||
ptr = Path(latest_ptr)
|
||||
if not ptr.exists():
|
||||
raise RuntimeError(f"latest.json not found: {ptr}")
|
||||
latest = json.loads(ptr.read_text(encoding="utf-8"))
|
||||
model_dir = latest.get("model_dir")
|
||||
if not model_dir:
|
||||
raise RuntimeError("latest.json missing 'model_dir'")
|
||||
self.model_dir = Path(model_dir)
|
||||
# Load artifacts
|
||||
self.feature_list: List[str] = json.loads((self.model_dir / "feature_list.json").read_text(encoding="utf-8"))
|
||||
thr = json.loads((self.model_dir / "thresholds.json").read_text(encoding="utf-8"))
|
||||
self.p_long = float(prob_long) if prob_long is not None else float(thr.get("p_long", 0.6))
|
||||
self.p_exit = float(prob_exit) if prob_exit is not None else float(thr.get("p_exit", 0.5))
|
||||
try:
|
||||
import xgboost as xgb
|
||||
except Exception as exc:
|
||||
raise RuntimeError("xgboost is required at runtime for xgb_signal.") from exc
|
||||
self._xgb = xgb
|
||||
self._booster = xgb.Booster()
|
||||
self._booster.load_model(str(self.model_dir / "model.json"))
|
||||
|
||||
self.size_mult = float(size_mult)
|
||||
self.cooldown_bars = int(max(0, cooldown_bars))
|
||||
self.cooldown_relaxed = int(max(0, cooldown_relaxed)) if cooldown_relaxed is not None else None
|
||||
self.min_atr_pct = float(min_atr_pct) if min_atr_pct is not None else None
|
||||
self.low_atr_pct = float(low_atr_pct) if low_atr_pct is not None else None
|
||||
self.prob_long_low = float(prob_long_low) if prob_long_low is not None else None
|
||||
self.cooldown_low = int(max(0, cooldown_low)) if cooldown_low is not None else None
|
||||
self.min_vol_24 = float(min_vol_24) if min_vol_24 is not None else None
|
||||
self.atr_relax_pct = float(atr_relax_pct) if atr_relax_pct is not None else None
|
||||
self.p_long_relaxed = float(prob_long_relaxed) if prob_long_relaxed is not None else None
|
||||
self.debug_log_hits = bool(debug_log_hits)
|
||||
self._block_until: int = 0
|
||||
self._debug_max = 0.0
|
||||
self._debug_none = 0
|
||||
|
||||
def _note_feature_miss(self, reason: str) -> None:
|
||||
self._debug_none += 1
|
||||
if self.debug_log_hits and self._debug_none <= 10:
|
||||
logger.warning(f"[xgb_signal] feature unavailable ({reason})")
|
||||
|
||||
def _features_from_state(self, state: Dict[str, Any]) -> tuple[Optional[np.ndarray], Optional[float]]:
|
||||
# Close history for returns/volatility
|
||||
ch = state.get("close_history")
|
||||
if ch is None:
|
||||
self._note_feature_miss("close_history missing")
|
||||
return None, None
|
||||
closes = np.asarray(ch, dtype=float)
|
||||
if closes.size < 80: # need at least slow window context
|
||||
self._note_feature_miss("insufficient history")
|
||||
return None, None
|
||||
close = float(state.get("close", closes[-1]))
|
||||
|
||||
# Returns & rolling vol
|
||||
def pct_change(arr: np.ndarray, k: int) -> float | None:
|
||||
if arr.size <= k:
|
||||
self._note_feature_miss(f"ret_{k} insufficient")
|
||||
return None, None
|
||||
a, b = arr[-k - 1], arr[-1]
|
||||
return (b - a) / a if a else None
|
||||
|
||||
ret_1 = pct_change(closes, 1)
|
||||
ret_3 = pct_change(closes, 3)
|
||||
ret_6 = pct_change(closes, 6)
|
||||
vol_24 = None
|
||||
if closes.size >= 25:
|
||||
rets = np.diff(closes[-25:]) / closes[-25:-1]
|
||||
vol_24 = float(np.std(rets)) if rets.size else None
|
||||
|
||||
# SMA diff
|
||||
sma_fast = _safe_get(state, "sma_fast")
|
||||
sma_slow = _safe_get(state, "sma_slow")
|
||||
if sma_fast is None or sma_slow is None:
|
||||
sma_fast = float(np.mean(closes[-20:])) if closes.size >= 20 else None
|
||||
sma_slow = float(np.mean(closes[-80:])) if closes.size >= 80 else None
|
||||
if sma_fast is None or sma_slow is None:
|
||||
self._note_feature_miss("sma missing")
|
||||
return None, None
|
||||
sma_diff = (float(sma_fast) - float(sma_slow)) / close if close else 0.0
|
||||
|
||||
# RSI (prefer engine state, else compute)
|
||||
rsi_val = state.get("rsi")
|
||||
if rsi_val is None:
|
||||
r = _rsi_from_series(closes, 14)
|
||||
rsi_val = r if r is not None else 50.0
|
||||
|
||||
# ATR normalized
|
||||
curr_atr = state.get("curr_atr")
|
||||
atr_norm = float(curr_atr) / close if (curr_atr is not None and close) else 0.0
|
||||
|
||||
# Time features from ts
|
||||
ts = state.get("ts")
|
||||
if ts is None:
|
||||
self._note_feature_miss("timestamp missing")
|
||||
return None, None
|
||||
try:
|
||||
import pandas as pd
|
||||
ts_pd = pd.Timestamp(ts)
|
||||
hour = float(ts_pd.hour)
|
||||
dow = float(ts_pd.dayofweek)
|
||||
except Exception:
|
||||
self._note_feature_miss("timestamp parse")
|
||||
return None, None
|
||||
hour_sin, hour_cos = np.sin(2 * np.pi * hour / 24.0), np.cos(2 * np.pi * hour / 24.0)
|
||||
dow_sin, dow_cos = np.sin(2 * np.pi * dow / 7.0), np.cos(2 * np.pi * dow / 7.0)
|
||||
|
||||
feat_map = {
|
||||
"ret_1": ret_1,
|
||||
"ret_3": ret_3,
|
||||
"ret_6": ret_6,
|
||||
"vol_24": vol_24,
|
||||
"sma_diff": sma_diff,
|
||||
"rsi": float(rsi_val),
|
||||
"atr_norm": atr_norm,
|
||||
"hour_sin": float(hour_sin),
|
||||
"hour_cos": float(hour_cos),
|
||||
"dow_sin": float(dow_sin),
|
||||
"dow_cos": float(dow_cos),
|
||||
}
|
||||
vec = []
|
||||
for name in self.feature_list:
|
||||
val = feat_map.get(name)
|
||||
if val is None or (isinstance(val, float) and (np.isnan(val) or np.isinf(val))):
|
||||
self._note_feature_miss(f"feature {name} invalid")
|
||||
return None, None
|
||||
vec.append(float(val))
|
||||
return np.asarray(vec, dtype=float), float(vol_24) if vol_24 is not None else None
|
||||
|
||||
def on_bar(self, state: Dict[str, Any]) -> Dict[str, Any]:
|
||||
bar_idx = int(state.get("bar_idx", 0) or 0)
|
||||
position_units = float(state.get("position_units", 0.0) or 0.0)
|
||||
default_qty = float(state.get("default_qty", 0.0) or 0.0)
|
||||
close = float(state.get("close", 0.0) or 0.0)
|
||||
|
||||
atr_pct = state.get("atr_percentile")
|
||||
if self.min_atr_pct is not None:
|
||||
if atr_pct is None or float(atr_pct) < self.min_atr_pct:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
# Determine per-bar entry threshold / cooldown after ATR gating
|
||||
effective_prob_long = self.p_long
|
||||
effective_cooldown = self.cooldown_bars
|
||||
if (
|
||||
self.low_atr_pct is not None
|
||||
and atr_pct is not None
|
||||
and float(atr_pct) < self.low_atr_pct
|
||||
):
|
||||
if self.prob_long_low is not None:
|
||||
effective_prob_long = self.prob_long_low
|
||||
if self.cooldown_low is not None:
|
||||
effective_cooldown = self.cooldown_low
|
||||
if (
|
||||
self.atr_relax_pct is not None
|
||||
and atr_pct is not None
|
||||
and float(atr_pct) >= self.atr_relax_pct
|
||||
):
|
||||
if self.p_long_relaxed is not None:
|
||||
effective_prob_long = self.p_long_relaxed
|
||||
if self.cooldown_relaxed is not None:
|
||||
effective_cooldown = self.cooldown_relaxed
|
||||
|
||||
# Cooldown gate (updated per bar)
|
||||
if bar_idx < self._block_until:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
feats, vol_24 = self._features_from_state(state)
|
||||
if feats is None:
|
||||
return {"action": "HOLD"}
|
||||
if self.min_vol_24 is not None:
|
||||
if vol_24 is None or vol_24 < self.min_vol_24:
|
||||
return {"action": "HOLD"}
|
||||
|
||||
dmat = self._xgb.DMatrix(feats.reshape(1, -1), feature_names=self.feature_list)
|
||||
p_up = float(self._booster.predict(dmat)[0])
|
||||
if p_up > self._debug_max:
|
||||
self._debug_max = p_up
|
||||
if self.debug_log_hits:
|
||||
logger.info(
|
||||
"[xgb_signal] new max prob %.4f (ts=%s close=%.5f position=%s)",
|
||||
p_up,
|
||||
state.get("ts"),
|
||||
close,
|
||||
position_units,
|
||||
)
|
||||
|
||||
# Long-only logic
|
||||
if position_units == 0.0:
|
||||
if p_up >= effective_prob_long and default_qty > 0.0:
|
||||
self._block_until = bar_idx + effective_cooldown
|
||||
if self.debug_log_hits:
|
||||
logger.info(
|
||||
"[xgb_signal] ENTER signal p=%.4f (thr=%.4f) ts=%s",
|
||||
p_up,
|
||||
effective_prob_long,
|
||||
state.get("ts"),
|
||||
)
|
||||
return {"action": "ENTER_LONG", "size": default_qty * self.size_mult}
|
||||
return {"action": "HOLD"}
|
||||
else:
|
||||
if p_up < self.p_exit:
|
||||
if self.debug_log_hits:
|
||||
logger.info(
|
||||
"[xgb_signal] EXIT signal p=%.4f (thr=%.4f) ts=%s",
|
||||
p_up,
|
||||
self.p_exit,
|
||||
state.get("ts"),
|
||||
)
|
||||
return {"action": "EXIT_LONG"}
|
||||
return {"action": "HOLD"}
|
||||
Reference in New Issue
Block a user