start experiment 60 (exp/60-scheduled-algo-retrain-on-2026-08-21-tac)

This commit is contained in:
zhaoli
2026-08-24 12:01:27 +00:00
parent 7460d5fe01
commit 891bd0743d
14 changed files with 1500 additions and 21 deletions
@@ -0,0 +1,361 @@
"""Long-short Top-K strategy for cross-sectional signals.
Each day the strategy ranks the cross-section by prediction score and rebalances
to an equal-weight two-sided book: the ``topk`` highest-ranked names go long and
the ``k_short`` lowest-ranked names go short. Net-new shorts are opened by
selling beyond current holdings, which requires a short-aware exchange such as
``tac_qlib.contrib.backtest.tradeac_exchange.TradeACExchange`` with
``allow_short=True`` (borrow limits, margin requirements and borrow fees are
enforced there, not here).
Sizing deploys ``equity * risk_degree`` as gross notional split evenly across
all long and short legs, so the book is approximately market neutral.
``allow_short=False`` disables the short side entirely (long-only ``topk``).
Short eligibility can be restricted further, with static or dynamic gates:
``short_whitelist`` limits shorts to an explicit symbol set; ``short_vol_top_pct``
requires a candidate's trailing realized volatility to rank in the top fraction
of that day's cross-section; ``short_max_mom`` (falling-knife filter) only
allows shorting names whose own trailing momentum is at/below a threshold;
``short_regime_sma`` disables shorts entirely while the benchmark trades above
its moving average (risk-on). Borrow availability itself is enforced by the
exchange (``borrowable`` whitelist / per-symbol caps via ``TradeACExchange``).
Wired into qrun workflows like any ``BaseStrategy`` (see ``PortAnaRecord``
config). Mirrors the API usage of qlib's ``TopkDropoutStrategy``: ``Order`` /
``OrderDir`` from ``qlib.backtest.decision``, ``trade_calendar`` /
``trade_exchange`` / ``trade_position`` injected by the backtest executor.
"""
from __future__ import annotations
import copy
from typing import Dict, List, Optional
import pandas as pd
from qlib.backtest import Order
from qlib.backtest.decision import OrderDir, TradeDecisionWO
from qlib.contrib.strategy.signal_strategy import BaseSignalStrategy
from qlib.log import get_module_logger
__all__ = ["LongShortTopkStrategy"]
class LongShortTopkStrategy(BaseSignalStrategy):
"""Equal-weight long-short Top-K strategy over a cross-sectional signal.
Parameters
----------
topk : number of long legs (highest-ranked names).
k_short : number of short legs (lowest-ranked names).
hold_thresh : minimum holding days before a leg may be closed/reduced.
only_tradable : only select candidates tradable on the trade date.
rebalance_tol : skip rebalances smaller than this fraction of a leg's
target notional (turnover control).
allow_short : enable/disable the short side. With ``False`` the bottom-ranked
legs are dropped and the book is long-only ``topk``; pair with
``allow_short=False`` on the exchange for a fully borrow-free run.
Legacy alias ``enable_short`` is accepted.
short_whitelist : optional list of symbols eligible for shorting; candidates
outside the list are skipped (``None`` = all names eligible).
short_vol_window : trailing window (trading days) for realized-vol estimation.
short_vol_top_pct : if set, a short candidate's trailing realized volatility
must rank at or above this percentile of that day's cross-section
(e.g. ``0.5`` = only the more volatile half may be shorted). Candidates
without measurable vol are never shorted.
short_mom_window : trailing window (trading days) for the candidate momentum
used by the falling-knife gate.
short_max_mom : if set, a candidate's trailing ``short_mom_window``-day return
must be <= this value to be shortable (e.g. ``0.0`` = only short names
that are actually falling). Candidates without measurable momentum are
never shorted.
short_regime_symbol : benchmark symbol for the regime gate (default SPY).
short_regime_sma : if set, shorts are only allowed on days where the regime
symbol's last close (strictly before the execution bar) is BELOW its
``short_regime_sma``-day moving average — i.e. shorts are disabled in
risk-on regimes and enabled in drawdowns.
"""
def __init__(
self,
*,
signal=None,
topk: int = 4,
k_short: int = 2,
hold_thresh: int = 1,
only_tradable: bool = True,
rebalance_tol: float = 0.05,
allow_short: Optional[bool] = None,
enable_short: Optional[bool] = None,
short_whitelist: Optional[List[str]] = None,
short_vol_window: int = 20,
short_vol_top_pct: Optional[float] = None,
short_mom_window: int = 20,
short_max_mom: Optional[float] = None,
short_regime_symbol: str = "SPY",
short_regime_sma: Optional[int] = None,
risk_degree: float = 0.95,
trade_exchange=None,
level_infra=None,
common_infra=None,
**kwargs,
):
super().__init__(
signal=signal,
risk_degree=risk_degree,
trade_exchange=trade_exchange,
level_infra=level_infra,
common_infra=common_infra,
**kwargs,
)
if allow_short is None:
allow_short = True if enable_short is None else bool(enable_short)
self.allow_short = bool(allow_short)
self.topk = topk
self.k_short = k_short
self.hold_thresh = hold_thresh
self.only_tradable = only_tradable
self.rebalance_tol = rebalance_tol
self.short_whitelist = set(short_whitelist) if short_whitelist is not None else None
if not 0 < float(short_vol_window) <= 1000:
raise ValueError(f"short_vol_window must be in (0, 1000], got {short_vol_window}")
self.short_vol_window = int(short_vol_window)
if short_vol_top_pct is not None and not 0.0 < float(short_vol_top_pct) <= 1.0:
raise ValueError(f"short_vol_top_pct must be in (0, 1], got {short_vol_top_pct}")
self.short_vol_top_pct = None if short_vol_top_pct is None else float(short_vol_top_pct)
if not 0 < float(short_mom_window) <= 1000:
raise ValueError(f"short_mom_window must be in (0, 1000], got {short_mom_window}")
self.short_mom_window = int(short_mom_window)
self.short_max_mom = None if short_max_mom is None else float(short_max_mom)
self.short_regime_symbol = str(short_regime_symbol)
if short_regime_sma is not None and not 1 < int(short_regime_sma) <= 1000:
raise ValueError(f"short_regime_sma must be in (1, 1000], got {short_regime_sma}")
self.short_regime_sma = None if short_regime_sma is None else int(short_regime_sma)
# per-day caches (keyed by trade date)
self._vol_cache_key: Optional[str] = None
self._vol_cache_val: Dict[str, Dict[str, float]] = {}
self._regime_cache: Dict[str, bool] = {}
# ------------------------------------------------------------------ utils
def _is_tradable(self, code, start, end) -> bool:
if not self.only_tradable:
return True
try:
return self.trade_exchange.is_stock_tradable(stock_id=code, start_time=start, end_time=end)
except TypeError:
return True
def _mark_price(self, code, start, end) -> Optional[float]:
try:
px = self.trade_exchange.get_deal_price(
stock_id=code, start_time=start, end_time=end, direction=OrderDir.BUY
)
except (KeyError, ValueError):
return None
if px is None or px != px or px <= 0:
return None
return float(px)
def _day_stats(self, codes: List[str], trade_start) -> Dict[str, Dict[str, float]]:
"""Per-day cross-sectional stats used by the dynamic short gates.
For each code, returns ``{"vol_rank": r}`` (percentile of trailing
realized vol over ``short_vol_window`` bars across that day's
cross-section) when the vol gate is on, and ``{"mom": m}`` (trailing
``short_mom_window``-bar return) when the falling-knife gate is on.
All series end on the last bar strictly BEFORE the execution bar (no
lookahead). Codes without measurable data are simply absent — such
candidates are never shorted (fail-closed).
"""
if self.short_vol_top_pct is None and self.short_max_mom is None:
return {}
key = str(pd.Timestamp(trade_start))
if self._vol_cache_key == key:
return self._vol_cache_val
out: Dict[str, Dict[str, float]] = {}
try:
from qlib.data import D
end = pd.Timestamp(trade_start)
buf = max(self.short_vol_window, self.short_mom_window) * 3 + 30
df = D.features(
sorted(codes),
["$close"],
start_time=end - pd.Timedelta(days=buf),
end_time=end - pd.Timedelta(days=1),
)
close = df["$close"].unstack(level="instrument") if isinstance(df.index, pd.MultiIndex) else df["$close"]
if self.short_vol_top_pct is not None:
vol = close.pct_change().rolling(self.short_vol_window).std().iloc[-1]
for code, rank in vol.rank(pct=True).dropna().items():
out.setdefault(str(code), {})["vol_rank"] = float(rank)
if self.short_max_mom is not None:
w = min(self.short_mom_window, len(close) - 1)
mom = close.iloc[-1] / close.iloc[-(w + 1)] - 1
for code, m in mom.items():
if m == m:
out.setdefault(str(code), {})["mom"] = float(m)
except Exception as e: # noqa: BLE001 - degrade to fail-closed (no shorts)
get_module_logger(self.__class__.__name__).warning(
f"short gates unavailable ({type(e).__name__}: {e}); no shorts this step"
)
self._vol_cache_key, self._vol_cache_val = key, out
return out
def _regime_ok(self, trade_start) -> bool:
"""True when shorting is allowed by the benchmark-regime gate.
With ``short_regime_sma`` set, shorts are permitted only while the
regime symbol's last close strictly before the execution bar sits below
its moving average (risk-off). Data failure fails closed (no shorts).
"""
if self.short_regime_sma is None:
return True
key = str(pd.Timestamp(trade_start))
cached = self._regime_cache.get(key)
if cached is not None:
return cached
ok = False
try:
from qlib.data import D
end = pd.Timestamp(trade_start)
df = D.features(
[self.short_regime_symbol],
["$close"],
start_time=end - pd.Timedelta(days=int(self.short_regime_sma * 3 + 30)),
end_time=end - pd.Timedelta(days=1),
)
s = df["$close"]
if isinstance(s.index, pd.MultiIndex):
s = s.droplevel("instrument")
sma = s.rolling(self.short_regime_sma).mean().iloc[-1]
px = s.iloc[-1]
ok = bool(px < sma)
except Exception as e: # noqa: BLE001 - degrade to fail-closed (no shorts)
get_module_logger(self.__class__.__name__).warning(
f"regime gate unavailable ({type(e).__name__}: {e}); no shorts this step"
)
self._regime_cache[key] = ok
return ok
# ------------------------------------------------------------- decision
def generate_trade_decision(self, execute_result=None):
trade_step = self.trade_calendar.get_trade_step()
trade_start, trade_end = self.trade_calendar.get_step_time(trade_step)
pred_start, pred_end = self.trade_calendar.get_step_time(trade_step, shift=1)
pred_score = self.signal.get_signal(start_time=pred_start, end_time=pred_end)
if isinstance(pred_score, pd.DataFrame):
pred_score = pred_score.iloc[:, 0]
if pred_score is None or len(pred_score) == 0:
return TradeDecisionWO([], self)
pred_score = pred_score.dropna()
if pred_score.empty:
return TradeDecisionWO([], self)
time_per_step = self.trade_calendar.get_freq()
current_temp = copy.deepcopy(self.trade_position)
# ---- signed current holdings ---------------------------------------
cur_amount: Dict[str, float] = {}
for code in current_temp.get_stock_list():
amt = float(current_temp.get_stock_amount(code))
if abs(amt) > 1e-6:
cur_amount[code] = amt
# ---- targets: top-k long, bottom-k_short short ----------------------
ranked = list(pred_score.sort_values(ascending=False).index)
longs: List[str] = []
for code in ranked:
if len(longs) >= self.topk:
break
if self._is_tradable(code, trade_start, trade_end):
longs.append(code)
shorts: List[str] = []
if self.allow_short and self._regime_ok(trade_start):
stats = self._day_stats(list(ranked), trade_start)
for code in reversed(ranked):
if len(shorts) >= self.k_short:
break
if code in longs:
continue
if not self._is_tradable(code, trade_start, trade_end):
continue
if self.short_whitelist is not None and code not in self.short_whitelist:
continue
st = stats.get(code)
if self.short_vol_top_pct is not None:
rank = None if st is None else st.get("vol_rank")
if rank is None or rank < self.short_vol_top_pct:
continue
if self.short_max_mom is not None:
mom = None if st is None else st.get("mom")
if mom is None or mom > self.short_max_mom:
continue
shorts.append(code)
# ---- marks & equity --------------------------------------------------
marks: Dict[str, float] = {}
for code in set(cur_amount) | set(longs) | set(shorts):
px = self._mark_price(code, trade_start, trade_end)
if px is not None:
marks[code] = px
equity = current_temp.get_cash()
for code, amt in cur_amount.items():
if code in marks:
equity += amt * marks[code]
if equity <= 0:
return TradeDecisionWO([], self)
n_legs = len([c for c in longs if c in marks]) + len([c for c in shorts if c in marks])
if n_legs == 0:
return TradeDecisionWO([], self)
per_leg = equity * self.risk_degree / n_legs
target_signed: Dict[str, float] = {}
for code in longs:
if code in marks:
target_signed[code] = per_leg / marks[code]
for code in shorts:
if code in marks:
target_signed[code] = -(per_leg / marks[code])
# ---- order generation -------------------------------------------------
sell_orders: List[Order] = []
buy_orders: List[Order] = []
def submit(code: str, amount: float, direction: int) -> None:
factor = self.trade_exchange.get_factor(stock_id=code, start_time=trade_start, end_time=trade_end)
amount = self.trade_exchange.round_amount_by_trade_unit(amount, factor)
if amount <= 1e-6:
return
o = Order(stock_id=code, amount=amount, start_time=trade_start, end_time=trade_end, direction=direction)
if self.trade_exchange.check_order(o):
(buy_orders if direction == Order.BUY else sell_orders).append(o)
# close holdings that are no longer targeted (frees cash / unwinds shorts)
for code, amt in cur_amount.items():
if code in target_signed:
continue
if marks.get(code) is None:
continue
if current_temp.get_stock_count(code, bar=time_per_step) < self.hold_thresh:
continue
submit(code, abs(amt), Order.SELL if amt > 0 else Order.BUY)
# rebalance targeted legs toward their signed target quantity
for code, tgt in target_signed.items():
cur = cur_amount.get(code, 0.0)
delta = tgt - cur
if abs(delta * marks[code]) < max(self.rebalance_tol * per_leg, 1.0):
continue
if delta > 0:
submit(code, delta, Order.BUY)
else:
if cur > 0 and current_temp.get_stock_count(code, bar=time_per_step) < self.hold_thresh:
continue
submit(code, -delta, Order.SELL)
return TradeDecisionWO(sell_orders + buy_orders, self)