"""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