add signal-quality gate strategy class (exp 57)

This commit is contained in:
zhaoli
2026-08-20 23:13:03 +00:00
parent c367e25889
commit 1603a063ba
+94
View File
@@ -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