Files
tac-exp-dev/code/tac-qlib/tac_qlib/contrib/strategy/momentum_gate.py
T

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)