"""TopkDropout with a 1-day momentum entry-confirmation gate. Gates NEW entries on short-term momentum: a name that is not currently held may only be bought when its trailing 1-day return is above ``min_momentum`` (Lag-1 autocorr ~ +0.45 in the time-series study => short-term momentum continuation). Held names are never force-sold by this gate — exits stay the pure TopkDropout rule. Implementation: override ``generate_trade_decision`` and zero out the signal score of any non-held name that fails the momentum check BEFORE calling the base TopkDropout decision, so it can never be selected as a buy candidate. This is a clean pre-filter: the rest of the strategy (top-k, n_drop, sizing, costs) is untouched. Wired into a workflow yaml like: strategy: class: MomentumGateTopk module_path: tac_qlib.contrib.strategy.momentum_gate kwargs: signal: "" topk: 10 n_drop: 2 only_tradable: true risk_degree: 0.95 min_momentum: 0.0 """ from __future__ import annotations import copy import pandas as pd from qlib.backtest.decision import TradeDecisionWO from qlib.backtest.position import Position from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy __all__ = ["MomentumGateTopk"] class MomentumGateTopk(TopkDropoutStrategy): """TopkDropoutStrategy gated on 1-day momentum for new entries.""" def __init__(self, *, min_momentum: float = 0.0, **kwargs): super().__init__(**kwargs) self.min_momentum = float(min_momentum) def _momentum_ok(self, code, trade_start, trade_end) -> bool: """True when the trailing 1-day return is above the momentum floor.""" try: cur = self.trade_exchange.get_deal_price( stock_id=code, start_time=trade_start, end_time=trade_end, direction=1 ) except Exception: return False if cur is None or cur != cur or cur <= 0: return False prev_start = trade_start - pd.Timedelta(days=5) prev_end = trade_start - pd.Timedelta(seconds=1) prev = self.trade_exchange.get_deal_price( stock_id=code, start_time=prev_start, end_time=prev_end, direction=0 ) if prev is None or prev != prev or prev <= 0: return False return (cur / prev - 1.0) >= self.min_momentum def generate_trade_decision(self, execute_result=None): trade_step = self.trade_calendar.get_trade_step() trade_start_time, trade_end_time = self.trade_calendar.get_step_time(trade_step) pred_start_time, pred_end_time = self.trade_calendar.get_step_time(trade_step, shift=1) pred_score = self.signal.get_signal(start_time=pred_start_time, end_time=pred_end_time) if pred_score is None: return TradeDecisionWO([], self) if isinstance(pred_score, pd.DataFrame): pred_score = pred_score.iloc[:, 0] current_temp = copy.deepcopy(self.trade_position) assert isinstance(current_temp, Position) held = set(current_temp.get_stock_list()) held = {c for c in held if abs(current_temp.get_stock_amount(c)) > 1e-6} # pre-filter: zero the score of non-held names that fail momentum pred_score = pred_score.copy() for code in pred_score.index: if code in held: continue # never gate exits / re-balancing of held names if not self._momentum_ok(code, trade_start_time, trade_end_time): pred_score[code] = -1e9 # cannot enter today return super().generate_trade_decision(execute_result)