From 1603a063baad186d8be6d01ebe31370ef938d5e6 Mon Sep 17 00:00:00 2001 From: zhaoli Date: Thu, 20 Aug 2026 23:13:03 +0000 Subject: [PATCH] add signal-quality gate strategy class (exp 57) --- book/strategies/signal_quality_gate.py | 94 ++++++++++++++++++++++++++ 1 file changed, 94 insertions(+) create mode 100644 book/strategies/signal_quality_gate.py diff --git a/book/strategies/signal_quality_gate.py b/book/strategies/signal_quality_gate.py new file mode 100644 index 0000000..bf78bb3 --- /dev/null +++ b/book/strategies/signal_quality_gate.py @@ -0,0 +1,94 @@ +"""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