"""Signal-quality gate strategy: gate trades based on hit-rate of topk predictions. Unlike the regime gate (which asks 'is the market calm?'), the signal-quality gate asks 'are my predictions accurate?' and works across ALL years. """ from __future__ import annotations import numpy as np import pandas as pd from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy class SignalQualityGateStrategy(TopkDropoutStrategy): """TopkDropout with signal-quality gate overlay. The gate computes the rolling hit rate of the model's topk picks: - For each day, check if yesterday's topk had positive returns - Compute rolling hit rate over lookback days - If hit rate >= threshold, trade; otherwise, go to cash Parameters ---------- signal_quality_gate_lookback : int Rolling window for hit rate computation (default: 5) signal_quality_gate_threshold : float Hit rate threshold to keep trading (default: 0.5) signal_quality_gate_topk : int Number of top picks to track for hit rate (default: 10) """ def __init__(self, *args, **kwargs): self.sg_lookback = kwargs.pop("signal_quality_gate_lookback", 5) self.sg_threshold = kwargs.pop("signal_quality_gate_threshold", 0.5) self.sg_topk = kwargs.pop("signal_quality_gate_topk", 10) super().__init__(*args, **kwargs) self._hit_rates = {} self._trade_dates = [] def get_kick_out_day_list(self, phase, **kwargs): """Override to compute hit rates and determine gate-open days.""" # Get the standard trade dates from parent trade_dates = super().get_kick_out_day_list(phase, **kwargs) if trade_dates is None: return trade_dates # We'll compute hit rates in the backtest loop # For now, return all dates (gate applied in get_gated_sp) self._trade_dates = trade_dates return trade_dates def compute_hit_rate(self, date_idx: int, pred_df: pd.DataFrame, ret_df: pd.DataFrame) -> float: """Compute rolling hit rate up to date_idx.""" if date_idx < 1: return 1.0 # default open when no history hit_count = 0 total_count = 0 for i in range(max(1, date_idx - self.sg_lookback), date_idx): if i < 1: continue # Get yesterday's topk prev_date = self._trade_dates[i - 1] if i - 1 < len(self._trade_dates) else None curr_date = self._trade_dates[i] if i < len(self._trade_dates) else None if prev_date is None or curr_date is None: continue try: prev_scores = pred_df.loc[prev_date] 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[:self.sg_topk]) # Get today's returns if curr_date not in ret_df.index: continue today_ret = ret_df.loc[curr_date] topk_rets = today_ret.reindex(topk_syms).dropna() if len(topk_rets) > 0: hit_count += (topk_rets > 0).sum() total_count += len(topk_rets) if total_count == 0: return 1.0 # default open return hit_count / total_count def is_gate_open(self, date_idx: int, pred_df: pd.DataFrame, ret_df: pd.DataFrame) -> bool: """Check if the signal-quality gate is open for this date.""" hr = self.compute_hit_rate(date_idx, pred_df, ret_df) return hr >= self.sg_threshold