92 lines
3.6 KiB
Python
92 lines
3.6 KiB
Python
"""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: "<PRED>"
|
|
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)
|