start experiment 57 (exp/57-signal-quality-gate-gate-trades-based-on)
This commit is contained in:
@@ -0,0 +1,215 @@
|
||||
"""Regime-gate TopkDropout strategy.
|
||||
|
||||
Subclass of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy`` that
|
||||
holds the book (issues NO orders) while a regime detector says the market is in
|
||||
an unfavorable state. When the gate is open it behaves exactly like the
|
||||
reference TopkDropoutStrategy.
|
||||
|
||||
Three detector types are supported (all causal — no lookahead):
|
||||
|
||||
* ``dispersion``: cross-sectional standard deviation of 22-day rolling returns
|
||||
across the universe. Gate closes when CS dispersion < threshold (low
|
||||
dispersion means the spread between winners and losers is too narrow for
|
||||
TopkDropout to exploit).
|
||||
* ``vol``: cross-sectional mean of 22-day rolling realized volatility. Gate
|
||||
closes when avg vol is outside a band ``[vol_low, vol_high]`` (strategy
|
||||
needs moderate vol — too calm or too turbulent both hurt).
|
||||
* ``hmm``: pre-computed HMM posterior for regime 1 (``sp_hmm_p_regime1``).
|
||||
Gate closes when posterior < threshold (model is not confident the calm
|
||||
regime is active).
|
||||
|
||||
The gate is provided as a precomputed ``pd.Series`` of booleans indexed by
|
||||
datetime (True = trade allowed). The companion ``compute_regime_gate``
|
||||
function builds this series from lake bars; call it once before backtesting
|
||||
and pass the result as the ``regime_gate`` parameter.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from qlib.backtest.decision import TradeDecisionWO
|
||||
from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy
|
||||
|
||||
__all__ = ["RegimeGateTopkDropoutStrategy", "compute_regime_gate"]
|
||||
|
||||
|
||||
class RegimeGateTopkDropoutStrategy(TopkDropoutStrategy):
|
||||
"""TopkDropout with a regime-gate circuit breaker.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable,
|
||||
forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``.
|
||||
regime_gate : pd.Series — precomputed per-date gate (bool indexed by
|
||||
datetime). True = trade allowed, False = no orders. Missing dates
|
||||
default to open (trade allowed).
|
||||
"""
|
||||
|
||||
def __init__(self, *, regime_gate=None, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._regime_gate = regime_gate
|
||||
|
||||
def _gate_open(self, trade_start_time) -> bool:
|
||||
if self._regime_gate is None:
|
||||
return True
|
||||
ts = pd.Timestamp(trade_start_time)
|
||||
known = self._regime_gate[self._regime_gate.index <= ts]
|
||||
if len(known):
|
||||
return bool(known.iloc[-1])
|
||||
return True # default open if no history yet
|
||||
|
||||
def generate_trade_decision(self, execute_result=None):
|
||||
trade_step = self.trade_calendar.get_trade_step()
|
||||
trade_start_time, _ = self.trade_calendar.get_step_time(trade_step)
|
||||
if not self._gate_open(trade_start_time):
|
||||
return TradeDecisionWO([], self)
|
||||
return super().generate_trade_decision(execute_result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Precomputation helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def compute_regime_gate(
|
||||
detector: str,
|
||||
threshold: float = 0.0,
|
||||
*,
|
||||
lake_root: str = "",
|
||||
market: str = "US",
|
||||
start: str = "2015-01-03",
|
||||
end: str = "2026-08-19",
|
||||
vol_low: float = 0.0,
|
||||
vol_high: float = 999.0,
|
||||
hmm_field: str = "sp_hmm_p_regime1",
|
||||
) -> pd.Series:
|
||||
"""Build a per-date regime gate series from lake bars.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
detector : str — ``"dispersion"``, ``"vol"``, or ``"hmm"``.
|
||||
threshold : float — for ``dispersion``: min CS dispersion to allow trading.
|
||||
For ``hmm``: min HMM posterior to allow trading.
|
||||
Ignored for ``vol`` (uses ``vol_low``/``vol_high`` band instead).
|
||||
lake_root, market : str — lake location.
|
||||
start, end : str — date window.
|
||||
vol_low, vol_high : float — annualized vol band for the ``vol`` detector.
|
||||
hmm_field : str — HMM feature column name for the ``hmm`` detector.
|
||||
|
||||
Returns
|
||||
-------
|
||||
pd.Series — bool, indexed by datetime. True = trade allowed.
|
||||
"""
|
||||
from tac_qlib.data.config import LakeConfig, resolve_lake_root
|
||||
|
||||
cfg = LakeConfig(resolve_lake_root(lake_root or None), market)
|
||||
symbols = _universe_symbols(cfg)
|
||||
close_df, vol_df = _load_daily_bars(symbols, cfg, start, end)
|
||||
if close_df.empty:
|
||||
return pd.Series(dtype=bool)
|
||||
|
||||
if detector == "dispersion":
|
||||
return _dispersion_gate(close_df, threshold)
|
||||
elif detector == "vol":
|
||||
return _vol_gate(close_df, vol_low, vol_high)
|
||||
elif detector == "hmm":
|
||||
return _hmm_gate(cfg, symbols, threshold, start, end, hmm_field)
|
||||
else:
|
||||
raise ValueError(f"Unknown detector: {detector!r}")
|
||||
|
||||
|
||||
def _universe_symbols(cfg) -> list:
|
||||
"""Read symbols from the lake symbols.parquet."""
|
||||
import pathlib
|
||||
|
||||
sp = cfg.lake_root / "symbols.parquet"
|
||||
if sp.exists():
|
||||
df = pd.read_parquet(sp)
|
||||
col = "symbol" if "symbol" in df.columns else df.columns[0]
|
||||
return sorted(df[col].astype(str).str.upper().tolist())
|
||||
return []
|
||||
|
||||
|
||||
def _load_daily_bars(symbols, cfg, start, end):
|
||||
"""Load daily close prices for all symbols into a wide DataFrame."""
|
||||
closes = {}
|
||||
vols = {}
|
||||
for sym in symbols:
|
||||
p = cfg.bar_path("1d", sym)
|
||||
if not p.exists():
|
||||
continue
|
||||
try:
|
||||
df = pd.read_parquet(p)
|
||||
except Exception:
|
||||
continue
|
||||
if not len(df):
|
||||
continue
|
||||
tcol = df["t"] if "t" in df.columns else df["date"]
|
||||
ts = pd.to_datetime(tcol)
|
||||
df = df.assign(_t=ts).set_index("_t").sort_index()
|
||||
df = df.loc[start:end]
|
||||
if len(df) < 22:
|
||||
continue
|
||||
closes[sym] = df["c"]
|
||||
if "v" in df.columns:
|
||||
vols[sym] = df["v"]
|
||||
close_df = pd.DataFrame(closes)
|
||||
vol_df = pd.DataFrame(vols) if vols else None
|
||||
return close_df, vol_df
|
||||
|
||||
|
||||
def _dispersion_gate(close_df, threshold):
|
||||
"""Cross-sectional dispersion of 22-day rolling returns."""
|
||||
if close_df.empty or close_df.shape[1] < 2:
|
||||
return pd.Series(dtype=bool)
|
||||
ret = close_df.pct_change(22)
|
||||
cs_disp = ret.std(axis=1)
|
||||
gate = cs_disp >= threshold
|
||||
gate.iloc[:22] = True # warmup: allow trading
|
||||
return gate
|
||||
|
||||
|
||||
def _vol_gate(close_df, vol_low, vol_high):
|
||||
"""Cross-sectional mean of 22-day rolling realized vol."""
|
||||
if close_df.empty or close_df.shape[1] < 2:
|
||||
return pd.Series(dtype=bool)
|
||||
import numpy as np
|
||||
log_ret = np.log(close_df / close_df.shift(1))
|
||||
rv22 = log_ret.rolling(22).std() * (252 ** 0.5)
|
||||
cs_mean_vol = rv22.mean(axis=1)
|
||||
gate = (cs_mean_vol >= vol_low) & (cs_mean_vol <= vol_high)
|
||||
gate.iloc[:22] = True # warmup
|
||||
return gate
|
||||
|
||||
|
||||
def _hmm_gate(cfg, symbols, threshold, start, end, hmm_field):
|
||||
"""HMM regime posterior gate from persisted SP features."""
|
||||
feat_root = cfg.lake_root / "features"
|
||||
all_posteriors = {}
|
||||
for sym in symbols:
|
||||
# check both ta and sp family paths
|
||||
for family in ("sp", "ta"):
|
||||
p = feat_root / f"market=US" / f"timeframe=1d" / f"family={family}" / f"symbol={sym}.parquet"
|
||||
if not p.exists():
|
||||
continue
|
||||
try:
|
||||
df = pd.read_parquet(p)
|
||||
except Exception:
|
||||
continue
|
||||
if hmm_field not in df.columns:
|
||||
continue
|
||||
tcol = df["t"] if "t" in df.columns else df["date"]
|
||||
ts = pd.to_datetime(tcol)
|
||||
s = pd.Series(df[hmm_field].values, index=ts, name=sym)
|
||||
s = s.loc[start:end].dropna()
|
||||
if len(s) > 0:
|
||||
all_posteriors[sym] = s
|
||||
break
|
||||
if not all_posteriors:
|
||||
# no HMM features found — default open
|
||||
idx = pd.date_range(start, end, freq="B")
|
||||
return pd.Series(True, index=idx)
|
||||
post_df = pd.DataFrame(all_posteriors)
|
||||
cs_mean = post_df.mean(axis=1)
|
||||
gate = cs_mean >= threshold
|
||||
return gate
|
||||
Reference in New Issue
Block a user