95 lines
3.6 KiB
Python
95 lines
3.6 KiB
Python
"""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
|