216 lines
7.7 KiB
Python
216 lines
7.7 KiB
Python
"""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
|