start experiment 59 (exp/59-signal-quality-gate-walk-forward-test-wi)
This commit is contained in:
@@ -28,26 +28,118 @@ __all__ = ["SignalQualityGateStrategy", "compute_signal_quality_gate"]
|
||||
class SignalQualityGateStrategy(TopkDropoutStrategy):
|
||||
"""TopkDropout with signal-quality gate overlay.
|
||||
|
||||
When ``lake_root`` is provided the gate is computed on-the-fly from the
|
||||
signal (``<PRED>``) and close prices — no precomputed gate file needed.
|
||||
This ensures the gate matches the model that is actually generating the
|
||||
predictions (critical when the model is retrained each year).
|
||||
|
||||
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).
|
||||
signal_quality_gate : pd.Series — precomputed per-date gate (bool).
|
||||
signal_quality_gate_path : str — path to pickled gate Series.
|
||||
lake_root : str — lake root for on-the-fly gate computation (preferred).
|
||||
gate_topk, gate_lookback, gate_threshold : int/float — gate params.
|
||||
gate_start, gate_end : str — date window for loading close prices.
|
||||
"""
|
||||
|
||||
def __init__(self, *, signal_quality_gate=None, signal_quality_gate_path=None, **kwargs):
|
||||
def __init__(self, *, signal_quality_gate=None, signal_quality_gate_path=None,
|
||||
lake_root=None, gate_topk=10, gate_lookback=5, gate_threshold=0.5,
|
||||
gate_start="2015-01-03", gate_end="2026-08-19", **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._sq_gate_computed = False
|
||||
if signal_quality_gate is not None:
|
||||
self._sq_gate = signal_quality_gate
|
||||
self._sq_gate_computed = True
|
||||
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)
|
||||
self._sq_gate_computed = True
|
||||
elif lake_root is not None:
|
||||
self._sq_gate = None
|
||||
self._lake_root = lake_root
|
||||
self._gate_topk = gate_topk
|
||||
self._gate_lookback = gate_lookback
|
||||
self._gate_threshold = gate_threshold
|
||||
self._gate_start = gate_start
|
||||
self._gate_end = gate_end
|
||||
else:
|
||||
self._sq_gate = None
|
||||
|
||||
def _compute_gate_on_fly(self):
|
||||
"""Compute gate from the signal (pred.pkl) and lake close prices."""
|
||||
import pickle
|
||||
from pathlib import Path
|
||||
|
||||
signal_path = self._signal
|
||||
if not Path(signal_path).exists():
|
||||
return
|
||||
|
||||
with open(signal_path, "rb") as f:
|
||||
pred = pickle.load(f)
|
||||
|
||||
# Handle MultiIndex DataFrame -> unstack to wide
|
||||
if isinstance(pred, pd.DataFrame) and isinstance(pred.index, pd.MultiIndex):
|
||||
pred = pred.iloc[:, 0]
|
||||
pred.index = pd.MultiIndex.from_arrays([
|
||||
pd.to_datetime(pred.index.get_level_values(0)).normalize(),
|
||||
pred.index.get_level_values(1)
|
||||
])
|
||||
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(self._lake_root, "US", self._gate_start, self._gate_end)
|
||||
if close_df.empty:
|
||||
return
|
||||
|
||||
ret_df = close_df.pct_change()
|
||||
ret_df.index = pd.to_datetime(ret_df.index).normalize()
|
||||
|
||||
pred_dates = sorted(pred.index.unique())
|
||||
if len(pred_dates) < 2:
|
||||
self._sq_gate = pd.Series(True, index=pd.DatetimeIndex(pred_dates))
|
||||
self._sq_gate_computed = True
|
||||
return
|
||||
|
||||
hit_rates = {}
|
||||
for i in range(1, len(pred_dates)):
|
||||
day = pred_dates[i]
|
||||
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[:self._gate_topk])
|
||||
|
||||
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:
|
||||
self._sq_gate_computed = True
|
||||
return
|
||||
|
||||
hr_series = pd.Series(hit_rates).sort_index()
|
||||
rolling_hr = hr_series.rolling(self._gate_lookback, min_periods=1).mean()
|
||||
gate = rolling_hr >= self._gate_threshold
|
||||
gate.iloc[:self._gate_lookback] = True
|
||||
|
||||
self._sq_gate = gate
|
||||
self._sq_gate_computed = True
|
||||
|
||||
def _gate_open(self, trade_start_time) -> bool:
|
||||
if self._sq_gate is None:
|
||||
return True
|
||||
@@ -58,6 +150,8 @@ class SignalQualityGateStrategy(TopkDropoutStrategy):
|
||||
return True # default open if no history yet
|
||||
|
||||
def generate_trade_decision(self, execute_result=None):
|
||||
if not self._sq_gate_computed:
|
||||
self._compute_gate_on_fly()
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user