start experiment 58 (exp/58-signal-quality-gate-5-year-walk-forward)
This commit is contained in:
@@ -0,0 +1,207 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user