"""Signal-quality gate TopkDropout strategy. Subclass of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy`` that holds the book (issues NO orders) when the model's recent prediction accuracy is below a threshold. When the gate is open it behaves exactly like the reference TopkDropoutStrategy. Unlike the regime gate (which asks "is the market calm?"), the signal-quality gate asks "are my predictions accurate?" — and works across ALL years. The gate is provided as a precomputed ``pd.Series`` of booleans indexed by datetime (True = trade allowed). The companion ``compute_signal_quality_gate`` function builds this series from a pred.pkl and lake bars; call it once before backtesting and pass the result as the ``signal_quality_gate`` parameter. """ from __future__ import annotations import numpy as np import pandas as pd from qlib.backtest.decision import TradeDecisionWO from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy __all__ = ["SignalQualityGateStrategy", "compute_signal_quality_gate"] class SignalQualityGateStrategy(TopkDropoutStrategy): """TopkDropout with signal-quality gate overlay. Parameters ---------- topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable, forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``. signal_quality_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, *, signal_quality_gate=None, signal_quality_gate_path=None, **kwargs): super().__init__(**kwargs) if signal_quality_gate is not None: self._sq_gate = signal_quality_gate elif signal_quality_gate_path is not None: import pickle with open(signal_quality_gate_path, "rb") as f: self._sq_gate = pickle.load(f) else: self._sq_gate = None def _gate_open(self, trade_start_time) -> bool: if self._sq_gate is None: return True ts = pd.Timestamp(trade_start_time) known = self._sq_gate[self._sq_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_signal_quality_gate( pred_path: str, *, lake_root: str = "", market: str = "US", topk: int = 10, lookback: int = 5, threshold: float = 0.5, start: str = "2015-01-03", end: str = "2026-08-19", ) -> pd.Series: """Build a per-date signal-quality gate series from a pred.pkl and lake bars. For each day, checks whether the model's topk picks from the previous day had positive returns. Computes a rolling hit rate over ``lookback`` days and opens the gate when hit rate >= ``threshold``. Parameters ---------- pred_path : str — path to pred.pkl (from rd_train / rd_predict). lake_root, market : str — lake location (for loading close prices). topk : int — number of top picks to track for hit rate. lookback : int — rolling window for hit rate computation. threshold : float — hit rate threshold to keep trading. start, end : str — date window for loading prices. Returns ------- pd.Series — bool, indexed by datetime. True = trade allowed. """ from pathlib import Path import pickle # Load pred.pkl with open(pred_path, "rb") as f: pred = pickle.load(f) # Handle MultiIndex DataFrame (datetime, instrument) -> unstack to wide if isinstance(pred, pd.DataFrame) and isinstance(pred.index, pd.MultiIndex): pred = pred.iloc[:, 0] # take score column as Series pred.index = pd.MultiIndex.from_arrays([ pd.to_datetime(pred.index.get_level_values(0)).normalize(), pred.index.get_level_values(1) ]) # Unstack to wide: dates x instruments pred = pred.unstack(level=1) elif isinstance(pred, pd.DataFrame): pred = pred.iloc[:, 0] if pred.shape[1] >= 1 else pred.squeeze() pred.index = pd.to_datetime(pred.index).normalize() # Load close prices from lake close_df = _load_close_prices(lake_root, market, start, end) if close_df.empty: return pd.Series(dtype=bool) ret_df = close_df.pct_change() ret_df.index = pd.to_datetime(ret_df.index).normalize() # Get sorted unique prediction dates pred_dates = sorted(pred.index.unique()) if len(pred_dates) < 2: return pd.Series(True, index=pd.DatetimeIndex(pred_dates)) # Compute hit rates hit_rates = {} for i in range(1, len(pred_dates)): day = pred_dates[i] # Get yesterday's topk prev_day = pred_dates[i - 1] try: prev_scores = pred.loc[prev_day] except KeyError: continue if isinstance(prev_scores, pd.DataFrame): prev_scores = prev_scores.iloc[:, 0] prev_scores = prev_scores.dropna().sort_values(ascending=False) topk_syms = list(prev_scores.index[:topk]) # Get today's returns if day not in ret_df.index: continue today_ret = ret_df.loc[day] topk_rets = today_ret.reindex(topk_syms).dropna() if len(topk_rets) == 0: continue hit_rates[day] = (topk_rets > 0).sum() / len(topk_rets) if not hit_rates: return pd.Series(dtype=bool) hr_series = pd.Series(hit_rates).sort_index() # Rolling hit rate rolling_hr = hr_series.rolling(lookback, min_periods=1).mean() # Gate is open when rolling hit rate >= threshold gate = rolling_hr >= threshold gate.iloc[:lookback] = True # warmup: allow trading return gate def _load_close_prices(lake_root, market, start, end): """Load daily close prices for all symbols into a wide DataFrame.""" from pathlib import Path lake = Path(lake_root) symbols_parquet = lake / "symbols.parquet" if not symbols_parquet.exists(): return pd.DataFrame() df = pd.read_parquet(symbols_parquet) col = "symbol" if "symbol" in df.columns else df.columns[0] symbols = sorted(df[col].astype(str).str.upper().tolist()) closes = {} for sym in symbols: p = lake / "market=US" / "timeframe=1d" / f"symbol={sym}.parquet" if not p.exists(): continue try: bar = pd.read_parquet(p) except Exception: continue if not len(bar): continue tcol = "t" if "t" in bar.columns else "date" ts = pd.to_datetime(bar[tcol]) bar = bar.assign(_t=ts).set_index("_t").sort_index() bar = bar.loc[start:end] if len(bar) < 10: continue closes[sym] = bar["c"] if "c" in bar.columns else bar["close"] return pd.DataFrame(closes)