"""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)