Add files via upload

This commit is contained in:
xiaochuan
2025-11-14 23:16:51 +00:00
committed by GitHub
parent 8c9371683e
commit 53c7aa9182
85 changed files with 6789 additions and 0 deletions
+37
View File
@@ -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)
+63
View File
@@ -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"}
+16
View File
@@ -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"}
+33
View File
@@ -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
+94
View File
@@ -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
+17
View File
@@ -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
+104
View File
@@ -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
+246
View File
@@ -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
+110
View File
@@ -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"}
+274
View File
@@ -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"}