From f215154cbb25391672cb6dcdf4e785be1561efa3 Mon Sep 17 00:00:00 2001 From: zhaoli Date: Thu, 1 Oct 2026 12:21:50 +0000 Subject: [PATCH] start experiment 86 (exp/86-scheduled-rolling-4y-retrain-for-round-3) --- code/MANIFEST.txt | 34 ++ code/tac-qlib/tac_qlib/contrib/__init__.py | 11 + .../__pycache__/__init__.cpython-312.pyc | Bin 0 -> 374 bytes .../tac_qlib/contrib/backtest/__init__.py | 1 + .../contrib/backtest/tradeac_exchange.py | 432 ++++++++++++++++++ .../tac_qlib/contrib/data/__init__.py | 3 + .../data/__pycache__/__init__.cpython-312.pyc | Bin 0 -> 216 bytes .../data/__pycache__/handler.cpython-312.pyc | Bin 0 -> 21533 bytes .../tac-qlib/tac_qlib/contrib/data/handler.py | 428 +++++++++++++++++ .../tac_qlib/contrib/model/__init__.py | 4 + .../__pycache__/__init__.cpython-312.pyc | Bin 0 -> 319 bytes .../__pycache__/rank_ensemble.cpython-312.pyc | Bin 0 -> 9461 bytes .../__pycache__/rank_gbdt.cpython-312.pyc | Bin 0 -> 12288 bytes .../tac_qlib/contrib/model/rank_ensemble.py | 189 ++++++++ .../tac_qlib/contrib/model/rank_gbdt.py | 238 ++++++++++ .../tac_qlib/contrib/strategy/__init__.py | 4 + .../__pycache__/__init__.cpython-312.pyc | Bin 0 -> 301 bytes .../__pycache__/long_short.cpython-312.pyc | Bin 0 -> 19672 bytes .../__pycache__/optimal_stop.cpython-312.pyc | Bin 0 -> 10428 bytes .../contrib/strategy/kelly_dropout.py | 201 ++++++++ .../tac_qlib/contrib/strategy/long_short.py | 361 +++++++++++++++ .../tac_qlib/contrib/strategy/optimal_stop.py | 217 +++++++++ .../tac_qlib/contrib/strategy/regime_gate.py | 231 ++++++++++ .../contrib/strategy/weekly_rebalance.py | 223 +++++++++ code/tac-qlib/tac_qlib/data/__init__.py | 25 + .../data/__pycache__/__init__.cpython-312.pyc | Bin 0 -> 522 bytes .../data/__pycache__/config.cpython-312.pyc | Bin 0 -> 11507 bytes .../__pycache__/providers.cpython-312.pyc | Bin 0 -> 13340 bytes code/tac-qlib/tac_qlib/data/config.py | 207 +++++++++ code/tac-qlib/tac_qlib/data/providers.py | 230 ++++++++++ 30 files changed, 3039 insertions(+) create mode 100644 code/MANIFEST.txt create mode 100644 code/tac-qlib/tac_qlib/contrib/__init__.py create mode 100644 code/tac-qlib/tac_qlib/contrib/__pycache__/__init__.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/backtest/__init__.py create mode 100644 code/tac-qlib/tac_qlib/contrib/backtest/tradeac_exchange.py create mode 100644 code/tac-qlib/tac_qlib/contrib/data/__init__.py create mode 100644 code/tac-qlib/tac_qlib/contrib/data/__pycache__/__init__.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/data/__pycache__/handler.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/data/handler.py create mode 100644 code/tac-qlib/tac_qlib/contrib/model/__init__.py create mode 100644 code/tac-qlib/tac_qlib/contrib/model/__pycache__/__init__.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/model/__pycache__/rank_ensemble.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/model/__pycache__/rank_gbdt.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/model/rank_ensemble.py create mode 100644 code/tac-qlib/tac_qlib/contrib/model/rank_gbdt.py create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/__init__.py create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/__pycache__/__init__.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/__pycache__/long_short.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/__pycache__/optimal_stop.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/kelly_dropout.py create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/long_short.py create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/optimal_stop.py create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/regime_gate.py create mode 100644 code/tac-qlib/tac_qlib/contrib/strategy/weekly_rebalance.py create mode 100644 code/tac-qlib/tac_qlib/data/__init__.py create mode 100644 code/tac-qlib/tac_qlib/data/__pycache__/__init__.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/data/__pycache__/config.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/data/__pycache__/providers.cpython-312.pyc create mode 100644 code/tac-qlib/tac_qlib/data/config.py create mode 100644 code/tac-qlib/tac_qlib/data/providers.py diff --git a/code/MANIFEST.txt b/code/MANIFEST.txt new file mode 100644 index 0000000..1a0422f --- /dev/null +++ b/code/MANIFEST.txt @@ -0,0 +1,34 @@ +# TradeAC custom-qlib-code snapshot (auto-generated) +# parent repo HEAD : fd5382caa4c4e96417a16d5a22006d9a7d625d00 +# tac-qlib/tac_qlib/contrib +# tac-qlib/tac_qlib/data +# per-file hashes (git hash-object): + 1b6298c4a5652f2e863cbdc385a1014a570fcd59 tac-qlib/tac_qlib/contrib/__init__.py + d1a8ec0c9e6c38839e966f687b08ec412f87ec20 tac-qlib/tac_qlib/contrib/__pycache__/__init__.cpython-312.pyc + 2224424d0ff193be4f55d1b791f8fce89439c5d2 tac-qlib/tac_qlib/contrib/backtest/__init__.py + 0bf40dee440ddbded357d7bbb4efc67c62c4b084 tac-qlib/tac_qlib/contrib/backtest/tradeac_exchange.py + c76a9f17f680e74eea766eff27f7624359749ed6 tac-qlib/tac_qlib/contrib/data/__init__.py + c1543fabfc22e07c7e9942015c4462f84acc1e91 tac-qlib/tac_qlib/contrib/data/__pycache__/__init__.cpython-312.pyc + 89f5cfc15b946b6ea1fb7724f021024967cdb6c5 tac-qlib/tac_qlib/contrib/data/__pycache__/handler.cpython-312.pyc + 3bba0f1696e4ab4b3deebec3f31f269b2e713899 tac-qlib/tac_qlib/contrib/data/handler.py + b151d139a0dcde87d74b21e7c4b729176ba5c39b tac-qlib/tac_qlib/contrib/model/__init__.py + f4e9bf0a78d22e256cade0d794a6958aaa7c6721 tac-qlib/tac_qlib/contrib/model/__pycache__/__init__.cpython-312.pyc + c35c166a8449c14fa3cd4972c363d3d9d58620ef tac-qlib/tac_qlib/contrib/model/__pycache__/rank_ensemble.cpython-312.pyc + 98be60fe19fedb3d9cb21c3f91ef0e176bd93287 tac-qlib/tac_qlib/contrib/model/__pycache__/rank_gbdt.cpython-312.pyc + d3f051f3a8650c42fedc7b367b966f7c74fb5789 tac-qlib/tac_qlib/contrib/model/rank_ensemble.py + d03e6611338918d4aac5eea4adf26f85a3763652 tac-qlib/tac_qlib/contrib/model/rank_gbdt.py + 184f80da8edf944bad3c8fb4d4d3d189bf4f082b tac-qlib/tac_qlib/contrib/strategy/__init__.py + 69684c9c9624cb16c63f3a18a864c49a7db37fa9 tac-qlib/tac_qlib/contrib/strategy/__pycache__/__init__.cpython-312.pyc + 4d3a5433e57e65e5dc96c6d2999f9d473b696dc4 tac-qlib/tac_qlib/contrib/strategy/__pycache__/long_short.cpython-312.pyc + 53e72f512b47207a5292a918efef7fee162deef3 tac-qlib/tac_qlib/contrib/strategy/__pycache__/optimal_stop.cpython-312.pyc + 896ef74ae47bcd1ed388e1e5d9c8d70c28097fe9 tac-qlib/tac_qlib/contrib/strategy/kelly_dropout.py + 9090fc6dfbd339f2f4df4b0c9b87f400ecb5c9d5 tac-qlib/tac_qlib/contrib/strategy/long_short.py + 79aaad9e39fcc740a773f4f63c512ce1086cfde0 tac-qlib/tac_qlib/contrib/strategy/optimal_stop.py + 5b9acfb4340111b204249add7760bd53c6ae03f1 tac-qlib/tac_qlib/contrib/strategy/regime_gate.py + 839abb89ad40cd516eabcfd91fe1be626b9f091f tac-qlib/tac_qlib/contrib/strategy/weekly_rebalance.py + 92e6e90eb0cd0a25142034560f27adb6b705b1a8 tac-qlib/tac_qlib/data/__init__.py + 0f1eb6f41d44bd9064b1f16f529dac6bb9d7c3fd tac-qlib/tac_qlib/data/__pycache__/__init__.cpython-312.pyc + cf99dd8a481f6e2d1f323f8d8d754730eadb357a tac-qlib/tac_qlib/data/__pycache__/config.cpython-312.pyc + ddac5bab6578399c0301390d71bae13511a56cb1 tac-qlib/tac_qlib/data/__pycache__/providers.cpython-312.pyc + 1953fb2a6371525db7f7b0e1c9dfbf3492d82110 tac-qlib/tac_qlib/data/config.py + 8d0644f6f0d1efb94798ed444cc73e63b643459b tac-qlib/tac_qlib/data/providers.py diff --git a/code/tac-qlib/tac_qlib/contrib/__init__.py b/code/tac-qlib/tac_qlib/contrib/__init__.py new file mode 100644 index 0000000..1b6298c --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/__init__.py @@ -0,0 +1,11 @@ +from . import data # noqa: F401 (registers tac_qlib.contrib.data) +from . import model, strategy # noqa: F401 +from .data import TACHandler # noqa: F401 +from .model import RankICLGBModel # noqa: F401 +from .strategy import OptimalStopControl # noqa: F401 + +__all__ = [ + "TACHandler", + "RankICLGBModel", + "OptimalStopControl", +] diff --git a/code/tac-qlib/tac_qlib/contrib/__pycache__/__init__.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..d1a8ec0c9e6c38839e966f687b08ec412f87ec20 GIT binary patch literal 374 zcmY+9y-ve06os#U@}ojKbOi>s1_}BC2njj>QUQrFd9h3rRRKE*t{FP=6y4bvcoNn&~>8JF8qa zTVT|=Ivky-BGs8i*Sl23?dfQIe01g~vD3e(TyB(}xUw3Rg|nqjm<{n>8+nOQ&Xc$X z%e>`Y0x$nZ>PSkZwUke=!W6!DhN`NDPEB|3bbjqYwlMWOupwn$*lJ+2f$atcU!1SehsdsD{sXhPT0H;& literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/contrib/backtest/__init__.py b/code/tac-qlib/tac_qlib/contrib/backtest/__init__.py new file mode 100644 index 0000000..2224424 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/backtest/__init__.py @@ -0,0 +1 @@ +from .tradeac_exchange import TradeACExchange diff --git a/code/tac-qlib/tac_qlib/contrib/backtest/tradeac_exchange.py b/code/tac-qlib/tac_qlib/contrib/backtest/tradeac_exchange.py new file mode 100644 index 0000000..0bf40de --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/backtest/tradeac_exchange.py @@ -0,0 +1,432 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +""" +TradeACExchange + +A short/borrow enabled Exchange implementation built on top of qlib.backtest.exchange.Exchange. +This exchange adds simple, configurable margin logic (initial/maintenance), borrowing support for +shorts, a borrow fee, and a lightweight SMA (Special Memorandum Account) concept to emulate +behaviors similar to brokers such as IBKR and Alpaca for backtesting purposes. + +Notes / limitations +- This implementation is intentionally lightweight and conservative: it implements the + key behaviors needed for strategy/backtest experiments (allowing short selling, computing + margin requirements, performing margin-call checks, and tracking SMA-like excess equity). +- It makes some simplifying assumptions compared to real brokers (no per-product house margins, + simplified SMA bookkeeping, borrow availability modeled only by a per-symbol boolean/limit). +- The Position class in qlib.backtest.position was not changed. To support shorts we update the + position.position dict directly when necessary. This keeps integration simple but bypasses some + internal Position helpers. Use with care. + +API additions +- allow_short: enable short selling (bool) +- initial_margin_long/short: fraction required to open a position +- maintenance_margin_long/short: fraction required to keep a position +- borrow_fee_rate: periodic borrow fee applied on short value (applied at trade time as additional cost) +- borrowable: dict mapping stock_id -> bool or float (max borrowable shares). Symbols missing from + the dict follow `borrow_default` (default True = unlimited; set False for a strict whitelist) +- get_sma(position): returns SMA-like excess equity available as "buying power credit" +- check_margin_call(position): returns True if position is below maintenance requirement + +""" +from __future__ import annotations + +from typing import Any, Dict, Optional, Tuple + +import numpy as np + +from qlib.backtest.decision import Order +from qlib.backtest.exchange import Exchange +from qlib.backtest.position import BasePosition + + +class TradeACExchange(Exchange): + """An exchange that supports short selling / borrowing and basic margin rules. + + The implementation aims to be compatible with the Exchange API used by Account and + Position classes in qlib.backtest. It overrides only the minimum methods required to + enable short/borrow behavior and margin calculations. + """ + + def __init__( + self, + *args: Any, + allow_short: bool = True, + initial_margin_long: float = 0.5, + initial_margin_short: float = 0.5, + maintenance_margin_long: float = 0.25, + maintenance_margin_short: float = 0.3, + borrow_fee_rate: float = 0.0, + borrowable: Optional[Dict[str, float]] = None, + borrow_default: bool = True, + sma_enabled: bool = True, + **kwargs: Any, + ) -> None: + """Create TradeACExchange. + + Parameters mirror Exchange with additional tradeac-specific options. + """ + super().__init__(*args, **kwargs) + self.allow_short = allow_short + self.initial_margin_long = initial_margin_long + self.initial_margin_short = initial_margin_short + self.maintenance_margin_long = maintenance_margin_long + self.maintenance_margin_short = maintenance_margin_short + self.borrow_fee_rate = borrow_fee_rate + # borrowable can be a dict with per-symbol max borrowable amount, or None (unlimited) + self.borrowable = borrowable or {} + # borrow_default: policy for symbols absent from `borrowable`. + # True -> unlisted symbols are unlimited-borrowable (legacy behavior) + # False -> unlisted symbols are NOT borrowable; only listed ones can be shorted + self.borrow_default = bool(borrow_default) + # sma_enabled: whether to expose lightweight SMA calculation + self.sma_enabled = sma_enabled + + # --------------------------- Helper calculations --------------------------- + def _initial_margin_requirement(self, position: BasePosition) -> float: + """Compute the initial margin requirement (money) for the given position. + + We treat longs and shorts separately and sum their required initial margins. + """ + im_req = 0.0 + for sid in position.get_stock_list(): + amt = position.get_stock_amount(sid) + price = position.get_stock_price(sid) + val = amt * price + if val > 0: + im_req += abs(val) * self.initial_margin_long + elif val < 0: + im_req += abs(val) * self.initial_margin_short + return im_req + + def _maintenance_margin_requirement(self, position: BasePosition) -> float: + """Compute the maintenance margin requirement (money) for the given position.""" + mm_req = 0.0 + for sid in position.get_stock_list(): + amt = position.get_stock_amount(sid) + price = position.get_stock_price(sid) + val = amt * price + if val > 0: + mm_req += abs(val) * self.maintenance_margin_long + elif val < 0: + mm_req += abs(val) * self.maintenance_margin_short + return mm_req + + def get_equity(self, position: BasePosition) -> float: + """Return account equity (position value + cash).""" + return position.calculate_value() + + def get_sma(self, position: BasePosition) -> float: + """Return a simplified SMA: excess equity above initial margin requirement. + + Note: This is a synthetic/Simplified SMA used for strategy/backtest logic. Real-broker + SMA accounting (e.g. credits/debits across days) can be more complex. + """ + if not self.sma_enabled: + return 0.0 + equity = self.get_equity(position) + im_req = self._initial_margin_requirement(position) + return max(0.0, equity - im_req) + + def check_margin_call(self, position: BasePosition) -> bool: + """Return True when the account is under maintenance margin (margin call). + + Margin call condition here is simple: equity < maintenance requirement. + """ + equity = self.get_equity(position) + mm_req = self._maintenance_margin_requirement(position) + return equity < mm_req + + def get_buying_power(self, position: BasePosition) -> float: + """Estimate buying power for new long positions assuming opening margin requirement. + + Simplified: the maximum notional long value = equity / initial_margin_long. + """ + equity = self.get_equity(position) + if self.initial_margin_long <= 0: + return 0.0 + return equity / self.initial_margin_long + + # --------------------------- Order / execution overrides --------------------------- + def _borrow_headroom(self, stock_id: str, current_short: float) -> float: + """Remaining borrowable shares for `stock_id` given an already-open short of `current_short` shares. + + borrowable values: bool (True=unlimited, False=not borrowable) or numeric max shares. + Missing symbols follow `borrow_default` (True = unlimited when allow_short is enabled). + """ + if not self.allow_short: + return 0.0 + v = self.borrowable.get(stock_id, self.borrow_default) + if isinstance(v, bool): + return float("inf") if v else 0.0 + try: + limit = float(v) + except (TypeError, ValueError): + return float("inf") + return max(0.0, limit - max(current_short, 0.0)) + + def _calc_trade_info_by_order( + self, + order: Order, + position: Optional[BasePosition], + dealt_order_amount: Dict[str, float], + ) -> Tuple[float, float, float]: + """Override to allow (optionally) short selling and to apply borrow fees. + + The original Exchange implementation forbids selling more than you own. Here we allow + sell orders to create/expand short positions when allow_short is True. We still rely on + most base logic (price discovery, impact, cost calculation) by calling super(), but we + adjust the sell-side clipping behavior before delegating to the base implementation. + """ + # When selling and shorts are allowed, temporarily relax the clipping logic in the base + # implementation by monkey-patching current position check. Simpler: replicate minimal + # parts of logic from Exchange._calc_trade_info_by_order with the key change. + + # Get basic trade price & volume info using Exchange helpers + trade_price = float(self.get_deal_price(order.stock_id, order.start_time, order.end_time, direction=order.direction)) + total_trade_val = float(self.get_volume(order.stock_id, order.start_time, order.end_time) or 0.0) * trade_price + + order.factor = self.get_factor(order.stock_id, order.start_time, order.end_time) + order.deal_amount = order.amount # attempt full + + # volume clipping (same as base) + self._clip_amount_by_volume(order, dealt_order_amount) + + # approximate adjusted cost ratio based on liquidity + if not total_trade_val or np.isnan(total_trade_val) or total_trade_val <= 0: + adj_cost_ratio = self.impact_cost + else: + trade_val_tmp = order.deal_amount * trade_price + adj_cost_ratio = self.impact_cost * (trade_val_tmp / total_trade_val) ** 2 + + # Differentiate buy / sell + if order.direction == Order.SELL: + cost_ratio = self.close_cost + adj_cost_ratio + current_amount = ( + position.get_stock_amount(order.stock_id) if (position is not None and position.check_stock(order.stock_id)) else 0.0 + ) + long_held = max(current_amount, 0.0) + short_open = max(-current_amount, 0.0) + + if position is not None: + if not self.allow_short: + # clip by current holdings only + if not np.isclose(order.deal_amount, current_amount): + order.deal_amount = self.round_amount_by_trade_unit( + min(long_held, order.deal_amount), order.factor + ) + else: + # allow selling beyond holdings up to the remaining borrow limit; + # later when updating the position we create/expand a short if necessary. + max_sell = long_held + self._borrow_headroom(order.stock_id, short_open) + if order.deal_amount > max_sell and not np.isclose(order.deal_amount, max_sell): + order.deal_amount = self.round_amount_by_trade_unit(max_sell, order.factor) + + elif order.direction == Order.BUY: + cost_ratio = self.open_cost + adj_cost_ratio + if position is not None: + cash = position.get_cash() + trade_val = order.deal_amount * trade_price + if cash < max(trade_val * cost_ratio, self.min_cost): + order.deal_amount = 0 + self.logger.debug(f"Order clipped due to cost higher than cash: {order}") + elif cash < trade_val + max(trade_val * cost_ratio, self.min_cost): + max_buy_amount = self._get_buy_amount_by_cash_limit(trade_price, cash, cost_ratio) + order.deal_amount = self.round_amount_by_trade_unit(min(max_buy_amount, order.deal_amount), order.factor) + self.logger.debug(f"Order clipped due to cash limitation: {order}") + else: + order.deal_amount = self.round_amount_by_trade_unit(order.deal_amount, order.factor) + else: + order.deal_amount = self.round_amount_by_trade_unit(order.deal_amount, order.factor) + else: + raise NotImplementedError("order direction {} error".format(order.direction)) + + # compute final trade_val & trade_cost + trade_val = order.deal_amount * trade_price + # base trade_cost + trade_cost = max(trade_val * cost_ratio, self.min_cost) + # apply borrow fee only on the net-new short portion of the sell + if order.direction == Order.SELL and self.allow_short: + new_short = max(0.0, order.deal_amount - long_held) + trade_cost += new_short * trade_price * self.borrow_fee_rate + + if trade_val <= 1e-5: + trade_cost = 0 + + return trade_price, trade_val, trade_cost + + def deal_order( + self, + order: Order, + trade_account: Optional[Any] = None, + position: Optional[BasePosition] = None, + dealt_order_amount: Dict[str, float] = None, + ) -> Tuple[float, float, float]: + """Deal order and handle short position bookkeeping. + + This method mirrors Exchange.deal_order but when a position is provided and shorts are + allowed it will update the Position.position dict directly to support negative amounts. + """ + if dealt_order_amount is None: + dealt_order_amount = {} + + if not self.check_order(order): + order.deal_amount = 0.0 + self.logger.debug(f"Order failed due to trading limitation: {order}") + return 0.0, 0.0, np.nan + + if trade_account is not None and position is not None: + raise ValueError("trade_account and position can only choose one") + + pos = position or (trade_account.current_position if trade_account is not None else None) + trade_price, trade_val, trade_cost = self._calc_trade_info_by_order(order, pos, dealt_order_amount) + + if trade_val > 1e-5: + if trade_account is not None: + cp = trade_account.current_position + if not cp.skip_update(): + held = cp.check_stock(order.stock_id) + # Account-level bookkeeping (turnover/cost/returns). Mirrors + # Account._update_state_from_order except for fresh short sales, + # where no prior price exists to compute order profit from. + if order.direction == Order.SELL and not held: + trade_account.accum_info.add_turnover(trade_val) + trade_account.accum_info.add_cost(trade_cost) + trade_account.accum_info.add_return_value(0.0) + if order.direction == Order.SELL: + # sell: update account state first (stock entry may be deleted) + if held: + trade_account._update_state_from_order(order, trade_val, trade_cost, trade_price) + self._position_sell(cp, order, trade_val, trade_cost, trade_price) + else: + # buy: update position first (entry may be created), then account state + # A buy that covers a short to exactly flat deletes the entry inside + # _position_buy; re-seed a transient zero-amount stub so the + # account's order-profit lookup still finds the trade price, + # then drop it (_update_state_from_order never mutates entries). + sid = order.stock_id + had_entry = isinstance(cp.position.get(sid), dict) + self._position_buy(cp, order, trade_val, trade_cost, trade_price) + covered_to_flat = had_entry and not isinstance(cp.position.get(sid), dict) + if covered_to_flat: + cp.position[sid] = {"amount": 0.0, "price": trade_price, "weight": 0} + trade_account._update_state_from_order(order, trade_val, trade_cost, trade_price) + if covered_to_flat: + cp.position.pop(sid, None) + elif position is not None: + if order.direction == Order.BUY: + self._position_buy(position, order, trade_val, trade_cost, trade_price) + else: + self._position_sell(position, order, trade_val, trade_cost, trade_price) + return trade_val, trade_cost, trade_price + + # --------------------------- Position mutation helpers --------------------------- + def _position_buy(self, position: BasePosition, order: Order, trade_val: float, cost: float, trade_price: float) -> None: + """Handle buy order bookkeeping against a BasePosition while supporting shorts. + + Rules implemented (simplified): + - If there is an existing short (amount < 0), the buy will first cover the short. + - If covering closes the short completely, the remaining buy becomes a long + - Cash updates mimic Position._buy_stock/_sell_stock (cash decreases by trade_val+cost for buys) + """ + trade_amount = trade_val / trade_price + sid = order.stock_id + current_amount = position.get_stock_amount(sid) if position.check_stock(sid) else 0.0 + + # covering existing short + if current_amount < -1e-12: + # amount is negative -> we are short. Buying reduces the short. + new_amount = current_amount + trade_amount + if abs(new_amount) <= 1e-8: + # short fully covered exactly -> remove entry + if sid in position.position: + del position.position[sid] + elif new_amount > 0: + # short fully covered with leftover buy amount -> leftover becomes a long position + position.position[sid] = {"amount": new_amount, "price": trade_price, "weight": 0} + else: + # partially cover + position.position[sid]["amount"] = new_amount + position.position[sid]["price"] = trade_price + else: + # normal or increasing long + if sid not in position.position or not isinstance(position.position[sid], dict): + # initialize stock + position.position[sid] = {"amount": trade_amount, "price": trade_price, "weight": 0} + else: + position.position[sid]["amount"] = position.position[sid].get("amount", 0.0) + trade_amount + position.position[sid]["price"] = trade_price + + # cash effect same as Position._buy_stock + position.position["cash"] -= trade_val + cost + + def _position_sell(self, position: BasePosition, order: Order, trade_val: float, cost: float, trade_price: float) -> None: + """Handle sell order bookkeeping against a BasePosition while supporting shorts. + + Rules implemented (simplified): + - If holding enough long shares, sell will reduce/close the long position normally. + - If not holding enough long shares and shorts are allowed, the remaining sold amount will create/expand a short position. + - Cash update for sells follows Position._sell_stock logic (cash increases by trade_val - cost) + """ + trade_amount = trade_val / trade_price + sid = order.stock_id + current_amount = position.get_stock_amount(sid) if position.check_stock(sid) else 0.0 + + if current_amount > 1e-12: + # we have long shares; sell from them first + if trade_amount >= current_amount - 1e-8: + # selling all or more than holdings + # remove long position + if sid in position.position: + del position.position[sid] + # remaining sold amount becomes short if allowed + remain = trade_amount - current_amount + if remain > 1e-8: + if not self.allow_short: + # should not happen due to clipping earlier, but guard anyway + raise ValueError(f"Attempt to short {sid} while shorting disabled") + # create short entry + position.position[sid] = {"amount": -remain, "price": trade_price, "weight": 0} + else: + # partial sell + position.position[sid]["amount"] = current_amount - trade_amount + position.position[sid]["price"] = trade_price + else: + # currently flat or already short + if not self.allow_short: + raise ValueError(f"Attempt to short {sid} while shorting disabled") + # expand short + new_amount = current_amount - trade_amount + if sid not in position.position or not isinstance(position.position[sid], dict): + position.position[sid] = {"amount": new_amount, "price": trade_price, "weight": 0} + else: + position.position[sid]["amount"] = new_amount + position.position[sid]["price"] = trade_price + + # cash effect same as Position._sell_stock + new_cash = trade_val - cost + if getattr(position, "_settle_type", None) == position.ST_CASH: + position.position["cash_delay"] = position.position.get("cash_delay", 0.0) + new_cash + else: + position.position["cash"] = position.position.get("cash", 0.0) + new_cash + + # --------------------------- Borrow availability helpers --------------------------- + def is_borrowable(self, stock_id: str, amount: float) -> bool: + """Check whether the requested amount is borrowable for the given stock. + + If a borrowable dict is provided, it may contain either booleans or numeric limits (maximum borrowable shares). + Symbols absent from the dict follow `borrow_default`. + """ + if not self.allow_short: + return False + if stock_id not in self.borrowable: + return self.borrow_default + v = self.borrowable[stock_id] + if isinstance(v, bool): + return v + try: + limit = float(v) + return amount <= limit + except Exception: + return True + diff --git a/code/tac-qlib/tac_qlib/contrib/data/__init__.py b/code/tac-qlib/tac_qlib/contrib/data/__init__.py new file mode 100644 index 0000000..c76a9f1 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/data/__init__.py @@ -0,0 +1,3 @@ +from .handler import TACHandler + +__all__ = ["TACHandler"] diff --git a/code/tac-qlib/tac_qlib/contrib/data/__pycache__/__init__.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/data/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c1543fabfc22e07c7e9942015c4462f84acc1e91 GIT binary patch literal 216 zcmX@j%ge<81kW_qW?2F0#~=<2FhLog#ej_I3@HpLj5!Rsj8Tk?3@J?Mj8ROL%$h7O z8G(|TjJLQ#9GyK9^HOqBi;9?mLVlXex7ag~1a7g%$0z3G#K*5>_zW`mm%e^tL4kfr zVzO>wPG%B_5f5f0=jW9a0R>VLOA__t<1_OzOXB183My}L*yQG?l;)(`6>$RfgX}Hl a2NEBc85tSxGRQyR7Vpq&WG`X|iU9z5do{KI literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/contrib/data/__pycache__/handler.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/data/__pycache__/handler.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..89f5cfc15b946b6ea1fb7724f021024967cdb6c5 GIT binary patch literal 21533 zcmeHvYj7Lam1Z~I#FGHuF9{MQijYJ~v}H-AY}tBHwk#T!NGg`;I3S2_lAu6paI@zXVIDNIICIQ0VI8zi*amGA6@wKM_Cb3g zuVc{3zHtq@aks=g6P1HrmS&ClCaMOj*t0EGJyA1QgXapt9;=&0vdAJGeopdzBm9 zC~Oio3-zy>2mJ=_DNbm36_iEq_+hR2#rxYv^?7hJ%h`yWO}{5+E6dr8oGrg6=N6W8 z19EN@tMrn#j&g$ktNdV_xLNduHfwbeTk+Paziq+W7VT}u7-$VYY~Z+oLr0zt$Ay?E zb^F8qmt)ak|Iu(VT!`yE=}(S@lm6kUXiSj(0#al#Df>sH!~{}A|9})0#6w5?vG7^3 z({3Lai^_idr=-o1DLIL3CFC~y#duPh@lPhA@ua`~xyfWS5%(Vsbh8h})OyhOk~ktt zVmu-WdQtZ2gmiW!mYDX>gePL#WdG1mp*o#~4~K^Qpnz|~sOoc1_a6C}e}{ik5<4P^ ziOH!XXeGlP)W8ujoSc$G*)J!haCCGmY1i9B?eI^ArI)9~q<8D;N zvQzO0x_ML-I_)3OBnt3~BOH$>l3^M{If!@nsF)m?ip9dxs5~h~0w%@Oi%Karb)3U7 z#XK-I850%r={UZ@7ZrMg@!~3(D7_~UQ3ZOTNEif%V`Av+G|~NlRu?GSdofT)67iAf zsNzdTC&UpcJRycg5>jYH5?@v-kM$fHINjG1I(F#U-DUDw1oJh+ zoNY~eE46%D-6f}BNtsX=TgoI@8i>A4uslY0<2a4$x+`TXmV4>eN8~Y;M*8gR}vCAZ?_;hd^;mB!nrf!^v=REMQmcQ8^kX&K41sN_`oJ)J#f6Vd*3! z$w3ronkgn)z=~yR5=|EsTPPH#B^y#qLNt<8Ea6Ek3n5UUScqIGR;++xj*HWwn8ITf zon=eZI;pNvi3FPBU>~4k!3|`ZjehO=ZRakK-31QaL7h(zj8*R{jKVG&gt(K$Jz zRF@kGX*Xzg;ZOc+1Q)pXo$d=gKXFy1`DIVT!luQRj-{4;*_M5ap8d-$yKj3J_Z-c& z^jr+4d#>zTuG(~CYp$w&sj4Ge)sa4$_iVheC+pdEGk8bH?mV2Y+nTBCx*f~x4X$to zuPdHXci#V5MIhSoCbZ zF?Wa0ZhJJ}(Rp(&+wu4n$HKGO%I$fl_j2#W-i4+gIk&!Fc=mq&xs?{qUA?k_b9yiI zuA)KQ8ZKf+1mhwptmV&ET;)nu5>}zMg1}phs&?r;`il}oE2QstCgR!Xndic zE0nHt3bpj|c>}iayBLC_Q8|)0CrX;u;*XGMh|8qe$WW)i5&YXlWSH5mXCcLfSPrRS z16E@!zp=kh!mb=l<8uz@JwO~s@Y_0N`^Ky;mC}C zBnB;z=|h1|JJA|A#FVZ6lX`uKe+W@|oauJMGl&O&kftI;m#HXJ8>nIv&|08GLd|1@ zbmeT4+Q_QqZ`ax=kAbS-2Q{bxetBkMI1y8;p*D{y#Xbq_WW-PeR3elEVjTry=%%R{ zbXKb06e%9iVf=j}G(dDB^e|EknWXj;KT<2zJ{?VtQAZH1p`k6{#pv#l==ngWzh`1H zIpdFxfLUrCt`bL?GV~o%k*J3xlB27g1T5B_}bTy9Kim}?0Mz2UEsmpCGZ5A-+^n${0663L9rp zn>jwkOZ=syJQwG&a`^MS48yHRM_Y0ITq{>xwOcsJ$a61nmxiXf^TrprX+B^G&Q_|! z`uO(Fg5DjlDi(EI6w_!dF{~IdT@(lD^&uJjc_IvfKtm;wTvsZ3&PPOMNd(M_6@wlU zqLN}!XPUHwvY0TRWW_M~vP}I&gL6UJMG5YZI!Qum?%6J;a#&; zD{EFfoUePOlB=wJ?bFw)zFd>BZTR&)2Up$lA;-Irp?<}P`)_69IbYp#=wZY69L4d224^nb%p691iNC9*us$ZA}UZJks=1IqF^tn9m^~Y zS@BHHge564t&N&uEes#~Rw6d6SfVn0Bkdt-cBO)NQAkQ8Anhl@(pijZg|53R#>mL1 zVxzLyHwW>JbbzQBQEgU7>2bWuG(~9$+4ZD4jtw6TEI8W*N-$T!4=P}a9IZOuD;mk(Y%_#;PC-szi<-*`S(+mUs4Ts*d76e3 zSug|2u*?_&Hi>3R_g$GFn~yAG)Ka5 z|5zdhUJ4U%rbAS@WE{F(I0*|=T=V+%tteN8Au`tcSX2Hp)VXdMtMpie-nRW}V z+fU}KMkIonH$(iF*~era`sL_oJRD;(X@b<=p`oJ1)`?mV4KcbfBIAAtAhmjOEDBqk z38!fYFX{?3a>73>LJLRhA^#GHfT^w>dQ!b&(S0RI5ahFI^U1kdW z*>*4zG!5DpKZpBxDHRZ*>y^h4Z&II2!&i1!=o;!AWEXTt$ikv-TAo^YTAJo87~e5z z(A}J&q|UnLyr#~Ea#rq9?y@NY{W)S_dawyYXCfsmJc9`!i02g> z8Gpp|k!+4(p2q?X#TXOg#i^szF+-Ud$P{3Wbxr~FlBEH(N-?ulD?Lij)&fUROi?)= zR*c|R(n%x)%+e`JHwYt&MYWvOkNJqM0u-jE5estGT{0{2{waJbe;)xRq{Ef&zE-u+ z_vN}9hyJ2DUsabrk@r@mEqRA$$+02p*pP8-znNUx^<;L}lbKz;D+a5p`KOgVf8Ve@ zUtN36ydW%X=*n*By1i%V;iK7ykLK!o?hf8F8XA39Ecfi3(~~~;xdRKEehAsYNVj)> z)BO#1uKCn0_l>SQZSQ&C4Scur-OgP1GfUmUYCQ? z>I@#RsqDEe!>;VKcco?JRA3~T{)m+^-zwgGLrPe*!);)L5 zT-leaI+k%BV^W?m7pS5^={bBBuu3FK6?>?FW&omahC(k-g#r38z`_~|2|y@9QXle3 zLllf4xWzF-jN>r@%Yp<_mx^~1dw-u|g!2oiuwVTd^-k|G8h zl1w!cGyv8i`yeI>NGBUf2jC&*0buSV(~33%*do@H0A@>(++#9|1RX#r=w#~0&D<0s0;hI!$M>A+wQ06ZYKAJhNMu`-{F@+td=mqXj=fuBrT{{SUSQsi&hq93xju8~* ze-^+v`igRh%vwC$bL_dk9#Rh&RkAgyL!gjI6Ohic3=(G&8+D>OQ-?&8Bp|63y2T4! zL2G)UvkTT|Lc;C=eKZn_PRhXcF^RK zV3ctW0lOiI!pN2b?l6X=mKIP{Y6eJaNGh-Z!dhc9g}w_uH_+3K3Zd(n=C^K^tsC8v zsfO*FKVYr}(DmpAea5K+GXbLu2Cb^Vu~F&QScN6w3Q$|O@WSkgX^>G zA$8}bLX}2w%Cqic`c5&?##8a@wLr8g6i8mo*0L_wx?5L*zkxLQ1cD3PN-fvenl@c> z>soNbX7*a!jGH6hu=$4#2N&D+->J^E9b^gjY0dr;tyx#;L0s3G{}Tl2N6GhBoaH5X zKbI_yQdxZ&-{~WQ(h8R*{RuT1*GE-L1xm%0+b=?_j`tSx>a?(P^gyiFT|`x}?_?m1 z`0YYZe1=czdpzU7f7f3n7`Cw_CkgR&SIQ_DnBGmnx+_VOgk4&>RW z^o>MXqX5gHpskH>T9=h_6`~mOzWM(2@vZmAY2TzwA>DQWZolpRoPxf)_Yvtj&Q18Z zU}eam8Cbh{yj)#KLtA$Vws$JDmU=j9H&(Hp=cDFiz#g3KShJ?OOT^D`7^pJA;k30i zk^p<}hE4&Z3ZJ|M6NFhM(ck{=cfb1pf4}>Hrify~CKIKa$ifcPNTkCtEUd^d6dNeB z%`lZ~Ft%`*KK5dJ7+dSfV^mlVskJ~J1;xqm%Yp<|Ed3DhB6h=2(99&Z#NpWx1}w&; zrK~u}t}f_Uik-kv)qIkYR6>QOXy7}>(G|OZjdem=L?K{N`P6Z$^9~9~e^z)9E#YP| zEmN51xKs<>8$tzyQieCXaZRUo-k0G&fQ0NuhiQs-*3q1Cw8M^Bs@jsR+LCc@S$5Yi zRA1ll#)g~K-)y?olyg4<^vvmDyHZ;GhFcqQ&i!dV@2R|e{^I#dv&+8LOk00u;B=<{ zW0@__=X_@}&NIs{&*jj?(88&lYh&85>~vo~e)0H~;f1RC(W~+K_}5S8YIiKvK9sF} z=sTOqV7lwOd*laox$aYo&b}q*OIhbjIcErk_^PgMo8PwNYt8yvbH1%<{*v|m+PbTw z^P>yHub;hf=&h&TeCp=mubs%&?z+7xTl;X@k#F91vuP>Noegv^1$wf9p1XT9fu3BT zFVo(aZSL31uqBT_>+xr7{`;W{np0q${3>-iG2MPuAW39gNg{P4t?I&1FoR;&ySDS8 zx6t0>D)Dc0H(7o6mqX-PIom;wF?T9?-K8yfKo2@IF9&^)_)-)xk%ES-u}#RKWslTBKE$_a~9MT$8qBy~|yG_ApCSz5Vr zi4~{pDGOPYxLZ>;db;!iwr;G<9G7StgAf^kxKW@|x`=m)*tX&&Xq~yRV3UOHJgV9c z<3yE|$WWY6fpb|Uk+D=MMIVz-4>J^xGw{a(og1E2*uF_5;Xv~WWx?JI_T3HfNyQqL zf$N56qzi}-IEu!xaDnmD+3bJB^NhQ-9B;HoI zRPAd7k@ss~Drk^Fn*0+4SX?gelB+fAYR$Q}g41|5E_u4Lp03-QZwpI%pUUoiD(87R zZOVfmow;~s$<>y1wcVJ_xgJUz^1kY;+vm4m>sxB+$hLIc9R6nXR`l(&xt8vn@4%Ap zNY;1cy~8=*Q`j-^)Lx#uIQO|vFL}0RJzGJ}U3q!*;^^n1Y2#0@i*yKp#g*{V`v6TwWSBo8p%Br-TF^qssG4`K)L85sUR1By4gNjiI&uG|(jbR_sC{m>{ z3ZfJ+oP$k?A$p~bRC!kyJ?)|(rUm~B5Aq`j2=DNkZ7bWks``9QGa3Q5MWM@=uWQQJ zG^|$n%$q)JcbPY>Zr}ljSe)i6R!2rG0uuT~>=75gdKBSsOd*dZaCTLmqK0bKyh(4dQgfq6O} zi(+RY?4OK*Scy@l*?Tan+hvjVF4i#rp&@mGN9{DnsZ#h1 z<-&ZU{SVsahqVdkpL5Z3*mhrwn5?yByI~;vrzSh5MG@`+M13~fn-#{dZ8cIJnQeHGZF?S)=I<5id!?FR-W5Isz5ncy5DjNMN^bFeEZ zmL-W{>?#$(FaJ5!JwhMXo&I5L0HHAehNiGN1h@(Yb~GM|O|c?|hOn`O(J7!1v?G;B zoDGjret>Ck2w;GVYD%lyMmlYBt?*CDY^YQtANyxgbQZG!8)4LkZ0fNpi|5JGmjQQ? zvuLRo;NJV40?`G)!vt|848I5exhTzy0%eA!3R_K7b~MbQjYM|EkWu}Xl93S+3Q zEYxmoBa2BLU_DrNv1uAiObEgziayH)XA~XwI^Ylgvrr_WTEK(+s%IH%hYB2_*B=N1 zGG~)A0H6-_XkbaHTO;XS_~TdqiE*@lX|N=zmz}b{UC;_k0h5w+-PN%aLMchtO*Eue zpz`d3R#CcT%p*%Gzp@Nvh;80>EBZD76qhZG5<%IYFp4G)LRLf)5usgu`5k* z=!g#a*2yvgAUYfVNJMA3wqc>|^|J)`kl)C@*Hd_OE?3trm|j1Q$BKN*cCtJg@7pk~ zZA_a!UoZ|z)qk=a*o2$ai~~fhzy3e+qb-iP|3{G@m|CRvue-|QLGS2@&mV!-Sc9`& z_pO;f2(ghQ)JmcSB~4Rs9zmd5x=GKpM2p~zG)u1x&DC)k>5~*kR*5u6!KWzrGzAwZ zAX`NG3wV2#gQc)c9xPS) z7of+ml*_R014+4hVE#bHx#6FWqdcEE^FrqN!OWHybH0yfoFD%e#8L9Kbqn6> zb#K%yHSEYX?6_Te(Zu+6=P2pJp6$eNzMHo z3!J1Gq@>S*%Ke4dv_dXm7vc8Tgvi@JG#inBE+LPWB74VL4)nbeHRE1FS zj$LcVoIPbPw@>g9XI>sSE^FvIlH0U0lL75%e7p8kyfff&XYpNo(rT!U`kWIpov?>> zm)>4i$_Y^JEWj1vs9?n3Ot>!R?g)h*?DP%3#?3l6Ql zY#*gERXI+e%DPLVH|IrPlvyF;d)Iw{*Oa$h&X5i~BE4KcxXOK{jm7w*Wy&e^xt1s8 zO1TNdyY!y{JGuZnI)l2>z37Pp_CE!ek$%Td5rf`pof$!7;H(bEFXNu4K17Q0L5 z6%()&872{K(jB7wE>SR%sZLIE?8q5|P^B+ZD%$`v!o*k1Fvz2VVxtpQgrCYjq^~h{ z*h>k5TGrxGQPho?xbj072xu{t)k~Fa*~+$DWqaDZV&Od1m(N~2yKwaS@i&g&G=0-~ z%b9K7bNl6N^ZrcZqd9l?ozvfa@!c2iZuy(Q4+3`^bB_;X-2<8D&t%qutz(ah1oOjAMI0ZMzta&d6U zwK413c*B`^UTSx8~?i~H@iFZ%j?aDsZo7?kjI+k+>m)w0> zcVDLe`5e_f@AmzHU*y~a$nNo7p1wGJMZV!pPv@vU?*sj9pKo7yn2Nk1y>L;;n%{c0W4;63 z;p8CFTF9sK#{`YHImTI0i_IWZ<-P9gjij2Eq@Xpnx<( z#xbO3JpBQPHCyFF^wkIj;}jGT>#x!)1%HP>`4EB!A=VymzS5g_x%0If3g~rHqq*tB z{SI@}Y9m9hn>F;BXh!g-Uj$>|;#dD1;b+(3*MdR7Ccj`3O`@5eL`yk`UzTD+XdZE= zxe@##!mIeb1ko-y1na9NO$-b=5aSeVB{5E+0^Wcw#JiB!{;FBe3p(fdK@S+QTd@}* zwQvxN7~3`_KLIcd4oGOX>SQ-iIP$MKgOwaGI~k@^*bMfZz_Aa2YC}VZVv}Rxu1EF{ z4G}(GWZF354!-U`U{ZMQPJ45Vt~qH z$S8^d8d>Ejr8btbb(9Vyj997w|-zX}Ixb|x%k#aSx@^i~_4 zxmt6Q*-Tq^y6?ET#CP1XCWUm@G8%Wv0qX!#`0axIGUrG+WJmGTR#78|=FgQqPv+F5 zIXnGbIp>8f0xBbhRAtILf~kDX04NZrGOuB(;vBZG=*4xhD#d*sI*MS0Hinj^e6(c= z?bCA2n(};|{{qkBm_fi6oPD`?vJSo*bcWs^J_it=9rVyTn5NwUc93xW1i(H7IYUIr zWl}wE^5QlXaNf4lzikT+>4m!Xut#|!9C~Q?wwD5m`xK4`)2vs2+W~&ae~EEnj0k5N zaF~qF7Ap=&WX$dm{c?i}9Fnxec{NALyoSGqKB*O(Hlq~>5aKZNp%5jda{cc;D3$Af z>OpZ9@Z%6|$ui~z*N)vNPiYA04j}7JsVVl4k-Md;yi}HCV{;bXB2ys6w(|g-oT3)y zC|IQ6M-(tLbcvqw6zCYJ^kaJcAp#saSJkz`!8q1OCD02-7{1f2yb!Gm#i2f{C-(l9 zQY+vzO2@cEA(?X1%3gO}C_z}Oio8jOR*#XQ@iTn;YhbHwoU<|G*pzSCmJa3}bs5LT zd{b+>_vcN!GP|FAj}A}gnnD>{{qp*V_ZwQ48un!y_T6#cHQkM7UI^tHhSKgIds>!# z%_|1p+r4TuR=HO=qsxsw)YffFoA+fm@4I8pZ9bT_ZA{y*oPb79+lFA%uG{X+#wYR{ z+Hda2wCw$*+1TV;v2p$#`TFKxx~z5XUz+W8_O#<(Emz^bXwBN{SGbA_d%nIkU$-e= z(?Sm|9rV9p>n}Y>!D&BF&8inqDvVJIIatwfEN_0*2~~?=j0?v~SOJz+(X#^ww$W8| z>1@MPT<}Q}sqE5nrA# z|NP0Z<2}7c`=w2kV1!iy#sH3)#0@`9FGM#~DhdFVsMi8WSjVdyQOqhe-t}=g_ilRkR89d;Ap@aU#{uAIdi9N z@$f+A^b5JegSna))5n%uHZ5-6lWTe82W`3i=YG)k-l@#M^O?aHvjZ=_lSg_;PqE7fi3s!6Le9o3|}^zv#$UR$aBv+ZVdN z{en)`O;k$sl$SP^`0XU&nvOy4*7C~Vw# zls82ROB>lre*NV&#IF6-hT`wNXqbWo4PlfQ_)&hu0Iy!14i>D=Xo?GdL_fbH{b}jq z)|iH1IE%?4Jyl9a^;<1d@}3la?S-(I(KRRE1$KVz_mu7_vMx|zUubopc4%>hRsrKx zuau6x7EkV#YllV)OoA`4K(N|5UNEQlS+Ghlk_yJ^)Lzq>7-5B9HkA4o7nUiH({NUt zb?QBIxTklwbQ{%SkOZ@0X00rGt1#OLlwoReQvE3cCN0T11PVvDABT}P(O18qfGuQm zl72%;dL9ul-svZql!^gp2u$WfZz0=B;iGFj3YQT#+f;tVvGV3I1LXgRTo<^X<8-um zt$z^$wYuTTzJ)Uj&t$6FGS0STmLPn+XW^wq$Bt!>@5&>en_W2d+NZzXeN)5=82gTK zvGuWheak}l4cnF3ZwK4qeANY>>JCXOcetUdz|M8svne6`Kx4UjM+&uNQmdyU+SwE`t zOvd@lvb!$Rxaaoi+s8BY2XpQxGPWmv^P!R3bb@ETt>545Ji=MltPq;zOl*=ph82QP zy$bG_@PBh1?KyVnbnifj4NZV&I+9Eu7Tjc2hc|MK0Q5pbw5D5`H(W{LHR@IyzP`XM zdzv%0W>uD>68fdf@=ippUA8!As^=`aq-%0qtF)oW>2W@{$dgHvTYCt%4!_R9eoADW z(m1=hWM+Isp6k>Sg4r=aJnc6dZXa8GEOqW?t zmoOh$OcVRzuNU?BoqqlCB?v8M6){U6u7_ei{wD>$MxeL~ z`e{U@Lxz|Ie~(x;3!YT<(*H(TDL6>oLO9~D2|F;u+s((-G-BTlKLtU4ZVg@ZpLx6G zYc?;{Y|qxf?y*M|+@N}4!;;mJbQ{$*k3b1FQ?%|@3|)fs8_K(lIr-6f$B`-e zT|69Yq8PSbOTOubYUGg2;XjMEWf~ z{cj5XkpedEyXmQyf>sLF^$PWrf-{}$Na;G1NU=Sp{-o7G$$_!Pe@=EH#1UJb|D~aV zH~p%X<6D2mRs59O{CC{09JlLdobM-G)jx6ve#$-kQ*QhJGMRb9Du;kn`lcJhSyLeI zZd`Y-y|FoK+PYF4NdSGG}S9&62YRGu{X^OMhwgFodhb()2+x0vztVD#V-Fwp8ac-}f5E!q List[str]: + """Discover feature columns present in *every* feature file of the lake. + + Walks the `family=ta|sp` partition layout (plus any legacy flat files). + TA and SP columns are disjoint by construction, so the common set is + computed per family (columns shared by all symbol files of that family), + then the per-family results are unioned. Returns sorted field names + (without the ``$`` prefix). Empty if no features are persisted. + """ + cfg = LakeConfig(lake_root, market) + feat_dir = cfg.features_dir(timeframe) + if not feat_dir.exists(): + return [] + import pyarrow.parquet as pq + + def _family_common(fam_dir: Path) -> set: + common = None + for p in sorted(fam_dir.glob("symbol=*.parquet")): + try: + cols = set(pq.read_schema(p).names) - set(NON_FEATURE_COLUMNS) + except Exception: # pragma: no cover - skip unreadable files + continue + common = cols if common is None else (common & cols) + if not common: + break + return common or set() + + common: set = set() + # family tier: features/market=*/timeframe=*/family=*/symbol=*.parquet + for fam in FEATURE_FAMILIES: + fam_dir = feat_dir / f"family={fam}" + if fam_dir.is_dir(): + common |= _family_common(fam_dir) + # legacy flat: features/market=*/timeframe=*/symbol=*.parquet + if (feat_dir / "family=ta").exists() or (feat_dir / "family=sp").exists(): + pass # family layout already covered + else: + common |= _family_common(feat_dir) + return sorted(common) + + +class DropAllNaN(processor_module.Processor): + """Drop feature columns that are all-NaN over the fit window. + + The lake can hold fully-empty indicator columns (e.g. a ta-lib output that was NaN + from the start). Such columns carry no learnable signal and make ``ZScoreNorm.fit`` + warn on empty slices, so we drop them before any other processor runs. The drop set + is fixed on the fit window once (during ``fit``), then applied consistently to every + segment so train/valid/test keep identical feature columns. + """ + + def __init__(self, fit_start_time=None, fit_end_time=None): + self.fit_start_time = fit_start_time + self.fit_end_time = fit_end_time + self.cols_to_drop = [] + + def fit(self, df=None): + if df is None or len(df) == 0: + return self + window = df + if self.fit_start_time is not None and self.fit_end_time is not None: + try: + from qlib.data.dataset.utils import fetch_df_by_index + + window = fetch_df_by_index( + df, slice(self.fit_start_time, self.fit_end_time), level="datetime" + ) + except Exception: # pragma: no cover - defensive + window = df + if len(window) == 0: + return self + self.cols_to_drop = [c for c in window.columns if window[c].isna().all()] + return self + + def __call__(self, df): + if self.cols_to_drop: + return df.drop(columns=self.cols_to_drop, errors="ignore") + return df + + +class BenchResidual(processor_module.Processor): + """Subtract a benchmark instrument's forward return from the label, per datetime. + + Turns the training target from an absolute-return rank into a *residual* rank: + ``r_i - r_bench`` is ranked cross-sectionally by the downstream ``CSRankNorm`` / + ``CSZScoreNorm`` processors instead of ``r_i`` alone. Must be inserted BEFORE any + per-date normalization so the ranking itself is computed on residual returns + (ordering flips exactly where the benchmark trends). + + Stateless: ``fit`` is a no-op and the benchmark forward return is recomputed from + the lake parquet on first ``__call__``. Rows whose benchmark value is missing are + left untouched. Accepts ``fit_start_time``/``fit_end_time`` (ignored) so + ``check_transform_proc`` can inject the fit window uniformly. + + NOTE: under any cross-sectional normalization downstream (``CSRankNorm`` / + ``CSZScoreNorm``) this processor is a mathematical no-op: subtracting the same + per-date constant preserves ranks, and z-scoring absorbs constant shifts. Use + ``BenchBetaResidual`` for a target that actually reorders. + """ + + def __init__( + self, + benchmark="SPY", + fields_group="label", + lake_root=None, + market="US", + timeframe=None, + freq="day", + fit_start_time=None, + fit_end_time=None, + ): + self.benchmark = benchmark + self.fields_group = fields_group + self.lake_root = lake_root + self.market = market + self.timeframe = timeframe or timeframe_for_freq(freq) + self.fit_start_time = fit_start_time + self.fit_end_time = fit_end_time + self._bench_label = None + + def _load_bench_label(self): + if self._bench_label is not None: + return self._bench_label + cfg = LakeConfig(self.lake_root, self.market) + p = cfg.bar_path(self.timeframe, self.benchmark) + if not p.exists(): + raise FileNotFoundError(f"BenchResidual: benchmark bar file not found: {p}") + df = pd.read_parquet(p) + s = pd.Series(df["c"].astype(float).values, index=pd.to_datetime(df["t"])).sort_index() + s.index = s.index.normalize() + # mirror Ref($close,-6)/Ref($close,-1)-1 on the benchmark's own calendar + bench_label = s.shift(-6) / s.shift(-1) - 1 + self._bench_label = bench_label[~bench_label.index.duplicated(keep="last")] + return self._bench_label + + def fit(self, df=None): + return self + + def __call__(self, df): + bl = self._load_bench_label() + cols = processor_module.get_group_columns(df, self.fields_group) + dt = df.index.get_level_values("datetime") + aligned = bl.reindex(pd.DatetimeIndex(dt.unique())).reindex(dt) + mask = aligned.notna().values + out = df.copy() + for c in cols: + vals = df[c].values + res = vals.copy() + res[mask] = np.asarray(vals[mask], dtype=float) - aligned[mask].values + out[c] = res + return out + + +class BenchBetaResidual(processor_module.Processor): + """Residualize the label against a beta-scaled benchmark move: ``r_i - b_i * r_bench``. + + Unlike a plain constant subtraction (see ``BenchResidual``), the name-specific rolling + beta ``b_i`` makes this survive cross-sectional normalization: in up-weeks high-beta + names lose rank, in down-weeks they gain — exactly the relative structure an absolute- + return ranking hides. + + Beta is estimated from *past* data only (rolling ``window`` trading days of daily close + returns of each instrument vs the benchmark, both read up to and including ``t``), so + no lookahead enters the target. The benchmark leg uses the same horizon as the label + expression (``Ref($close,-6)/Ref($close,-1)-1`` by default via ``horizon``/``base``, + matching the yaml's 6-day label). Rows with missing beta or benchmark values keep + their raw label. + + Requires ``$close`` to be present in the feature group (it always is for TACHandler). + Stateless; accepts ``fit_start_time``/``fit_end_time`` (ignored) for uniform kwargs + injection. Must be inserted BEFORE any per-date normalization processor. + """ + + def __init__( + self, + benchmark="SPY", + fields_group="label", + lake_root=None, + market="US", + timeframe=None, + freq="day", + window=63, + horizon=6, + base=1, + feature_field="$close", + fit_start_time=None, + fit_end_time=None, + ): + self.benchmark = benchmark + self.fields_group = fields_group + self.lake_root = lake_root + self.market = market + self.timeframe = timeframe or timeframe_for_freq(freq) + self.window = int(window) + self.horizon = int(horizon) + self.base = int(base) + self.feature_field = feature_field + self.fit_start_time = fit_start_time + self.fit_end_time = fit_end_time + self._bench = None + + def _load_bench_close(self): + if self._bench is not None: + return self._bench + cfg = LakeConfig(self.lake_root, self.market) + p = cfg.bar_path(self.timeframe, self.benchmark) + if not p.exists(): + raise FileNotFoundError(f"BenchBetaResidual: benchmark bar file not found: {p}") + df = pd.read_parquet(p) + s = pd.Series(df["c"].astype(float).values, index=pd.to_datetime(df["t"])).sort_index() + s.index = s.index.normalize() + self._bench = s[~s.index.duplicated(keep="last")] + return self._bench + + def fit(self, df=None): + return self + + def __call__(self, df): + bench = self._load_bench_close() + + # benchmark forward return over the same horizon as the label expression + fwd = bench.shift(-(self.base + self.horizon - 1)) / bench.shift(-self.base) - 1 + + px_col = ("feature", self.feature_field) + if px_col not in df.columns: + raise KeyError(f"BenchBetaResidual: {self.feature_field} not found in features") + px = df[px_col].unstack("instrument").sort_index() + rets = px / px.shift(1) - 1 + bret = bench.reindex(px.index).pct_change() + + # rolling beta per instrument using data <= t (no lookahead) + cov = rets.rolling(self.window, min_periods=max(10, self.window // 2)).cov(bret) + var = bret.rolling(self.window, min_periods=max(10, self.window // 2)).var() + beta = cov.div(var, axis=0) + + contrib = beta.mul(fwd.reindex(px.index), axis=0) + cols = list(processor_module.get_group_columns(df, self.fields_group)) + out = df.copy() + for c in cols: + lab = df[c].unstack("instrument").reindex(px.index) + resid = lab - contrib.where(contrib.notna() & lab.notna(), 0.0) + new_vals = resid.stack() + new_vals.index.names = df.index.names + # residual where available, raw label otherwise (e.g. beta warm-up rows) + out[c] = new_vals.reindex(out.index).fillna(df[c]) + return out + + +class TACHandler(DataHandlerLP): + """DataHandlerLP backed by the TradeAC parquet lake. + + Parameters mirror ``Alpha158``: ``instruments``/``start_time``/``end_time``/``freq`` define + the queried window; ``feature_fields`` selects the features (default: raw OHLCV + all common + ta-lib columns found in the lake); ``label`` is a qlib expression for the target. + """ + + def __init__( + self, + instruments="all", + start_time=None, + end_time=None, + freq="day", + infer_processors=DEFAULT_INFER_PROCESSORS, + learn_processors=DEFAULT_LEARN_PROCESSORS, + fit_start_time=None, + fit_end_time=None, + process_type=DataHandlerLP.PTYPE_A, + filter_pipe=None, + feature_fields=None, + label=DEFAULT_LABEL, + lake_root=None, + market="US", + **kwargs, + ): + # default the processor fit window to the queried window (like Alpha158 without a split) + if fit_start_time is None: + fit_start_time = start_time + if fit_end_time is None: + fit_end_time = end_time + + infer_processors = check_transform_proc(infer_processors, fit_start_time, fit_end_time) + learn_processors = check_transform_proc(learn_processors, fit_start_time, fit_end_time) + + feature_fields = self._normalize_feature_fields(feature_fields, freq, lake_root, market) + if not feature_fields: + raise ValueError( + "no feature fields available for the lake; set `feature_fields` explicitly " + "(e.g. ['$close', '$rsi_14', '$sma_20'])" + ) + + label_expr, label_names = self._normalize_label(label) + + data_loader = { + "class": "QlibDataLoader", + "kwargs": { + "config": { + "feature": (feature_fields, feature_fields), + "label": (label_expr, label_names), + }, + "filter_pipe": filter_pipe, + "freq": freq, + }, + } + super().__init__( + instruments=instruments, + start_time=start_time, + end_time=end_time, + data_loader=data_loader, + infer_processors=infer_processors, + learn_processors=learn_processors, + process_type=process_type, + **kwargs, + ) + + # ------------------------------------------------------------------ config + @staticmethod + def _normalize_feature_fields(feature_fields, freq, lake_root, market) -> List[str]: + if feature_fields is None: + common = get_common_feature_fields(lake_root, market, timeframe_for_freq(freq)) + feature_fields = list(RAW_FEATURE_FIELDS) + ["$" + f for f in common if "$" + f not in RAW_FEATURE_FIELDS] + elif isinstance(feature_fields, str): + feature_fields = [f.strip() for f in feature_fields.split(",") if f.strip()] + fields = [f if f.startswith("$") else "$" + f for f in feature_fields] + # de-dup while preserving order + seen, out = set(), [] + for f in fields: + if f not in seen: + seen.add(f) + out.append(f) + return out + + @staticmethod + def _normalize_label(label) -> Tuple[List[str], List[str]]: + if isinstance(label, str): + return [label], ["LABEL0"] + if isinstance(label, (list, tuple)): + if len(label) == 2 and isinstance(label[0], str): + return [label[0]], list(label[1]) if isinstance(label[1], (list, tuple)) else [label[1]] + return list(label), ["LABEL%d" % i for i in range(len(label))] + raise TypeError(f"unsupported label config: {label!r}") + + # ------------------------------------------------------------------ utils + def get_label_config(self): + return DEFAULT_LABEL + + @staticmethod + def discover_feature_fields(lake_root=None, market="US", freq="day") -> List[str]: + return get_common_feature_fields(lake_root, market, timeframe_for_freq(freq)) + + +__all__ = ["TACHandler", "DropAllNaN", "BenchResidual", "BenchBetaResidual", "get_common_feature_fields"] + + +# Make `DropAllNaN`/`BenchResidual`/`BenchBetaResidual` resolvable by bare name from processor +# configs (e.g. the default ``infer_processors`` and workflow yamls that reference them without a +# ``module_path``), mirroring how qlib registers its own processors in ``qlib.data.dataset.processor``. +processor_module.DropAllNaN = DropAllNaN +processor_module.BenchResidual = BenchResidual +processor_module.BenchBetaResidual = BenchBetaResidual diff --git a/code/tac-qlib/tac_qlib/contrib/model/__init__.py b/code/tac-qlib/tac_qlib/contrib/model/__init__.py new file mode 100644 index 0000000..b151d13 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/model/__init__.py @@ -0,0 +1,4 @@ +from .rank_ensemble import RankICEnsembleLGBModel # noqa: F401 +from .rank_gbdt import RankICLGBModel, rankic_feval # noqa: F401 + +__all__ = ["RankICLGBModel", "rankic_feval", "RankICEnsembleLGBModel"] diff --git a/code/tac-qlib/tac_qlib/contrib/model/__pycache__/__init__.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/model/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..f4e9bf0a78d22e256cade0d794a6958aaa7c6721 GIT binary patch literal 319 zcmX@j%ge<81kW_qW}O1kk3k$5V1hC}YXBM38B!Qh7;_kM8KW2(L2RZRrd;MI=3JI2 z7Dk42h7{&Sj8UwWESjt@8G*_*8E=UNCFW&&I=ki-r{*T*r24o!`R1pj=4dkA;)BS* zL~ijE0aa!u$ET&1CFW={7qI{hC}IT>%s|3VlkFBSNJV@q)F`ma44;8o8Gae)Cl(awmn0_Z7UpCoff(^%Msj{$NfA&W7vv=U`1s7c%#!$c uy@JYH95z6)(wtPgA|9XtAiop~0ErLGjEszT84T|;7(e22Yh*71g&P1&QB*?! literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/contrib/model/__pycache__/rank_ensemble.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/model/__pycache__/rank_ensemble.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c35c166a8449c14fa3cd4972c363d3d9d58620ef GIT binary patch literal 9461 zcmb_idvFv-dY{?ZmsY!ag2cnX0K(D=Y4rdBmJOVESlGO5gYB|xZZ+DOl~&B|ikVpp zDXTl;qKaLPJ4BsbEWsD!N-8G@r;b$Rf%A`jRhPP}O8$@#4qJ0v#TDOO#s3M*7yGVK zmHfV*on1)?DoJGoJw4OiUw41cUw^az6bSe@T)%DltKr@2Iqq|MurHTY;pfkxa+Q<0 zK2GLkhswtteGZ=Lf-3Y0xI0y6+|}n|?_8=o?&}uiC-&aOh~Hul9V|8-2RZy*WYhd z4;JVg0fnTjsQvxos3u{s-C;pP(}<$A%F?q5E1p%Kx>>9B*t{n(4<%b zT`<4ZcQj@Ui7{D87_kVJFf>Lh8<8{#(>3@Os8$*mgrR^)kWOi!wkp~%>grbHD@NkkM;OC`k7 zn5KYO!bplSLl;Mr+Ucm89Q8GUMM*iOD&Y~y80r>5HGD>m4TK`e1UNSkViSjGmaxr- z{{A53bP}tw-+=Itq7h+Qt0lh$0wW1=-X zV|Amxbgje-+9FAp)(#<%3{_Pjx-!^Ai=t_Z+JJ>pZ*T%?V+?7E1g%IWRq@}>y$xxp zg99Pf0g?f={ad$^CO}B4BI$*d_C9x{MbwibSRYAgkey*^f+#1&MA8rk)MVr|Y{y5l zqqEr5Rp{`Y?d|bcqOFSz%OI^8NS+Ob#G{~{g09hPN!-yoGBl=>TDC@#(7HCUqf7tGp3hGT8 zQ>4g{xGxFHu#5Y-?l62-^T>u=%b9+xHDv1axp9(P%KE1jF?R z6jOOxR3&W?*2OkMd$3nNpct5y>_-tvmql4JB;QDE1k!^6gFT~?7H4dMbugAGSfh&8 zL0fc9(U_`Q$~T1lsB9OJAth*8Fxws(5#B=@2@zTb4Z(hBo6XJ1(L{4I_DEtsfZ27$ z!1kJA30WCYP(ZX6SrWJ}u}J2`8EC1NiWp5njAFvmhd^B}3}TY7y(_f0zrQp`Ob0om zu#Hw-io!;c;y@~<%Fqa62OIDb2x|&hCu0ZTh@^t|O9ryFk*<*svpgXTG{x*sgocuJ zlKWDVgXFxHXQUJ*4PaDY4ZbK0N*_{mSl6p?NwVdVq@qI-$E3I_sDRAW#t%|~;ph*aHTTvv*{*jpF{IfpPDlF}`XFpL{a>K3(7jIuv2E&p%G9=N!hnuWLVNhuCVo3wB$xYikTXwc|ws-DmdFrW_j?QhH7Ol$m z)7@g1&({M--{0TdY)O@D35rPCC?)Wo5$x}8>+iP*p$92+G^s%iWN2q95kX(lZX@Ya z5T=Qq72g)Xk75uIo4^s+3t~kav0jcvSqvs`N0YI@`31qIAa9mS3D_GmCjA(mYfiGD z@dB*m;8`V-f^q~Mrt?TlH%#yG5ep`vpSNJ^hfoht{gCCKj=X>u6^qwq)>+R>4JBeh zH>-47H-274O8mJRj(DDlWN97S20w7RL9v@$5YB?DLdy z_MhytYt~&ZmtE&weO}ow2V^&DWf|@)-{(8Q1uHbdx6)nSf{rCP0fb;Jyp?G>LyZhX z*iO%K3*<$}6ki~IGknD!TZ2|f*49iX)v{m-NbQ%b6}DqE$os*vEXWBVjI0f01=x*z#&R<$#(5&owd|!M1D~?McC+LGV7cSTZ z$ECSpuT7_P*UJs;$GT*PP2nwm!JNgnpysyw8kKe}E!d;bUYxe@)*jfB?!tO-pK~*ESB*N#_|e5N5g+w%jd7 zq+m;?H;nibGs0ohW3}r+x9QXsHLB58wUsDLcVP^#n`PlJi}K;HPD~@45zqaht4$gi zX`@3AVHm2gb4nZS2XL^B&D^FfJg0<4#HaVNM@?yT?{H_8WvV_%TTYcwacL5NFdO+i%>rX#7F)jEK-93E5*DN1_;WF!@i038`x1xU=(!}JIB{OO6IKj6uRn%QJmxEoI+XK z{o1<8Gw*JGci`Uo&YAU{w>oCmcg?QZF;ljicSt~#5g9ne@5jhnKbT8FPT4m0DB28 zCrFG~MwGSym#1e7-~{)7qKK@{(#sZciE_|!%#x(bGG+n5Q`?MjmVpOR3-a1p+;rlH zXh|v%whmC19LF0TC9s9PWTLST#W^ltUU_+Uc6YAhPWclM-`jd_`+MEjx_|iHDd$fD z9|UG=cTD?t+^?*8d)uUOb^OZs5BJ_U@Z-bpAO5qWvz49G-cAjkm5USu1RhA8;YRuS z-$F_kxC{IR=LPo#&xP6xD=w_OPzQ-Qpw>u|WtWxQ#T99x*j_LuP%K^-YfB~x3XhQ_ zbO4Lfo;8Zx2Z(YkKSmbF#l?8EE@O@qASS#HF6|xm7Z*~jWcetcU6HNLI@7MKTXw$d z%6hDjlJ*XQVz|hiG1vdoTFTO8!_~zF6f0Ri>$dfY@wXTRi(QI={bHR#A%k5q?YrK7 zD&1`z&DM=W5;oHn?lc~+D$gIq`GhD`*sf(5YlzJz82C7Mp2FByAQo7RM1cF~{xp69;n2c+>{4nzb7YT$I=faXnjCjt|Jv=?B+ z2QVVw$-7V%Y1bc1X^cGeMzMdzL}j`XZ(hXcs&vILdBP=Ex++%W!|=u>jd-*<1}Xqo zZ-8qP)e|-8nzU<}EP2V5uDZ_2UT9Urk_J4=4r1kX?ni(v_FmJ~&)^+_t4k;F|JzQE z?O^M|9e|6C4^vB5q$_2|<7bk~ek6S8vlG%hcd6CMjkz{+9q>1!2t%)NJpL}-wt^3X zM~tfgNPg)Da7irD$B>ZagzU3jDX&XV_q3vn5KfO1=BA8KA&FZ!g%D0TxiIAnKZ&)rCg(R?4We6D6zvPB_<1LG@B^0F_elzBC1I4 zp>$O7G_^DhNt74mEO9D0h}@4}q_s5fQ`*hqgpb&?DCLB13l9ty16Y;{DP?QZqNtI2 z5<*nuLRc{?DVwq&Ct|K#HU(n(aYP?T>PnbC0GNI%^)TQ74q~QHQ$`i~B%x@>7R6=q za_~3&{J*2}PT9hv3G9_mlXvwGFXPdfcjdDcvV+C!b+VA=hn)}W@SD6=ne{E-CJSU4 zXe;{K()g@9TW(v>gacljU+y51cEF35Wxd$|Jh)(aaOOvZ3Fm|>jWvafkS)!p;p?LU z4&Ux$S_o5F%q}PXTGG&ig4RY2?NoG7K@cnG)+icj>0(kcT?w4wb<;x$D~!{gpoyFbq{>WB z!b)+O{2;vh2{B zJt4dIRZOA3jp7@e4yxAAQ~^j#SM8lTaIa<04E|N^&Gh6e>vGi>|0p+ctLN_z{q>=l z4bSA8yC(Uo-YeeRx|yoReC>*>d#~)hv2*HqV9Lyad_(t*o_h`5GY#GOsugpBdrQS< zVq^Zv=KQ+ggEBnL`MBEE56V%S@o&uZPQIKwI9I{dug~;c>IJb$CHL~=-f90PhGLNT zDF=cs;3<9F6{205dt^)R~t)YgUR3H#=+9oOp!fB-0 zrca~Yj_|v9vS{XA1*tqsBdfh|CJ}UO3&`DE~JQ4fh z%52wy`?qj0eb5e}!}yvT?shKTLWjy)HhC7w} zrkwYh_RTcy%Q$~sUXd{_k7vifw>Kx{t}(^ZFVoxIo8JA;2! zhoCe5>P*k&gV}@ms`@#>8z^Z2lpw!$%Yy*kLFl!szpSKM4YzK6Zd>NSr6YNN%~xKo zdIQKjlDsSL4czmto$;>C_1@^a>uvw6rtV%%!%R)XY|WDy&sT1)eoby;?&N!~Uwi%f zsj0!4bzQS{J2K^8I|F5#zv8G^&sB(1y`KiQ-*()uTa)X2Z^yMAH=3pn{N(u$o}XR2 zbGB|*ruS=stLS>*=9)Y1)^~hRe{C=)-Kf6N^C$dN{hj&_9DV(BZqC1Ix_t95H#c9Y zoa86Fa$A16di`Wiz9IPI=iYyAO1f1u+pyiTg@0*y=9)Jr!V$tJUjFB;py%p)9pupq0O#(Y5K$~cUHanK;X9QMzLnYoWM0~o%eEL zWA=EaCzF~~^8TvJ$Fs+C2XA!U+J2k=*!z+9le(E*ho*Nv_lxT3j>DPbpZbqrtl0Q< zJ$5o3*#4jMZf;elUIXF(%=u*RM&V~W1HG;8pKo;aHhUgQX$cfJ+|ZFno#&7yLLciY zrLx{aR9Wuy2RtYgQe42e6GB>8F2w~e?)cDQM+V@;dMDCh6T*Vjlk+ClM3*eEvxF<{ zl%0gja1Jb%fPB;m`Q||FyJ=WK3LNBE+MQ0#zwd?PoWnJR~9PZO{a7=rUyNi z$6A2}nx4Wu?DT0mMr6}X$yv!TopD7<*Kjj${fF6RdcoMN?C5p4e4%GQ^u9A zTtWUNcP2AFTe)%CyYcs5mT}cbc%9<=U##ib@qwCnQFNkx;js?9MV+-LB=B57-u)jLIj{Ml9uXD+2|L8f0MLTUOAyU~E6S zjF*kbp%eH?0J7NSlmVpRjtrnLOesM6xk1?Uhr?%5l3Hl-gu`+&5)QMBo$1F%eEfm} zN&TUuY(*ayjL2m(fdG0~S;B~-s^Kuw3;2Ny;693icc;@Z`@V3H8b(okbDn$P@VUJI zvC8XepR42CHD5ZMuC{prwK-Pv3$AVOv(2G=+qV1@>+>sD=WFZdDi86lqx@&ptLI#- z29LI;an57a;cQyx%B;E%K6SpF)xUnl;dO1A<4{;3x>g&+5{LoO6xa_2Xyg~P26|gV zN@{)b4_VE+6>Z(d7=8E9NWdCp3-fg~|re;NgHVLbo9(aAf1hfF`e y`E#!9S6s)xavMMAYJbB$`73VQZ@8yFE8jfrYdD7}be88g<-Ywpd=cU;M*TMzHLx!L literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/contrib/model/__pycache__/rank_gbdt.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/model/__pycache__/rank_gbdt.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..98be60fe19fedb3d9cb21c3f91ef0e176bd93287 GIT binary patch literal 12288 zcma)Cdu$xXdEdSFi#$FhilXGzgA$LAlVn-8te2w)^%iZtVuO}^T;49pBk$$iJxXMc z#W^)vGI1g@b<1<@B2??v6r|KD&>{*Dpl(s1E?S^GpBi2^3?MXV3;#z~P8;~2_V>-+ z-ks=kla;tTH}lOmGvD{k_nP^qni>y>=RZ2WKl=T4j{7S4`+fM@qcySG{#pj(i27r7 z{dKJDjMm2*`Wsl;6>W?)^*5pHmOR6j{$|N9dEd14uaoK|AHM6Qnm4)r77KTZlWN}t zl-zLFYRs~K!-OqRFTUa6xbvruT}ViBR2UbPh#1!dZA8=rSyZAEJ*t*SszM?zgq4J< z_Na1LizMP=RJbI@$IcuVB5_rdMM+2u3DKv69?xY>j7y>-39m;ZgTlaoiF;sREkbQX z78H3{kyRBKg(%Pq9Whl71cj3c#j|!iF`66^1vN5^!Bxyq7G6(bt`TiQxFHK-SWAh~ z=!6g!75o}Zq%KF)(l{PWM-;*-P&;N)c1RuvnWK>*j8P=^?%?ico)dbG z2q-+aPZ(E&j1ZVN<5NK z1QY+5c#|<)Ww?qA0GkLDjEP9-tQoV0NQ4=Ijinym*_VjRJ2B@jP-iHjs9MNq2_@pu ziNjZvluYDGtPwUg5T+X*D>0nlg$|NG8g|WY34qNnM`-c1WYFwJLKz#1Lc%7*7)j2U z%oqw5kWd_8_=zJ-EDjWf%G{u(Rsg4>awsWkBL^S~;SlZQU^o#68wZ1id=Y;`!-JAm z>NR#lREE{^7&H!T=>xQ@N>Hp8oQs3&Ay8XBK%%_&iS~FZ7K+N^xQxDg_dMAYl2vR& zQA;Rj+_!u86Rp7@lY7vK_O-Enun8di67AYG)GqR#6ai8BAfs?%kmN5iE<;Yl3DrYX z$I6Ey!`Q@9QAbJgkeG^Uv@qzZ%d!S8GlI~DJ$w@CiP#H{3_&9agWxCq04h&13Rx6P zieo#Bh~tq&N)bBB%9eC34NJr$4xf_aP&MHLjJqmZh!}^STceJdf`J}OK!-8aOZyuV z3fYk4o5F~QP30LFK&=c4(gdUOI8;7H#y&U7n0`tY4r2zSTxoN&t&K%q!Gz4M$LfGh zkv$S-HV(RoDd;U|kua&?Zj$VHLP*8YrjTk%%HhZmsemFS4SfUF#AR8MrJ(0iGD!hm z_lj}II%`$~R^4_!qH4P9MW|jfIlBEyDjAjg0O%^Jq|+$ZQ28{hPE=OTzeJ5Dfl=l! z>9xk*H7j)=@hTDyF#!o%4Qt5WWWVv7f81vw`>pY2TbjFDEKgarQjbxa*fjSgaR{*J>wx=3ga7*02>qGq%lZ^9E=2vnQClaAtD_FSYl|1 zaE>IRlHnDFQ?V5JK~Qz*!jKV+jDFa@%mT7O(KmvQiEucTjKIhNm|Tfh3Ha5>^==qo zDs=6+9vB!13YTQKJ8{Eik``0N7}+hA^d(f#zyS1=B&xy7&>XTFq^$_Uu8kz5!#kn! zi8wofND!0BXasbE(qVHe5&%956dF@dNpfvE1_t8EU{t=TCKU2diY^~BoAE@vCmM-E zvqG86N$66d*d1QCOWH(I9yUZ{^@S@6?Xgd3gZF*I9y)~!)Fu3O?s z-B}z)x684lHlaIVYDMkYXLJWVsQ9p~+XfSfsO|t|)l^J(4MyVOL@KTYe7eJ6s%{0h zbQ@7j=i|CPp+JG^E_`4M<72u-9me0d?uyGd2BV4anC@gBbn>tQ+z2tQTjAWQwBu+! z=JUm#U7*=6*y|or|5RXxdsoqL?_!2~7mdRVcQ84jH-#8kLq)v0$B@%KO?gy(|9uTQ_HJX8HFUK4`t&`q1BW$AeAs%Y_2If_?rRT@e`NpR+1zuda_dfgye+rxLgw`J)i+)&0Q1$E*37Yj z+dpfc*_1g}M%eVwzaw+H(6A1m7vDHnXj(V>`m8d0GCR1?xUe-loNH=d+`4#i@zUas z`^WC94=ndj=`{Dh@KkyJ4t56}b+X4C4;ammc2t!7vX zLsxErDXM63e8ngYCMI(>PZ+Ewv$gyMR{K4xX$GcT73;1J?WlmRG&I$r?P@&5W*eQ6wD+L`tcp0q3NrY|#%HK_F(wUTWTZU(EZMXit3x+`k^6{D^z zQ_L)-Ytw$}Q=)Vonenk}7_kl#ege25eN*)nSX#`! zNGQgjYAFa+%jUTl8yYVsg|4bu5xZcAz+|ub14210}D%6i^>NEW8ApsZv}t z#3B*X@n`GCJmV^&shHZ$W(9N3fXqIzgpn12Wrh_T80d!0X9dERR1jMtzTc%HcLT2= zAc(_A01YE3F_s7qz#s<_HpU{Efj3Pl^Dt6zRru{ayKf@Af?Y=vVFYOa@->R{WSHP+ zNKZH!rV0BQr;`N3(n1w|wBf?OtsEu=`H;U#a=dgL&c5QtRQo z_ekahnef0&;O%4cXXehl7hY(8cQoG=SlsvFp*x3uWX*RS%6krHj<2{kW#xs-i`(<= zU2r?_YiJ?7_)Olt8{UY=f2(_@`|VWTvk8#E*EQ2M8(w(&p=U>x-pWO~0t8syUrzPjxlIi_&4f$cGW=_ra zW_#x@*Ue%wjYe5u_Z_l19pWxVBLvy2_;dK=kZH9Kvh=cVpN(f6Kp2V3b|73^*RbF z!DhHeHBHMkn{zdr^EE=|WWm?C?Aw&{ZOZ$e0&hGUv%L$?F7C^F_JHHw#@Wl+?MvRx zOSa7hrGcUJ6n#YgaRybt2U^QqKaZmYi^TU^C8y+?=7w;{@McZFjSc}M>zg=SaKkO7 zf0EZMmA{ftvH`{+)kw9H9VKT)o4;J95{?lZfWR0x`F=NeP_K9~jma0xY>J^y%h@ZI z4mS-(kul>DkxY}3+-6EP_$65aR{rtRK>Y`NW#VPVn&!R@3P^kea`8U@9iFEI1s5vG z;sh74^xc7Q=yo-Qfpu3X6p2T)P)K(Zt3#BO3)rB%pZ&zOwXk99wC&rq1`{g< zo^m_*4JP6}&kb`^RtPpfit4q`IHy|9a^C<;PjIvR^*@6-QqU?2A{S|Rx(2P7PCai7&%9)gcYMz zf(gmV!l6T7(W6QRM{?0umAJiTJ#cfQBsFUu<+rqwQwsW%DPHqvf>|-T)hw02 zrJ21f`KtFH-Cl+^OXHmK0-$%Azsq0Acdy2r^C;QsHIJ#XN4wTE;HOp!#74Wn+&Ew#$SX-EEf2-?XEeEFow?Xe0UW)!i+cj44mrI=Z*q&nzJeXDer7dF-XP z7H!Crk`n9F?aZMISd@OiE2r>6KF2t#)9uDVl0|oh6VVXVzwRPC5E_#wCJ*#d?#DPh z8&YssSF{jt`fz4}+2B7NRegaoo>Jn)5Jil?o#Dp=DX`o<5*BuNNwr16X!zb71|HF=j(Jr(5oo$@JMWR zTU?63@l!`cs9xs2I^o*O@}#UgapWLsnxcdNXeeG!M5W_mF&>t6J3E)qElK3IF3_US z&^d)}k0W8I>P}{h2jRZbp`z~65+TG*NLitqw)bfy<%S1!=Ly3)!$C%H6PE&B3F}p5Oo7{VT28m)gaA>)=w;U}61+Y+!C` zVQ_K%?XlTCvwWekX@39Q{`tdmhZl~1(%8PZ<-?9U9Ulho1RvBRANJt&PddN0($e}p z&pV!lJs<49y+7a5J$LjMgz`$h_3Bd7)klq+mKxg&Ep2m0SDM?h{fn&+o1ZDPY+Y{Y z%C&SY9?Q4vDt0)YZ#}WpbYj)PHMPti`R_vz@165zFD_WK zr+-}6zPgud+Ohf!SJyP{HEeFBg>!=`;4J?*2xkEkEayoDRBW;xjzjAZ|cR>vVssjR!FJn;k7AxCM0uX!Xm=%x*VYgTpqs{I_5{Ujde zD`k3DOFUJ@jO3G(%<+%~mYeudCD{aiSeVtosf^@$56L-lBXLUOO;^Z5`86M;Oqq{3G=Q5j#HuM*|2oBUpxyTl-dZ*0tXG9vtsLHB&nnz5qd3)NkZmz554)x z-^LKCUP~;XWu@4uvtDnIu%fCMGtf1en98C&-I`1!4aebkF;u{3n4Jlf6{=9aLlf{J zH8hCRc!o;%RgzA5nE))3sw4^goAjEYDrb?>Upn0}Oj1*b2~mgyf>|XZ&E&SVmZ=KR z_kdk(V+Lwj^wjaez%ns3@ZQdvTqAA6${;Y->y}!quDd6vrc6U%M8KWW>&ba4Q- zve2@r(6GL+?y1lG1XRcQn?I{ZY1+Bk#MQ0O`relo_x+}}kkmAkio|xmu}u)nE4ZIS%A%-$aQ}`SB^MP?Q0MXgYdRl-2ZzJg zKIIZ-5^&%I6z;SgSFK{I?v}CR>CP4O8d`O4g*w!o$a6-;upIChyX$R0u}+>rKTKus)RVh{k(qN z^wC0F$8uYDuC06V;(bfLZSQj1^SQR?f82I(&NXYDytn>>0 z#%2GOoPW#0&cy?H|FiJdJC84JU+z4f>qP#wZq@4S@;qwmDr|kOu%R9QKl7qu)kn&s z2Bn;9=d|rD&+Peyc4!oL%`N|oKkHkN|JBvaI2YirAh5hj6Q_tj!1sX!Y?VYt%HY!- zi7z(M?WR}Rm8MrKJP6oTO}TzOVC(Dqlu{dgpc7mF_X#1<_lp|yM%1f_hORSnqkKrM zjOKf&bPun|ZDr&|7p7Z~B4MptU_%?@@H+M-+gX$DQnm6Y=#LWsgl~jN`5n9h_P##l zFHq>~3%JVbF<8fjC%d<1A^{I2ut`Le0 zy+~Y`oG5ou(rb#zy09FLhC%^ec?Gp9N&nZ3q#bEHmYgJFsnq*dZ*o7eo%wH@#lG!x z*Cu<@>It6f48j5-Bq=m+SgqM@fBI2v*Qy;~*pJqot4=DpIH70NO(hSvzID}0r8>^j z^SB_=Ow}^~N1~MuUvj6d-WXOn%1<+r3aD9dD zODp08JBLwzK+9mqGOWa?PbH9?Qqo5I^QF>(Qi(xQ9wgAJ{53vA?V%K1J_cE_t08#E z5+I$RkcCso#3&A`wzJ6Eh~6E>^&BNm4X!9<+lL2ZhDTLTsy;}`0`{V`IFsiMU#bt{ zM2>Fb8To_tti5;*rFegou8t}cpy@TN)x5cC7y?KGJUdLglO}nFUOsv;7Ly@iBs7$| z>yRN*N0h%q1IdGW8Lw3<&-0&I+IibA`*@Du_Mcowj_de2*YH!W@u%Fnf9H1mjNAJ& xuKQ=)#z)?b+1Hjl-5KX&M>Fq#yn8c$_`iry=5GyBRfIr_3e{|`kIE%E>W literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/contrib/model/rank_ensemble.py b/code/tac-qlib/tac_qlib/contrib/model/rank_ensemble.py new file mode 100644 index 0000000..d3f051f --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/model/rank_ensemble.py @@ -0,0 +1,189 @@ +"""Seed-ensembled LightGBM that early-stops on cross-sectional RankIC. + +``RankICEnsembleLGBModel`` wraps ``RankICLGBModel`` (per-day RankIC feval + +``metric='None'`` + ``first_metric_only`` early stopping) over a seed ensemble: +one sub-model is trained per seed with identical hyper-parameters, and +predictions are averaged across seeds. This is the model class the +``tac-rd-rank-ensemble-isolated`` reference run wires into its workflow +(``module_path: tac_qlib.contrib.model.rank_ensemble``). + +The ensemble inherits the RankIC early-stopping behaviour of the single-seed +model (valid RankIC drives the stopping iteration) while the seed averaging +stabilizes the prediction against any single seed's early-stopping path. + +Training is parallelized: the seed sub-models train in a thread pool — +``lgb.train`` is C++ and releases the GIL, so concurrent seeds do not block on +the GIL (5 seeds ~40min/5 on this box). Measured on a 6-physical-core / 12 SMT +host: the seeds scale ~2x, not linearly — the runs are memory-bandwidth bound +and each Booster caps its threads at ``cores // workers`` so 5 concurrent +boosters don't oversubscribe; larger-core hosts scale better. The qlib data +pipeline is warmed once on the calling thread (fills the handler cache), and +each worker then prepares its **own** ``lgb.Dataset`` (independent handle, so +no concurrent ``construct()`` on a shared handle — LightGBM's ``Dataset`` is +not thread-safe to build). qlib's ``R`` recorder is also not thread-safe, so +the per-seed evaluation curves are logged on the calling thread after the pool +finishes. + +Wired into a workflow yaml like: + + model: + class: RankICEnsembleLGBModel + module_path: tac_qlib.contrib.model.rank_ensemble + kwargs: + loss: mse + learning_rate: 0.02 + num_leaves: 31 + n_estimators: 3000 + num_boost_round: 3000 + early_stopping_rounds: 200 + min_data_in_leaf: 20 + lambda_l2: 0.5 + colsample_bytree: 0.8 + subsample: 0.8 + subsample_freq: 1 + reg_alpha: 0.1 + reg_lambda: 1.0 + seeds: "42,7,2026,99,123" + parallel: 5 + +Any ``**kwargs`` other than ``seeds``/``parallel`` are forwarded unchanged to +every ``RankICLGBModel`` sub-model (same params, different ``seed``). +""" + +from __future__ import annotations + +import os +from concurrent.futures import ThreadPoolExecutor +from typing import List, Optional + +import pandas as pd + +from qlib.data.dataset import DatasetH +from qlib.data.dataset.handler import DataHandlerLP + +from tac_qlib.contrib.model.rank_gbdt import RankICLGBModel + +__all__ = ["RankICEnsembleLGBModel"] + + +class RankICEnsembleLGBModel(RankICLGBModel): + """Seed ensemble of RankIC-early-stopping LightGBM models. + + Parameters + ---------- + seeds : comma-separated integers, one sub-model per seed. + parallel : number of seeds to train concurrently. ``0`` (default) = auto + (all seeds, bounded by the available cores); ``1`` = sequential. + **kwargs : forwarded to every ``RankICLGBModel`` sub-model (model + hyper-parameters). ``seeds``/``parallel`` are consumed here and not + forwarded. + """ + + def __init__(self, seeds: str = "42", parallel: int = 0, **kwargs): + self.seeds = [int(s.strip()) for s in str(seeds).split(",") if s.strip()] + if not self.seeds: + raise ValueError("seeds must contain at least one integer") + self.parallel = int(parallel) + # drop seed/parallel handling from the base kwargs, keep everything else + self._model_kwargs = dict(kwargs) + super().__init__(**self._model_kwargs) + self._models: List[RankICLGBModel] = [] + + # --------------------------------------------------------------- helpers + @staticmethod + def _cores() -> int: + try: + return max(1, len(os.sched_getaffinity(0))) + except AttributeError: + return max(1, os.cpu_count() or 1) + + def _worker_count(self) -> int: + if self.parallel > 0: + return min(len(self.seeds), self.parallel) + return min(len(self.seeds), self._cores()) + + # ------------------------------------------------------------------ fit + def fit( + self, + dataset: DatasetH, + num_boost_round: Optional[int] = None, + early_stopping_rounds: Optional[int] = None, + verbose_eval: int = 20, + evals_result=None, + reweighter=None, + **kwargs, + ): + """Train one RankICLGBModel per seed and keep them for prediction. + + The qlib data pipeline is warmed once on this thread (handler cache), + then each seed sub-model trains in a parallel worker thread on its own + ``lgb.Dataset`` (LightGBM releases the GIL in ``lgb.train``). Evals + are logged on this thread after the pool (qlib's ``R`` is not + thread-safe). + """ + n_round = num_boost_round or self.num_boost_round + n_es = early_stopping_rounds or self.early_stopping_rounds + + if len(self.seeds) == 1: + m = RankICLGBModel(seed=self.seeds[0], **self._model_kwargs) + m.fit( + dataset, + num_boost_round=n_round, + early_stopping_rounds=n_es, + verbose_eval=verbose_eval, + evals_result=evals_result, + reweighter=reweighter, + **kwargs, + ) + self._models = [m] + return + + # Warm the qlib handler cache once on this thread so the workers' + # concurrent prepare() calls only hit cached frames (no first-write race). + proto = RankICLGBModel(seed=self.seeds[0], **self._model_kwargs) + proto._prepare_data(dataset, reweighter) + + workers = self._worker_count() + # Cap per-Booster threads so concurrent seeds don't oversubscribe + # (LightGBM's num_threads=0 uses ALL cores per Booster). + per_booster = max(1, self._cores() // workers) + + def fit_seed(seed): + m = RankICLGBModel(seed=seed, **self._model_kwargs) + if workers > 1 and "num_threads" not in m.params: + m.params["num_threads"] = per_booster + ds_l = m._prepare_data(dataset, reweighter) + booster, evals, names = m._train_from_datasets( + ds_l, + num_boost_round=n_round, + early_stopping_rounds=n_es, + verbose_eval=verbose_eval, + **kwargs, + ) + m.model = booster + return m, evals, names + + with ThreadPoolExecutor(max_workers=workers) as ex: + results = list(ex.map(fit_seed, self.seeds)) + + self._models = [m for m, _, _ in results] + + # Merge + log evals on the main thread (qlib's R is not thread-safe). + if evals_result is not None: + for m, evals, names in results: + for k in names: + for key, val in evals.get(k, {}).items(): + evals_result.setdefault(f"{k}.seed{m.params['seed']}", {})[key] = val + for m, evals, names in results: + self._log_evals(evals, names, prefix=f"seed{m.params['seed']}.") + + # -------------------------------------------------------------- predict + def predict(self, dataset: DatasetH, segment="test") -> pd.Series: + """Average the per-seed predictions over the given segment.""" + if not self._models: + raise ValueError("model is not fitted yet!") + preds = [m.predict(dataset, segment=segment) for m in self._models] + if len(preds) == 1: + return preds[0] + frame = pd.concat(preds, axis=1) + return frame.mean(axis=1) diff --git a/code/tac-qlib/tac_qlib/contrib/model/rank_gbdt.py b/code/tac-qlib/tac_qlib/contrib/model/rank_gbdt.py new file mode 100644 index 0000000..d03e661 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/model/rank_gbdt.py @@ -0,0 +1,238 @@ +"""LGBModel variant that early-stops on cross-sectional RankIC instead of l2. + +Standard qlib ``LGBModel`` early-stops on the regression loss (mse). For +cross-sectional alpha signals the quantity we actually care about is the per-day +rank correlation (Rank IC), which mse early-stopping does not optimize for. +Experiments on the 50-ETF lake (SP-5d 55-feature panel) show that early-stopping +on a custom RankIC feval lifts RankIC 0.047 -> 0.075 vs. the mse-stopped model. + +This class reuses ``LGBModel``'s data preparation but: + + - tags each ``lgb.Dataset`` with per-day query ``group`` sizes so a ranking + metric can be computed per trading day; + - injects a custom ``feval`` (mean per-day Spearman of pred vs label) into + ``lgb.train``; early stopping then selects the iteration that maximizes + RankIC on the valid set; + - forces ``metric='None'`` + ``first_metric_only=True`` so early-stopping + tracks RankIC only (not the regression loss). + +Wired into a workflow yaml like: + + model: + class: RankICLGBModel + module_path: tac_qlib.contrib.model.rank_gbdt + kwargs: + loss: mse + learning_rate: 0.03 + num_leaves: 31 + n_estimators: 500 + ... + +The rank feval is used for early-stopping selection only; the objective stays +the configured loss (default mse). Set ``rank_eval=False`` to fall back to the +plain LGBModel behaviour (early-stop on the loss). + +Generic: works for any cross-sectional panel whose qlib dataset index has a +``datetime`` level (each level value = one query group). The per-day groups are +derived automatically, so no universe-specific configuration is needed. +""" + +from __future__ import annotations + +from typing import List, Optional, Tuple + +import numpy as np +import pandas as pd +import lightgbm as lgb + +from qlib.data.dataset import DatasetH +from qlib.data.dataset.handler import DataHandlerLP +from qlib.contrib.model.gbdt import LGBModel +from qlib.workflow import R + +__all__ = ["RankICLGBModel", "rankic_feval"] + + +def _group_averaged_rank(values: np.ndarray, gid: np.ndarray, offs: np.ndarray) -> np.ndarray: + """Averaged (tie-corrected) rank of ``values`` within each group, vectorized. + + ``gid`` maps each row to its group id; ``offs`` holds the cumulative row + offsets so that group ``i`` occupies rows ``[offs[i], offs[i+1])``. Returns + the same result as ``pandas.Series.rank(method='average')`` applied per + group, but in one pass (``np.lexsort`` is the only non-linear step). + """ + n = len(values) + order = np.lexsort((values, gid)) + ord_rank = np.empty(n, dtype=np.float64) + ord_rank[order] = np.arange(n, dtype=np.float64) - offs[gid[order]] + 1.0 + sg = gid[order] + sv = values[order] + newblock = np.empty(n, dtype=bool) + newblock[0] = True + newblock[1:] = (sg[1:] != sg[:-1]) | (sv[1:] != sv[:-1]) + blockid = np.cumsum(newblock) - 1 + block_mean = np.bincount(blockid, weights=ord_rank[order]) / np.bincount(blockid) + out = np.empty(n) + out[order] = block_mean[blockid] + return out + + +def _per_day_spearman(preds: np.ndarray, labels: np.ndarray, group: np.ndarray) -> float: + """Mean per-day Spearman rank correlation of preds vs labels. + + ``group`` holds the number of rows of each trading day (query group), in + order. Days with <3 valid rows or a constant pred/label are skipped. + + Vectorized: per-day Spearman == Pearson of the per-day rank transforms, + and the Pearson moments (``sum``, ``sum`` of products/squares) aggregate + over each day with ``np.bincount``. Runs ~10x faster than the per-day + ``pd.Series.rank()`` loop that preceded it — this feval is invoked on the + train and valid panels every boosting round, per seed. + """ + if group is None or len(group) == 0: + return 0.0 + offs = np.concatenate([[0], np.cumsum(group.astype(int))]) + gid = np.repeat(np.arange(len(group)), group.astype(int)) + rp = _group_averaged_rank(preds, gid, offs) + rl = _group_averaged_rank(labels, gid, offs) + n_g = group.astype(float) + s_p = np.bincount(gid, weights=rp) + s_l = np.bincount(gid, weights=rl) + s_pl = np.bincount(gid, weights=rp * rl) + s_pp = np.bincount(gid, weights=rp * rp) + s_ll = np.bincount(gid, weights=rl * rl) + cov = n_g * s_pl - s_p * s_l + var_p = n_g * s_pp - s_p ** 2 + var_l = n_g * s_ll - s_l ** 2 + denom = np.sqrt(var_p * var_l) + valid = (n_g >= 3) & (denom > 0) + corr = np.where(valid, cov / np.where(denom == 0, 1, denom), 0.0) + return float(corr[valid].mean()) if valid.any() else 0.0 + + +def rankic_feval(preds, dataset): + """LightGBM feval: mean RankIC (higher is better in lgb convention).""" + labels = dataset.get_label() + group = dataset.get_group() + ric = _per_day_spearman(preds, labels, group) + return "rankic", ric, True # (name, value, higher_is_better) + + +class RankICLGBModel(LGBModel): + """LGBModel that early-stops on per-day RankIC via a custom feval.""" + + def __init__(self, rank_eval: bool = True, **kwargs): + super().__init__(**kwargs) + self.rank_eval = rank_eval + + def _prepare_data(self, dataset: DatasetH, reweighter=None) -> List[Tuple[lgb.Dataset, str]]: + ds_l = [] + assert "train" in dataset.segments + for key in ["train", "valid"]: + if key in dataset.segments: + df = dataset.prepare(key, col_set=["feature", "label"], data_key=DataHandlerLP.DK_L) + if df.empty: + raise ValueError("Empty data from dataset, please check your dataset config.") + x, y = df["feature"], df["label"] + if y.values.ndim == 2 and y.values.shape[1] == 1: + y = np.squeeze(y.values) + else: + raise ValueError("LightGBM doesn't support multi-label training") + + if reweighter is None: + w = None + elif hasattr(reweighter, "reweight"): + w = reweighter.reweight(df) + else: + raise ValueError("Unsupported reweighter type.") + + # per-day query groups: each trading day is one group + if self.rank_eval and isinstance(df.index, pd.MultiIndex) and "datetime" in df.index.names: + group = df.groupby(level="datetime").size().to_numpy(dtype=np.int32) + else: + group = None + + d = lgb.Dataset(x.values, label=y, weight=w, group=group, free_raw_data=False) + ds_l.append((d, key)) + return ds_l + + def _train_from_datasets( + self, + ds_l: List[Tuple[lgb.Dataset, str]], + num_boost_round: Optional[int] = None, + early_stopping_rounds: Optional[int] = None, + verbose_eval: int = 20, + evals_result=None, + **kwargs, + ) -> Tuple[lgb.Booster, dict, List[str]]: + """Train a Booster from already-prepared ``lgb.Dataset`` objects. + + Pure training — no ``R.log_metrics`` — so it can be called from worker + threads (qlib's ``R`` recorder is not thread-safe; the caller decides + when/where to log). Returns ``(booster, evals_result, segment_names)``. + """ + if evals_result is None: + evals_result = {} + ds, names = list(zip(*ds_l)) + + callbacks = [ + lgb.early_stopping( + self.early_stopping_rounds if early_stopping_rounds is None else early_stopping_rounds + ), + lgb.log_evaluation(period=verbose_eval), + lgb.record_evaluation(evals_result), + ] + if self.rank_eval: + # early-stopping must be driven ONLY by the RankIC feval, not l2. + # metric='None' suppresses the default l2 metric; first_metric_only + # makes early_stopping track the single remaining (rankic) metric. + self.params["metric"] = "None" + self.params["first_metric_only"] = True + feval = rankic_feval + else: + self.params.pop("metric", None) + self.params.pop("first_metric_only", None) + feval = None + + booster = lgb.train( + self.params, + ds[0], + num_boost_round=self.num_boost_round if num_boost_round is None else num_boost_round, + valid_sets=ds, + valid_names=names, + feval=feval, + callbacks=callbacks, + **kwargs, + ) + return booster, evals_result, list(names) + + def _log_evals(self, evals_result, names: List[str], prefix: str = "") -> None: + """Log recorded evaluation curves to qlib's active recorder.""" + for k in names: + for key, val in evals_result.get(k, {}).items(): + name = f"{prefix}{key}.{k}" + for epoch, m in enumerate(val): + R.log_metrics(**{name.replace("@", "_"): m}, step=epoch) + + def fit( + self, + dataset: DatasetH, + num_boost_round: Optional[int] = None, + early_stopping_rounds: Optional[int] = None, + verbose_eval: int = 20, + evals_result=None, + reweighter=None, + **kwargs, + ): + if evals_result is None: + evals_result = {} + ds_l = self._prepare_data(dataset, reweighter) + self.model, evals_result, names = self._train_from_datasets( + ds_l, + num_boost_round=num_boost_round, + early_stopping_rounds=early_stopping_rounds, + verbose_eval=verbose_eval, + evals_result=evals_result, + **kwargs, + ) + self._log_evals(evals_result, names) diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/__init__.py b/code/tac-qlib/tac_qlib/contrib/strategy/__init__.py new file mode 100644 index 0000000..184f80d --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/__init__.py @@ -0,0 +1,4 @@ +from .optimal_stop import OptimalStopControl # noqa: F401 +from .long_short import LongShortTopkStrategy # noqa: F401 + +__all__ = ["OptimalStopControl", "LongShortTopkStrategy"] diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/__pycache__/__init__.cpython-312.pyc b/code/tac-qlib/tac_qlib/contrib/strategy/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..69684c9c9624cb16c63f3a18a864c49a7db37fa9 GIT binary patch literal 301 zcmX@j%ge<81kW_qW^Dq}k3k$5V1hC}s{k3(8B!Qh7;_kM8KW2(L2RZRrd;MIW+0n6 zg(aOSilvfOlkFuVP^l*4Eg}DclFZ!1oZyoD0_Xg^lA`<^ps1)%eqMTTMt)IANPaIU&i{01qJ#giOIT!IhjcyMm(6246+$0Pz<$6KR!M)FS8^*Uaz3?7Kcr4eoARh rs$CH`&E=~#5BL*zH2->;`Jo8G>Nmqe$P92lD61sAiP+w=kA_d}=eVEIjUH@T z!gPNMiL0Ez^>czrFvm>^bHCX{X-nMFZ^6|Xw5&cQ^}#8^!b#O;ZLSUdyeqwj1Dl%Kus<&TM(o}_q@XSZoS0tQo~Vv_0}cc#U7 zJO;7@;&=*NN5>>dOlJ7=skng0>2Chw`543@f%}*QHX<6Cp2$U_BVLN0k0gg6rZHrS zq@8F(BpQA>9vkS5rji*6*MUfMBm?$(A$>u_TfI>4NqVsFz`-IBG~{9|a~?DzakNNl z-R_gpn20x>11U*LUF74jL@Wcg5)o-Amel!5fM@WA1{AXmio~lVI>jV9JBqeKW|G*= zCsSw~QXzEoiP!`+T@Xj(sqr)y%sdyNJJUkjjac9|}0$V~%Kuzn=MkxYb(J_)3iN_|u zB8q`4#y={dX~CoKNfsc+4@wa>`JAaiJd0##?t;vcX=laXmjx5`BWN{=OXIXc?2nR!dQa!Y-YoNh;o=Dne&h^ zO`=_f;3RP=P*$OVWnQtWb+5aWZypfMSsv`~dppNV|uxpUgYp-~3*j85h^ z@S~#CQ|!xVWHilRh(*w!CzsD%8g8dC7zK0%4^7lP+m>)}1lhM_rWYEoE z1zNLDN{yz*G9|h@@#1Mo5G7Qx0e!&IeKEQnlu`-2WlG9WPJ$SXr7^kBo$JOsCh>4I z5*L#~L?@v=)MZ9=jhv&Ybc|F39>$WxuqhCOCSzW&0lHF*j%89(uk%wfRUwn=iX@>? zFtSt0^s$gdwf12~s@9{JT&n%)QBBt%pN%meGUFk&2(3`vO#XdErcXbOTwW$vX2oLI

D4CvlssyxH{Prv*3Ei=iq+Ze}Z?48#D=3U#sc23O=FwYu0|7;1_Ce ztq=l&8`nxa_uy(5f18cPs;GNO1)!#ksF5)c4^(^{WYQzCQRr|T z2Jz5*@i>G79ZdwFk95sV@?u=B(_~AM$-uyqVbo1-XJ$-FvVLN~A{8&nfoW zNp09rw3keeYT_hw*~_0M6L;>MZdRF^O=U8vL=kJTKB3x~^av=;e z(69`4kRTWpk`Q$>NeDLNSSm@@V>h^r7fq&4yx3bBTLc3x#*zX|C7`cvN{dfAX(W<) zqcI#}OzA+N2L*Y|2Lew}n&mAmLl&SJL+iv$HwKl~`^G5#5_n;>G%*GlF`xlF0kJUK zs|{M%!zh_+2n<4e7~n7{b&9=1z38xwz1x^**yJJsCTMOJk&VcokHAD1=P2V@6kz-g z7_-lyM<@;bz$e5=dQ4(7mZU;Gkpv(u_TbWC0Fd$<$(z_SLX&(5}Ag1?YiN#_+E`zN9 z?eaJIST8hk8QVprO#srMS46f&xuAJa!nQF3DKw)NOEY9g};A1#bH6jj`$~v#o0N66>f|8wQQ693Wh2 zsj60$D`nNqnpRT5sz-Z3no3i3YHjHWIG|OLo~bks)M_nLaVn`k?N%2s8&P^G{b3@a zy2_TYY7MDW(O7t^t~fMBJRD2H0IOBeR3ec|YIji*8wDT!9Qbo_oaw&P!g1d+UF9Q>N&hK!)XV&xot|`-` zzbjEeJvGKt*Q;xUin6*~dfBqNDp#z_F1wc2CA%PB4)QH=Zcsy=ZoSSWZL3pbL+XvWVABRF4?`5 zXW3oksnjn9HSkni=b-;-Gz4_MVssZ=a%aeRoC{SS3)v)MS8AiEouV}qQNvD3q<&)m z69ajot2Doh;6}G!&|RGyyZKEUH*P%J4c}gd|2vdop_o{Fd?HjTUQ`eK^76&PdQBy@ z@j%W}O%nB(UZSKjmXz4Egtzy(P`zqPj{*9q_AoR>EE5i^R(Q1H5={n)WLCQLdr z3b(8n9udcv+yelp)!||*Se8&wwUWa_>!F8nrtytA%`x^?Vxz>tO8u->jKBptlvW+# za1^F99S*0d@i0RVhRa`U+W^1c1~^B1$nQW2<|W-w^pb8siH3E_H$8 z#Ok9fE!cbYj(7%n(m#M={pDru^NpON=1QevuV3UU9nONUp-|smXbMr>(D9iI&lamV zci`vMl=5&@by?TyLR)XP?U8G1e$aQTIrrGHsU!EATIO82rtT^0U3-1OUq5Ai(^05f zjZ0ObX|48j)q1)%7WC`(L)q=e-$~Csez#>y_W2j-Wplyp`R2B4Ln!B4f7iWvUO;ru z+njBADtqLG>{BmhSDnpyU&>a!bdPV(uGyO7w^6B81y^&{)w$sN{#bVRvE25j?{0c7 zd*V67tn8E7!~NMOU&yX{G3Py-tvXu>?aYSuU-$o|Eqg4Sd-5ESSp8d8+_5>^vGdx= zKegRT}K4?AonedpAy5KejjCM<;So!`|VBtR0ky>HZ}ou0kUX zH0#?XC%ASlV_4EPoMbU^XSi|GMedU240q8K zvK%}4zD2dCGpXoEI3}o0IEbZ87)vlwwZl8cQX#80oNp*wHT27S?)sl|CP)2GP{4xg z7in@JeSK9A_eYLLEjLY%T9)(-$&2;OO(d=YCav5z!^pEdX!6e({XA^QzLfJr(+9=( zlUCXKs#8WhvK`H{$yUKk*g!IkbD3(rglyIABkN5A`(zAsGiWw?mMiO7bjD*KBJHMV z4CB3SDw%OpYsiU+y_wv-bL5e(vCE8{S!;b$2Wq1q4}sw`X4^_rpr}z6Q$=>rIHG<=sjUnGJ;n zwC!P&^m5J06ZEn<#+~3wKsaN7*kPj+LBAMe17_;oJqb`eY;>}o4*=FO(h0Fw;Vq5h zW#I%D<0h>R4sKdYyfR^*v;oronl(>Wl!0k(*nm%^Dp02ZoJvnWD5@vHIsrKewh6qN ztdwoT29z%4Edzcgy`@6MkV&p68Sjk+q+q8prq?l9jox{v%vOFUSAzx!yEu~TgbNa8G$(81n&Dq9rOk$v+%o>r zY)_lsESFemk@MR~%)=zM6%w<^)+jB+NprJ+v|+MjD>FEhnFSb`9$fYUw?b{^3H_<;F>5ez06gTP*F6W$=EGf-M$!m|ekitQ{T|Da1) zt4QGoB`jLy0{EN*0NE8j{k}Yi|+T$WmwA@OQNA;cKlefxW|ueh|Vv}JNYkypR?;qBr53`-CN*u)E36nBGo=9M&OJW)9{;O^H?lW z5yW^VqB^mAj`q-jC)I|n9r2Rt{0f}cvBNA~IYuv|BdQY(3u}jl)JiGE4q(uFUr?>F zcq*z|X_}}NwC$FRfKouXjmQ*d+omU+=J%nv9c%>z1P);KAKt2lc z^TS6`i8Xo>n@Mo-VhIwe)m9%inxSdcMGFj^$j}xcGtNo+VRnjyQfa)PYo!i2JY^SPVb+- zG;=iX?^OJq^K#bTdDs6~Ha3!tXO!63)IqkC>rmVsS@)(wuqhv0s|45Pf}Q!`MkTm$ z!JprBNZE8K7knZgJfZ}THzd;hCB|DI_RSiCwsGd$;=NzRSE zH*tMpVc$*D_YN*}D(iMB!Cf~aO7L+o)zbFv__gtW`O3@_sJ*pNS3kS&?F;uF@4NNv zzrQqVdAH_T&5iW@*6)rlwB_ozEr@sOA6v9>eAi+Pw{`nXOV$pL0Xzt^C-1J>Sa5r% zoWK0s#d#Wj&RHBa_o`OS3hySbCG)GdDXX{TsvgN#?NO@s6uey@_AS(Y@I=1*F{S&l zY}f9a$8VW#rnCE>&w9@M!eXwj`HVyKzZWf_@L5YIu-g3Oa2;CrD^7 zvxIhy1BgAEd!1p%GQ3+nOrvd5l=cyTOWD~A_FRKeK)X1#ONZjv=?CMCqf@t)#R^v1 ziZ;OD!*-%^avLz69lSK?J^Ya5;2^g4`v?sZm;;(&SDY4Q=b{K$77=C->?56I=LRBVrCyy*KHq1E1o! zNR2!c_!v2UM0CE1(Fx#F!+;O3xOi*@khg9HtV8akDR8_IgtFGV%q-?EU&Z$GMRKYIIRW&1OPC$^u; zZF%l)>uIIw^smG^*k7sjP4Aegn%Z~I-7wqw?z(I1Zft&U_x0U(0b=}tt4C&z%vI(5 zJ-`cI-__7eXtpokyh&-^l=E)Rd$%jz?K$tGdG8*@yXU5)c=v1gD>xIJHP19o?Y~zW zm>!tjGIRdz)y%iq-Z5p(+Zz>o8h8Je83o_N0tQvd(;2iELaE4%(B^J+p9HYPFTayffEX_N>WuOpc5E^ zymjg)I8d+C-;}O~oI!WEzVuvAJ)GV{WJ`O!%YsoyMTYiySs11j<+1>Z{IW;5)$QbD zCEV6_x$;*b0yq$%7bhKMCk!fAa+;RTff&@C8N-2EI=La(M*5g*ts6MOE;|STyt&!R zmDzjSaLH1}11lB{s2Uf{-cyC~Z zQYP7Pg}>6G@AAPpxJj-WHXN$u`E*`P+?xli9QdgQ#Yajx^whLT^5Ko)D%GEJ+!bD@ zfjlPUG2Er4JY^CyT(l@xoNtpg8QOzlp@erJmBKQq80Se)%O1T%>1s&jKgc+}asm5f zZ?+-luOpvFkZ*5^lAe0i4N0!}<^tOK=Cg3s9?R4Z*S zc{7J~Zc1mq%AOgkf(dLql|485m!6~lp)V#${zr6brR&lWKIxNvmrZXr$UdAz)8=Ul z+}D-E1lg7aLm&CEN+2M&Ea>l1KXfUx1rBU`IAreyaf^p1h?3q=8;vjlWew&D5$XZff;Z%M?Fo!)KmVlMP?mU+E?NRnIhlOZKib!%A}$ z<7O2`%83VCj`mex)m+q4kPbooWuZ6&_n0}kL3XTUMPw%=Pcz7qc}_1p{JfqjU6byz z_vBby7IbY6O+5y=)_=*^{R2Z5jgZ@lS_-RkT?04CjdBxt6(*a@+R$wDD#m~@5^15r z0whA~m1RM$W1fp@opw<|hvnv?92zuP$<4oZdxzwvV%gUJ=bm~?uieF!)xSc2l*#)& z*n9(5U0Ho)70sz~bQOgyrrCNCv7QRKkFkXzdH2)jV}luOg-845**Lo*t&`1V+v*sJ z%LBu=<}0?O=}P4=bnfy{ItOFo#_=y?Jth$pk={atQz?Yr)QU(VHHKq{gg4lk7m4hq ze$d592VDm7F#t{{HcH4Sgth$P0qqMZF+7IDj}oTcB4n?Ktn>^`BjX1u9NeUsS zwn;Eebr5;AqDa$}Qk?^1<64=twi)4~Pq2$mREM@D0U{DD#M2}OwyeHPL0eQ;q{jx} zPMO%f@=n4KKJBKngzWf~h+{k2hmF0DVS5Fi6YuHOiXX;UO1}Z}zr=~DKtn#zsRTOn zfh|g4%gvLwN3ud98`zQyB&QCQy*;1=4%|9)`%o@$0?*xnS^w)_EqH?WJb~HF{Gpq# zWKX@69Zui%WTy5p*L|DfZp*uS6?gCa%e4A!>HP5EzdHQ>;X5tcW*Tqq%6pC}o@2LP z{@C*@-0Yss5}`}l{lYGwv<_(rOBya@*ThMiR}zycxzm_nnHxl{&^{7Tt@Qm-Xrq4I zUH^PjJHDjv^@yWtHT5#>jj_XP?3}2|dS3e+Ky#~*l2aYzDLO>o{b_ELm#JbJ#!`Y_ z+r-3ZW?Z$iQG>6U=~&w+?gddCWvSx9d;RlJjT7x~9DP5fZESfocIbeeHrP}PwkK69 zlD3;NVQA!5fJvQwd1(79FzVTMSxP;=~}6W zQb!PJj7KBs^HP$Ywb3(oiJK^!DRi7i8}QZ>2ag_QBV9U24{5y^YSl(@jkZB6;wWw! z+kUFEOx{lVfrR2Vo1LXGzOfU?p3)B{sMU0yFbp>#>lS>Hp#2O%7`uO?suOob$*_R{ z@gy?nboT2A*f_=ESlW`L4_edgNFgp}2$%Ngv<*5N9U@CJZh7^ZOm34R*{R7++ECltYgYGXT{(p4fmOVKro*fecL>KvP=4~h07XtTHO zA@g&h=Qo5O z4RqxKr>30uT!HD~8*T4(U+=!_+E8d{f4AjY%OBRx@0<544ZTxO6${Ztu6+Gg zrG9I!{*kHbdv#4$Cub(-PUbsyC>=X;b&pO}Emm-y>p{fN=lvTL|Ars?H%;~3^VZE) zeDk%-{_kUe)s4?7d8*44&JNh=Yl`3@0_;Y^983byuE8~4GH>=ulFbZ2Dm_Hx6SM> z_?ruLP5HV`rLJ?{KmW?j!*}Zr(L*}nx=!(}%lfw5wBG8=_dT!R|BmPBU7+#mnVB=Q znYY8!W|lFe_(ECVwnBTy+|FxT7dHRNBj0=E=K9>0gVTMWJ}u8>XsfwT>Fv9-uCGwf zvqBw;uOsW*RA^f>xBl9ig|%2x}k`Dxp<@C_Gr+v5Rq z4o~$>x1(>S$6lWV3$yKWwmIRtOKI#w(#Q^R4q^f4Ke|k-2~p*p{n(?^Y)N#A)S; z(|6pbk>Mv0+tzn{vnRe&zu;SV`Cr%F*!JGe-{1KI@ju4?ZS3Q!Cl}o$UDUcR-_Wfz zbpN=ax7Z=x=BvFky_hO#oK%268-t*I<}k)QO|02%bN*|)a6CKbUiV?Ue#+t2cBS{J z+Xr&%pSk0HMq_vS<=4l5f-?x2*;6+TzIXWg;g3DNpZJ2awcpr2E9YBwDJ{Eh_T5Tk zTXyAIPG(QPkUu@7oF2NOb`FP; zB}@K=+sMPdt2wKEJ;UC!lRke=)Od;*Y$5|SQEimLhc{=nk6H*yh0M}Bl#@1z^fZ;n z`S1(4OgIcSSiXrS`ke}5tD7&SG&atLOvgelwzZ@>!}g zW}mlcScdQjLogbaxs6+V*Q0$|6D5Eg4w=}tRT{XJ(`+gGnFkgTq)Go1qF;WE`^?PS z?2AvCIFBE@EjV3KXlgF3ZiCe;tX*GN*HzfCrNFlpdNxrU>M68$7S{I`R^c2*6Vx(n z-)CJrZJQSDoZGwTpcM2?V9`aXDlXWtSWT%KuD*HEO(_poziP3TQeG~&X3<9}KUdeW z7@$;;YwTI9qf|ZDw0g0DQjJ`5*J2Z;n(5gpO0DMV*DbbCisx3hF1Avtjce*$Y^T&3 zF4*)B9Z21uFjd&SEsaHePKy})^MC%#V&hp~dNZ4kiS{V*w=NWv`K8X~fh}HobT(Uo={dmb9 zDUWB=_i1wc;xi4C=`(Ya$@&jFIaB-3IM?5B>lAL?&p6+I=UV=Td+euN-%q)l*qVz*l;NdMOJxAQ;gps1f?K}$xJpzmJ=;s(W1 zA&RA0U4mZJg>*F1`h4mUbKWPi`I~plxq{VMSI8& zv?)f19IQFyWGx{VYYo-0wvd~(hdit!QyO1zGXwu6?w=&Q4?t~ zz8FdL3t}prjwfTxAOmD{E+UBWXn!=76!}yFs`!Y=#g>`b6wgHYlpypATvUvwk`ZW( z$KW9hSS)9u_c9X~mHB$mi5=TZq4%@-Fe$rO`FC1Ws=OE9@PGEF`%EClZYfYDb?lRJ>U=( zL3+50OW?zoxOi+%EWf;uSxP423t;^@&N2nY9N|;x)Dm=v2y^|Q<5W~t;vBaG`T>83 zmp&8c;Z^aZm|`ySOG#AW!YmjlnBZBCpPnAHl*o}d zl$cpQwFr$kd|)OLT@X1z9ALR3UwARWtvqmcxcWFt5RS*VrVWmw+eA6QxL zsTl|=ng+n|M+Wg@cn28O%LMVG`BKVbP7aU@79)<_a;LbNDA(J%&T90$Tk}hXtP$_B|w^l)pthqq8I_KZ2na2jn^@Kd{V8U?w$M? z0rz7aWU;Nbul&asx4!d(DCu3w?CO)gzY4@x0e3knL<9EfSURL|%2w8$(a9sw{QxP8U z4OX~7iv8t8iil{;Q3jSilU!Py0d|0o;HJ?fo(B{yFDO7cWL5@N91~A|_RZkv00$Ll z=3yO@^#iMC;;FOaOway-fdj!3w-Sy4vj@--HUr?I;UhsAVJG8b;Cm_xs=OWmFJu_u z0nW<+)}!KAe+j=Bg1IFz9>5uOV)4bLMMeWqghN0m*yvQD`U#)`f^&$qunV%P2pSdv zi_Sj;jnq*>0{Y`zRkq?B&k6Wu1~(r*Lbg;i3Q;|i=HpS0={a_s*$)`Z#)ZgC0?}!f z*{`x2=WzkVW=Z<5lR=wOqi&dr0d1|A%O+Gzng$j71KNXPHS#JZZDNW}NGMj^IKnI! z<2g=oJcu8PHNm~eCBpF}NUqqTsl~-qQY}U4(rU25(**ym6s0c0zo2hWqSi`Xb4b)n za0+2{5-nLIt7K#8Hw}{goAlRdnv!%96~@I)lB9={zFaau2}Y_(GD4|z#G)vv21+&M zwpu9FmP;lmp+qDtQ$hWB(7+>3@b#GaF>8SA9V%#6jKUJQw_*yzw7~zuyaSt^n7J^k z!_1AD2Qx2ZL8oF6xWuf=^Y{1XeK_*LcPljumm+*jP|V?QG!YSma9F^X;-ajiE5AK5 z7)hrG#YnUtw*n->gdp25RA|sKSURVK(enn<%X|wkTWS6*)DCD9eHc0jGm!n}3iat8 z%IvsSE1Md&sT#AT=xHoA2a63I#m25r%uumyp`5;7T9I;4zQ(NKy1Cf62Of@M0|O6x zv2|b8aL3eCewM1;-mKxJ+WYVps_ZaWKqJF{`hGVMHzfSBY1D5K? zr(j3X!zmE=Aw28hHbj6#z+Mc`Kn@Cp9${==Y1wW1TI&VqETB!VQ1_hPmshT>tg+dZ zyt5-`>fqajp#TEY^YGl6=F7fYX<=z?3K^ZGyHW=u%U-;DFR|ju)zjCF zD@KAAL4#t__HsoZ<3z;;85%yi5RNRSmXe}kgab`P4C?t-cn9b#CGb5!2`C`Am#QgS zxP%>uK`VjXlvUN?1ss1EGSHXHbHnzEt>9v07qjlqyMkF`(dx<`FM9n2Z;$Nl+3@GR z2eN1HI=!pszb55O&8kML^;6d92=M$o7mzC2(@T2Z2neBHI0CDQUZw#%^eccmtZoG$ zk~qk}11pkB$A|5cL7E?c8UO&)1m6#tVpxpu3yP7$Kol+wU zV+3S%g(|wd1y{T5YG3cjySlT+kDT5$NbtpH46`jX)${VvV`(j?VEi6xq%Iw`&#;PsE>2kWh5qNwfy zaJ)oa)?c75(XjU30k2bCEY2zxL5%QX7`J2Cb&{+~nOQg)kp1l)J^v7h#Dk?M9)Sit z5%Eu{iwh-3#f=8{I>kJ-oaWBr@i}Nz-PI6Hu^g;ye21b3_WUC-RzS_6rOMt5yfiMm zdv8?-)GI-|E7Z2mXr7~s9)H2pEql5Po`+@6!^OtdjjpYZ4;r`il=qSQl+NqergV14 zZ$7P|T#dh^=+6$h)tQ@oHW!`Ec^=MtVz+IvZ9Vq-Ey3|WYip<8-aBm7zi*)-e*jBx zZz5(F!F!44L&eUto6)rcC?Lu1hOKt45<{tsAI(yz&(t zYnatFLQAGrs{OJNP%$cF7x!x8wfAHkAd6G7&Ih%Y%2VqzLaov- zgeZ#mi7n#-NqQ^2G)l5b=9{&e1u}J@gFdaT@|+*6&>96J0r15AT6um*qbkpvCJhcV zZpr;+qXfCSUtFJ%O!K&X?MQGQHJ7FFNGK0N*N(LCJO!V0{C)zd2xvPJ_!)RP^azk1 zkM3xOqU4g?+I!#~yYN_O0!w;{1Za#M2_+EfBoB-|v7-r!_;x`jdHL7=cu&aq!~+^f z(6R3i(6Z&-5?W;HKgWK<8pq00s@HsvACT0S?ydmkVYTpHWGkdCWq)LYNvJBvH(w>|DOY{J5 zd|ld8CCe_dl3iaXwI%8QPo%m2Gw_r~Smi$gJvM>I?1oso?PPzz+N;{hI;4&-TO`PF zClicP2iObt3Rv66!R8&rik+YHcKmAFk<^uc)2SWxGhLrEHoRjX6s1n7OY`qcXO+F1 zpyk2cc^9kvA6Wh$rM6vhsJyoT4t-#tD(5}GnYI;_yz0~!bd7(E5d#uGn2XPfL0#oj zg~#Kn0^cC%<=9LA`48*A{?X& z%p4KmRtxWKVqfxp=dpiXdF>>B6ua!#Vr^OgUvc19yc&NKL@uq^)vtOkO}?&3ut~9y zIucJ;EN~Z)gpZ_Tlv8wRRxyu6#KfLf#-@0s zi3hpaQsg8&vs|*l5`+VT-Rc*qL%kQ_!Vt?{N{EW3%1l;`#;;YgS|CEF;hG6PcCp$- z=&D+be)D){SM2x^9flAjbOoc^N`d+?C^xN`N+D~3KaPE^5Oqwbx)Go!xn>NDDTrq) zh6OGymckntd*|^;$9Li={A^ESTmvk)C@OSHF=KQLLcT&&v0!Jg9268Ixwr-EE-xyz zAU9Z6P^=={{v=dgD^9gGTn->xW|o#Uf%G7{VrS#Cvz2z6x<24t2^flN16=hoKZ-<+ zs(%5UlYH0z-b{_r_+wm>!^`TDp|o4Em3_NJus)}+0PhsO0A2|Z9A9I>*CYFS3ce$< z??~SFNY+wp?<%w(lG_gz+E2^vr*9q2w~x!FwybsaLNU-$2n@-Ap+ex495{8WDK|Nl z3!KUa&da8jtY!6a;CJJ>SDq{QdSzel=E?_;=X||+-}$WNuGP0X|IIz$>U+EIj&-ou z*zs22^}x65HikF6a^pbuiK5k)v$lTIvmW}Jz~=Dn_QQ8Qd)9~6#k{94XY0G?aKD_q zmRt+v9X(lnv7x2V&?`6e78(x84Ttg#4`uCl{q5`Bd4I6rACmn;n}-SqN9BW~dH-Ws z+dYpjYbtt~f_G5%4*nEEb0c@1?$u)j&wkmnf76@yJd}4H+3L(YPZZnxS50f)Z#cG% zly3+Ezn$G#>s^1-t4mp1(OH)@6dT*$YJ0t{&=|aSE<3W?QS<}~o?=3*#7W?DH_ZM?h&*Uz=kRSci z+{|oYCN9I@XgvEk2H?bkKPdZyx1GU_sm+n^KKb60xBG|RfA(&D<62*#>4@BPWXq6m zI-ajTv8uo8ZCpRJslR0{cJzJU_yf!LET8IW-w3@*->XE6hGfss=7BpfDQ}?Q?UudW z>lg2MgGFyc!Mj)X?k#wSWH0pCl5*amJKnLPtEChu_5Re=eXq5B)v!ACs%5*2YTCEm zNBMxsKH0l(^L);`FYkRcYrgAiSo6Mm0ld)VhPbMy<;Ka^PG--b!47|IWo_!M3$I_e z?d<-@At$nZW`|G}IPk<|Q7H$4F53Wgtz%e;+Y-=hXIGt}k zbBmRmpO9@&T>Ao8y?HONXzt!H6nY+(;jii8?;qT{_`{A{v485zKRO{dO=O+WqoWIE z{~+(XZr^I6P=8RaKe#!PuYV-(KDsp}yN9!9iY>hxQMu(n_DNLH@w~HreL{8~++=S% zkK!bodo~Wq&4byoV%Gs!S6zqSG`?ki!~FfjKREimquFzJUA{H(=9x`fzVG-4y~RM! z<~gu>-O23mJ*TJO+#@^p6r2OHb714*hoI0v&&{)cckZ2Yw*!Y?Y1(?U;2M`*;~!kS z8V>#ErhpxxB{V?zQEtIzrgst-x>N@4R&O+Tmx$a=Gv1Oxq zYx&O9^M$Etd1@L|=38}5yZ-}d8le4XpO>*Ch{pmxdLln+6c*kBU@g8jhyKOiEjs8WFuDe@cRn&77fqD>Qx${$X5U7 zSVztPgu=ibPg=>5fKWcI4np@`K>ya9W(OuYJ`NjI4bO>7e3Jhv)Wgv&Th*0pdxSxy zR_&n}W>Tp{kRA`Z@wN|tBm~DFI0`MMSolc~(6(?Ge`Hc>sSStOR1~V~M$V27pC6kF zpB*4VuONQUPw;X9E1<|_4x&T2GmgAKh!Z+)@3uh_J&*gAxNKzQ4Vt?huifN0=ipR~3a zecL9=>E1RYWdZANTamI+o`!8ZQVy!2dE1GU3#mG!+}P?t%1hbZ+dib~sru$$`H?zJ zQ#I~i>S$xzT6d_@HQNi5hO`00##0Ih$ZuKaTva#kbWw>` z(A3|^@fg8Zz(L>_Mj**Ipu~S2vzIYj#q1T#Ud0UK;yiq=!G#QF3}lMwnCf&V_`iZ8 zT1&uS16&}`^d~wGZTQtbithLY<^MV5`2`jD7wXv0sF8ctK+eKkslDfE&QUI~PL8UB ieHQ74rX1xfHZ)>xcVDr5T9eYzCv%OxzoL*JZ~kv3dJ2vJ literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/kelly_dropout.py b/code/tac-qlib/tac_qlib/contrib/strategy/kelly_dropout.py new file mode 100644 index 0000000..896ef74 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/kelly_dropout.py @@ -0,0 +1,201 @@ +"""Fractional-Kelly dropout strategy for cross-sectional signals. + +Sizing rule variant of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy``: +the topk/n_drop SELECTION is identical to the reference, but the buy size is +proportional to the score MAGNITUDE (edge) instead of equal-weight, capped at a +fraction ``cap_frac`` of the equal-weight notional so a single name cannot +over-concentrate the book. + +``cap_frac`` is the fraction of the equal-weight per-name notional that a top +signal can deploy at most (e.g. 0.5 = at most half the equal-weight size). +Names whose score is below the median of the buy set get a proportionally +smaller slice; the residual stays in cash (that is the point of the rule: +throw away less edge per name, deploy less capital when conviction is low). +""" + +from __future__ import annotations + +from typing import List + +import numpy as np +import pandas as pd + +from qlib.backtest import Order +from qlib.backtest.decision import OrderDir, TradeDecisionWO +from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy + +__all__ = ["FractionalKellyDropoutStrategy"] + +DEFAULT_CAP_FRAC = 0.5 + + +class FractionalKellyDropoutStrategy(TopkDropoutStrategy): + """TopkDropout selection with score-magnitude (fractional-Kelly) sizing. + + Parameters + ---------- + topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable, + forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``. + cap_frac : max buy notional as a fraction of the equal-weight notional. + """ + + def __init__(self, *, topk, n_drop, cap_frac: float = DEFAULT_CAP_FRAC, **kwargs): + super().__init__(topk=topk, n_drop=n_drop, **kwargs) + self.cap_frac = cap_frac + + def generate_trade_decision(self, execute_result=None): + import copy + + 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 isinstance(pred_score, pd.DataFrame): + pred_score = pred_score.iloc[:, 0] + if pred_score is None: + return TradeDecisionWO([], self) + + if self.only_tradable: + + def get_first_n(li, n, reverse=False): + cur_n = 0 + res = [] + for si in reversed(li) if reverse else li: + if self.trade_exchange.is_stock_tradable( + stock_id=si, start_time=trade_start_time, end_time=trade_end_time + ): + res.append(si) + cur_n += 1 + if cur_n >= n: + break + return res[::-1] if reverse else res + + def get_last_n(li, n): + return get_first_n(li, n, reverse=True) + + def filter_stock(li): + return [ + si + for si in li + if self.trade_exchange.is_stock_tradable( + stock_id=si, start_time=trade_start_time, end_time=trade_end_time + ) + ] + + else: + + def get_first_n(li, n): + return list(li)[:n] + + def get_last_n(li, n): + return list(li)[-n:] + + def filter_stock(li): + return li + + current_temp: "object" = copy.deepcopy(self.trade_position) + sell_order_list: List[Order] = [] + buy_order_list: List[Order] = [] + cash = current_temp.get_cash() + current_stock_list = current_temp.get_stock_list() + last = pred_score.reindex(current_stock_list).sort_values(ascending=False).index + + if self.method_buy == "top": + today = get_first_n( + pred_score[~pred_score.index.isin(last)].sort_values(ascending=False).index, + self.n_drop + self.topk - len(last), + ) + elif self.method_buy == "random": + topk_candi = get_first_n(pred_score.sort_values(ascending=False).index, self.topk) + candi = list(filter(lambda x: x not in last, topk_candi)) + n = self.n_drop + self.topk - len(last) + try: + today = np.random.choice(candi, n, replace=False) + except ValueError: + today = candi + else: + raise NotImplementedError(f"This type of input is not supported") + + comb = pred_score.reindex(last.union(pd.Index(today))).sort_values(ascending=False).index + + if self.method_sell == "bottom": + sell = last[last.isin(get_last_n(comb, self.n_drop))] + elif self.method_sell == "random": + candi = filter_stock(last) + try: + sell = pd.Index(np.random.choice(candi, self.n_drop, replace=False) if len(last) else []) + except ValueError: + sell = candi + else: + raise NotImplementedError(f"This type of input is not supported") + + buy = today[: len(sell) + self.topk - len(last)] + for code in current_stock_list: + if not self.trade_exchange.is_stock_tradable( + stock_id=code, + start_time=trade_start_time, + end_time=trade_end_time, + direction=None if self.forbid_all_trade_at_limit else OrderDir.SELL, + ): + continue + if code in sell: + time_per_step = self.trade_calendar.get_freq() + if current_temp.get_stock_count(code, bar=time_per_step) < self.hold_thresh: + continue + sell_amount = current_temp.get_stock_amount(code=code) + sell_order = Order( + stock_id=code, + amount=sell_amount, + start_time=trade_start_time, + end_time=trade_end_time, + direction=Order.SELL, + ) + if self.trade_exchange.check_order(sell_order): + sell_order_list.append(sell_order) + trade_val, trade_cost, trade_price = self.trade_exchange.deal_order( + sell_order, position=current_temp + ) + cash += trade_val - trade_cost + + if len(buy) == 0: + return TradeDecisionWO(sell_order_list, self) + + # ---- fractional-Kelly sizing -------------------------------------- + # equal-weight notional (reference baseline) + eq_notional = cash * self.risk_degree / len(buy) + buy_scores = pred_score.reindex(buy).astype(float) + lo, hi = buy_scores.min(), buy_scores.max() + if hi == lo: + w = pd.Series(1.0, index=buy_scores.index) + else: + w = (buy_scores - lo) / (hi - lo) # [0,1] edge magnitude + w = w.clip(lower=0.0) + w_max = w.max() + w = w / w_max if w_max > 0 else w # max == 1.0 + for code in buy: + if not self.trade_exchange.is_stock_tradable( + stock_id=code, + start_time=trade_start_time, + end_time=trade_end_time, + direction=None if self.forbid_all_trade_at_limit else OrderDir.BUY, + ): + continue + buy_price = self.trade_exchange.get_deal_price( + stock_id=code, start_time=trade_start_time, end_time=trade_end_time, direction=OrderDir.BUY + ) + notional = eq_notional * min(self.cap_frac, float(w.get(code, 0.0))) + buy_amount = notional / buy_price + factor = self.trade_exchange.get_factor( + stock_id=code, start_time=trade_start_time, end_time=trade_end_time + ) + buy_amount = self.trade_exchange.round_amount_by_trade_unit(buy_amount, factor) + buy_order = Order( + stock_id=code, + amount=buy_amount, + start_time=trade_start_time, + end_time=trade_end_time, + direction=Order.BUY, + ) + buy_order_list.append(buy_order) + + return TradeDecisionWO(sell_order_list + buy_order_list, self) \ No newline at end of file diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/long_short.py b/code/tac-qlib/tac_qlib/contrib/strategy/long_short.py new file mode 100644 index 0000000..9090fc6 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/long_short.py @@ -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) diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/optimal_stop.py b/code/tac-qlib/tac_qlib/contrib/strategy/optimal_stop.py new file mode 100644 index 0000000..79aaad9 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/optimal_stop.py @@ -0,0 +1,217 @@ +"""Optimal-stopping / stochastic-control strategy for cross-sectional signals. + +Entry is a control policy: a symbol opens a position only when its cross-sectional +signal percentile is at or above ``entry_pct`` (i.e. it is one of the top-ranked +names) and the portfolio has fewer than ``topk`` open positions. + +Exit is an optimal-stopping rule: a held position is stopped (closed) when its +signal percentile falls below ``exit_pct`` (the continuation value of holding is +no longer worth the risk), OR after ``max_hold_days`` (time stop / finite +horizon), OR when the position P&L breaches ``sl`` (loss control) and the +position has been held at least ``min_hold_days``. + +Sizing is fixed ``notional`` per position (equal-weight control), unlike the +TopkDropout cash-allocation heuristic. + +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 + +from typing import List + +import pandas as pd + +from qlib.backtest import Order +from qlib.backtest.decision import OrderDir, TradeDecisionWO +from qlib.contrib.strategy.signal_strategy import BaseSignalStrategy + +__all__ = ["OptimalStopControl"] + +DEFAULT_NOTIONAL = 20_000.0 +DEFAULT_ENTRY_PCT = 0.80 +DEFAULT_EXIT_PCT = 0.50 +DEFAULT_MAX_HOLD_DAYS = 10 +DEFAULT_MIN_HOLD_DAYS = 2 +DEFAULT_SL = -0.06 + + +class OptimalStopControl(BaseSignalStrategy): + """Optimal-stopping long-only strategy over a cross-sectional signal. + + Parameters + ---------- + topk : max number of concurrent positions. + entry_pct : min cross-sectional score percentile required to OPEN (0..1). + exit_pct : held positions are stopped when score percentile < exit_pct. + max_hold_days : hard time stop (finite-horizon close). + min_hold_days : minimum holding days before stop-loss is evaluated. + notional : $ per position (equal-weight control). + sl : stop-loss threshold as fraction of entry price (<= 0), disabled if 0. + """ + + def __init__( + self, + *, + signal=None, + topk: int = 10, + entry_pct: float = DEFAULT_ENTRY_PCT, + exit_pct: float = DEFAULT_EXIT_PCT, + max_hold_days: int = DEFAULT_MAX_HOLD_DAYS, + min_hold_days: int = DEFAULT_MIN_HOLD_DAYS, + notional: float = DEFAULT_NOTIONAL, + sl: float = DEFAULT_SL, + risk_degree: float = 0.95, + trade_exchange=None, + level_infra=None, + common_infra=None, + **kwargs, + ): + super().__init__( + signal=signal, + trade_exchange=trade_exchange, + level_infra=level_infra, + common_infra=common_infra, + **kwargs, + ) + self.topk = topk + self.entry_pct = entry_pct + self.exit_pct = exit_pct + self.max_hold_days = max_hold_days + self.min_hold_days = min_hold_days + self.notional = notional + self.sl = sl + + # ------------------------------------------------------------------ utils + @staticmethod + def _pct_rank(score: pd.Series) -> pd.Series: + return score.rank(pct=True) + + def _entry_price(self, pos) -> float: + # Position stores avg entry price under key "price" (see Position.position) + price = pos.position.get("price") + if price is None: + price = pos.get_stock_amount("price") + return float(price) + + def _pnl_pct(self, pos, mark: float) -> float: + entry = self._entry_price(pos) + if not entry or entry != entry: + return 0.0 + return mark / entry - 1.0 + + def _is_tradable(self, code, start, end, direction) -> bool: + try: + return self.trade_exchange.is_stock_tradable( + stock_id=code, start_time=start, end_time=end, direction=direction + ) + except TypeError: # some exchanges take no direction kwarg + return self.trade_exchange.is_stock_tradable(stock_id=code, start_time=start, end_time=end) + + # ------------------------------------------------------------ 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) + + pct = self._pct_rank(pred_score) + time_per_step = self.trade_calendar.get_freq() + current_temp = __import__("copy").deepcopy(self.trade_position) + + holdings = {} + for code in current_temp.get_stock_list(): + if abs(current_temp.get_stock_amount(code)) > 1e-6: + holdings[code] = current_temp + + # ---- optimal stopping: close held positions ----------------------- + sell_orders: List[Order] = [] + closed_today = set() + kept = {} + for code, pos in holdings.items(): + held = current_temp.get_stock_count(code, bar=time_per_step) + mark = self.trade_exchange.get_deal_price( + stock_id=code, start_time=trade_start, end_time=trade_end, direction=Order.SELL + ) + if mark is None or mark != mark: + continue + rank = pct.get(code, 0.0) + stop_pnl = held >= self.min_hold_days and self.sl < 0 and self._pnl_pct(pos, mark) <= self.sl + if held >= self.max_hold_days or rank < self.exit_pct or stop_pnl: + amt = abs(current_temp.get_stock_amount(code)) + o = Order(stock_id=code, amount=amt, start_time=trade_start, + end_time=trade_end, direction=Order.SELL) + if self.trade_exchange.check_order(o): + sell_orders.append(o) + self.trade_exchange.deal_order(o, position=current_temp) + closed_today.add(code) + else: + kept[code] = mark + + # ---- equal-weight control: target notional per name ----------------- + # candidate opens: top-ranked names whose signal pct >= entry_pct + rank_desc = pred_score.sort_values(ascending=False) + held_codes = set(kept) + opens = [] + for sym in rank_desc.index: + if len(opens) >= self.topk: + break + if sym in held_codes: + continue + if pct.get(sym, 0.0) < self.entry_pct: + continue + if not self._is_tradable(sym, trade_start, trade_end, OrderDir.BUY): + continue + opens.append(sym) + + targets = held_codes | set(opens) + if not targets: + return TradeDecisionWO(sell_orders, self) + + # total value (cash + marked positions) -> per-target notional + total_value = current_temp.get_cash() + for code, mark in kept.items(): + total_value += abs(current_temp.get_stock_amount(code)) * mark + + target_notional = total_value * self.risk_degree / max(1, len(targets)) + + # ---- rebalance kept positions toward target weight ------------------ + buy_orders: List[Order] = [] + for code, mark in kept.items(): + cur = abs(current_temp.get_stock_amount(code)) * mark + diff_notional = target_notional - cur + if abs(diff_notional) / target_notional < 0.02: + continue # skip tiny rebalances + amount_delta = diff_notional / mark + direction = Order.BUY if amount_delta > 0 else Order.SELL + o = Order(stock_id=code, amount=abs(amount_delta), 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) + self.trade_exchange.deal_order(o, position=current_temp) + + # ---- open new positions at target weight ---------------------------- + for sym in opens: + px = self.trade_exchange.get_deal_price( + stock_id=sym, start_time=trade_start, end_time=trade_end, direction=OrderDir.BUY + ) + if px is None or px != px or px <= 0: + continue + amount = target_notional / px + factor = self.trade_exchange.get_factor( + stock_id=sym, start_time=trade_start, end_time=trade_end + ) + amount = self.trade_exchange.round_amount_by_trade_unit(amount, factor) + o = Order(stock_id=sym, amount=amount, start_time=trade_start, + end_time=trade_end, direction=Order.BUY) + if self.trade_exchange.check_order(o): + buy_orders.append(o) + + return TradeDecisionWO(sell_orders + buy_orders, self) diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/regime_gate.py b/code/tac-qlib/tac_qlib/contrib/strategy/regime_gate.py new file mode 100644 index 0000000..5b9acfb --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/regime_gate.py @@ -0,0 +1,231 @@ +"""HMM-regime overlay TopkDropout strategy. + +Regime-gate overlay on ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy``: +selection and sizing are identical to the reference, but a name is only BOUGHT +(entry gate) when its per-symbol HMM regime posterior ``sp_hmm_p_regime1`` on +the signal date is >= ``regime_threshold``; otherwise it is held in cash instead +of being opened. + +The regime posterior is read from the lake feature provider on the fly via +``qlib.data.D.features`` (field ``$sp_hmm_p_regime1``) for the signal window, so +no regime column needs to enter the model's ``feature_fields`` — the gate is a +pure overlay (book ch.01: regime flags regressed as model features, survived +only as an overlay). The HMM itself was fit with ``fit_end=`` when +the lake features were backfilled, so there is no lookahead. + +Names already held are NOT force-sold when the regime turns unfavourable +(entry gate only, matching the queue-10 design). +""" + +from __future__ import annotations + +from typing import List + +import numpy as np +import pandas as pd + +from qlib.backtest import Order +from qlib.backtest.decision import OrderDir, TradeDecisionWO +from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy + +try: + from qlib.data import D +except ImportError: # pragma: no cover - qlib always present in this stack + D = None + +__all__ = ["RegimeGateDropoutStrategy"] + +DEFAULT_REGIME_THRESHOLD = 0.5 +REGIME_FIELD = "$sp_hmm_p_regime1" + + +class RegimeGateDropoutStrategy(TopkDropoutStrategy): + """TopkDropout with an HMM-regime entry gate on buy candidates. + + Parameters + ---------- + topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable, + forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``. + regime_threshold : minimum ``sp_hmm_p_regime1`` posterior required to open a + new position (default 0.5). + """ + + def __init__(self, *, topk, n_drop, regime_threshold: float = DEFAULT_REGIME_THRESHOLD, **kwargs): + super().__init__(topk=topk, n_drop=n_drop, **kwargs) + self.regime_threshold = regime_threshold + + def _regime_for(self, codes, pred_start, pred_end) -> pd.Series: + """Return {code: sp_hmm_p_regime1} for the signal window (last day).""" + if D is None: + return pd.Series(dtype=float) + try: + df = D.features(list(codes), [REGIME_FIELD], start_time=pred_start, end_time=pred_end, freq="day") + except Exception: # noqa: BLE001 - a regime read failure should gate open, not crash + return pd.Series(dtype=float) + if df is None or len(df) == 0: + return pd.Series(dtype=float) + # df index is MultiIndex (datetime, instrument); take the last day's values + df = df.reset_index() + ts_col = "datetime" if "datetime" in df.columns else df.columns[0] + sym_col = "instrument" if "instrument" in df.columns else df.columns[1] + last_ts = df[ts_col].max() + last = df[df[ts_col] == last_ts] + out = {} + for _, row in last.iterrows(): + sym = str(row[sym_col]).split("/")[-1].upper() + val = row.iloc[-1] + out[sym] = float(val) if val == val else np.nan + return pd.Series(out) + + def generate_trade_decision(self, execute_result=None): + import copy + + 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 isinstance(pred_score, pd.DataFrame): + pred_score = pred_score.iloc[:, 0] + if pred_score is None: + return TradeDecisionWO([], self) + + if self.only_tradable: + + def get_first_n(li, n, reverse=False): + cur_n = 0 + res = [] + for si in reversed(li) if reverse else li: + if self.trade_exchange.is_stock_tradable( + stock_id=si, start_time=trade_start_time, end_time=trade_end_time + ): + res.append(si) + cur_n += 1 + if cur_n >= n: + break + return res[::-1] if reverse else res + + def get_last_n(li, n): + return get_first_n(li, n, reverse=True) + + def filter_stock(li): + return [ + si + for si in li + if self.trade_exchange.is_stock_tradable( + stock_id=si, start_time=trade_start_time, end_time=trade_end_time + ) + ] + + else: + + def get_first_n(li, n): + return list(li)[:n] + + def get_last_n(li, n): + return list(li)[-n:] + + def filter_stock(li): + return li + + current_temp: "object" = copy.deepcopy(self.trade_position) + sell_order_list: List[Order] = [] + buy_order_list: List[Order] = [] + cash = current_temp.get_cash() + current_stock_list = current_temp.get_stock_list() + last = pred_score.reindex(current_stock_list).sort_values(ascending=False).index + + if self.method_buy == "top": + today = get_first_n( + pred_score[~pred_score.index.isin(last)].sort_values(ascending=False).index, + self.n_drop + self.topk - len(last), + ) + elif self.method_buy == "random": + topk_candi = get_first_n(pred_score.sort_values(ascending=False).index, self.topk) + candi = list(filter(lambda x: x not in last, topk_candi)) + n = self.n_drop + self.topk - len(last) + try: + today = np.random.choice(candi, n, replace=False) + except ValueError: + today = candi + else: + raise NotImplementedError(f"This type of input is not supported") + + comb = pred_score.reindex(last.union(pd.Index(today))).sort_values(ascending=False).index + + if self.method_sell == "bottom": + sell = last[last.isin(get_last_n(comb, self.n_drop))] + elif self.method_sell == "random": + candi = filter_stock(last) + try: + sell = pd.Index(np.random.choice(candi, self.n_drop, replace=False) if len(last) else []) + except ValueError: + sell = candi + else: + raise NotImplementedError(f"This type of input is not supported") + + buy = today[: len(sell) + self.topk - len(last)] + + # ---- regime gate ----------------------------------------------------- + if buy: + regime = self._regime_for(buy, pred_start_time, pred_end_time) + gated = [c for c in buy if regime.get(c, np.nan) >= self.regime_threshold] + else: + gated = [] + + for code in current_stock_list: + if not self.trade_exchange.is_stock_tradable( + stock_id=code, + start_time=trade_start_time, + end_time=trade_end_time, + direction=None if self.forbid_all_trade_at_limit else OrderDir.SELL, + ): + continue + if code in sell: + time_per_step = self.trade_calendar.get_freq() + if current_temp.get_stock_count(code, bar=time_per_step) < self.hold_thresh: + continue + sell_amount = current_temp.get_stock_amount(code=code) + sell_order = Order( + stock_id=code, + amount=sell_amount, + start_time=trade_start_time, + end_time=trade_end_time, + direction=Order.SELL, + ) + if self.trade_exchange.check_order(sell_order): + sell_order_list.append(sell_order) + trade_val, trade_cost, trade_price = self.trade_exchange.deal_order( + sell_order, position=current_temp + ) + cash += trade_val - trade_cost + + if len(gated) == 0: + return TradeDecisionWO(sell_order_list, self) + + value = cash * self.risk_degree / len(gated) + for code in gated: + if not self.trade_exchange.is_stock_tradable( + stock_id=code, + start_time=trade_start_time, + end_time=trade_end_time, + direction=None if self.forbid_all_trade_at_limit else OrderDir.BUY, + ): + continue + buy_price = self.trade_exchange.get_deal_price( + stock_id=code, start_time=trade_start_time, end_time=trade_end_time, direction=OrderDir.BUY + ) + buy_amount = value / buy_price + factor = self.trade_exchange.get_factor( + stock_id=code, start_time=trade_start_time, end_time=trade_end_time + ) + buy_amount = self.trade_exchange.round_amount_by_trade_unit(buy_amount, factor) + buy_order = Order( + stock_id=code, + amount=buy_amount, + start_time=trade_start_time, + end_time=trade_end_time, + direction=Order.BUY, + ) + buy_order_list.append(buy_order) + + return TradeDecisionWO(sell_order_list + buy_order_list, self) \ No newline at end of file diff --git a/code/tac-qlib/tac_qlib/contrib/strategy/weekly_rebalance.py b/code/tac-qlib/tac_qlib/contrib/strategy/weekly_rebalance.py new file mode 100644 index 0000000..839abb8 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/weekly_rebalance.py @@ -0,0 +1,223 @@ +"""Weekly-rebalance TopkDropout strategy. + +Turnover-reduction variant of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy``: +the topk/n_drop selection and sizing are identical to the reference, but the +target book is recomputed only on the first trading day of each ISO week; on the +other days the strategy issues NO orders (holds the book untouched). + +The weekly cadence is derived from the qlib trade calendar: a rebalance happens +when the current trade step's date belongs to a different ISO ``(year, week)`` +than the previous trade step. ``hold_band_pct`` (default 0) optionally skips +tiny rebalances: when a name's existing position differs from the new target by +less than this fraction, no order is generated for it. +""" + +from __future__ import annotations + +from typing import List + +import numpy as np +import pandas as pd + +from qlib.backtest import Order +from qlib.backtest.decision import OrderDir, TradeDecisionWO +from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy + +__all__ = ["WeeklyRebalanceDropoutStrategy"] + +DEFAULT_HOLD_BAND_PCT = 0.0 + + +class WeeklyRebalanceDropoutStrategy(TopkDropoutStrategy): + """TopkDropout rebalanced once per ISO week; holds otherwise. + + Parameters + ---------- + topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable, + forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``. + hold_band_pct : skip order for a name whose deviation from target weight is + below this fraction of the target (no-trade buffer band). + rebalance_every_n_weeks : rebalance every N ISO weeks instead of every week + (default 1 = weekly; 2 = biweekly). Ignored when + ``rebalance_every_n_days`` is set. + rebalance_every_n_days : rebalance every N trading days (daily when N=1). + When set, overrides the weekly gating logic entirely. + """ + + def __init__(self, *, topk, n_drop, hold_band_pct: float = DEFAULT_HOLD_BAND_PCT, + rebalance_every_n_weeks: int = 1, + rebalance_every_n_days: int = 0, **kwargs): + super().__init__(topk=topk, n_drop=n_drop, **kwargs) + self.hold_band_pct = hold_band_pct + self.rebalance_every_n_weeks = rebalance_every_n_weeks + self.rebalance_every_n_days = rebalance_every_n_days + + @staticmethod + def _iso_week(ts) -> tuple: + return (ts.year, ts.week) + + def generate_trade_decision(self, execute_result=None): + import copy + + trade_step = self.trade_calendar.get_trade_step() + trade_start_time, trade_end_time = self.trade_calendar.get_step_time(trade_step) + + if self.rebalance_every_n_days > 0: + # daily gating: count trading steps since last rebalance + step_num = trade_step + if hasattr(self, "_last_rebal_step"): + if (step_num - self._last_rebal_step) < self.rebalance_every_n_days: + return TradeDecisionWO([], self) + self._last_rebal_step = step_num + else: + cur_week = self._iso_week(trade_start_time) + prev_week = getattr(self, "_last_week", None) + self._last_week = cur_week + + if prev_week is not None and prev_week == cur_week: + return TradeDecisionWO([], self) + + if self.rebalance_every_n_weeks > 1: + week_num = cur_week[1] + if prev_week is not None and (week_num % self.rebalance_every_n_weeks) != 1: + return TradeDecisionWO([], self) + + 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 isinstance(pred_score, pd.DataFrame): + pred_score = pred_score.iloc[:, 0] + if pred_score is None: + return TradeDecisionWO([], self) + + if self.only_tradable: + + def get_first_n(li, n, reverse=False): + cur_n = 0 + res = [] + for si in reversed(li) if reverse else li: + if self.trade_exchange.is_stock_tradable( + stock_id=si, start_time=trade_start_time, end_time=trade_end_time + ): + res.append(si) + cur_n += 1 + if cur_n >= n: + break + return res[::-1] if reverse else res + + def get_last_n(li, n): + return get_first_n(li, n, reverse=True) + + def filter_stock(li): + return [ + si + for si in li + if self.trade_exchange.is_stock_tradable( + stock_id=si, start_time=trade_start_time, end_time=trade_end_time + ) + ] + + else: + + def get_first_n(li, n): + return list(li)[:n] + + def get_last_n(li, n): + return list(li)[-n:] + + def filter_stock(li): + return li + + current_temp: "object" = copy.deepcopy(self.trade_position) + sell_order_list: List[Order] = [] + buy_order_list: List[Order] = [] + cash = current_temp.get_cash() + current_stock_list = current_temp.get_stock_list() + last = pred_score.reindex(current_stock_list).sort_values(ascending=False).index + + if self.method_buy == "top": + today = get_first_n( + pred_score[~pred_score.index.isin(last)].sort_values(ascending=False).index, + self.n_drop + self.topk - len(last), + ) + elif self.method_buy == "random": + topk_candi = get_first_n(pred_score.sort_values(ascending=False).index, self.topk) + candi = list(filter(lambda x: x not in last, topk_candi)) + n = self.n_drop + self.topk - len(last) + try: + today = np.random.choice(candi, n, replace=False) + except ValueError: + today = candi + else: + raise NotImplementedError(f"This type of input is not supported") + + comb = pred_score.reindex(last.union(pd.Index(today))).sort_values(ascending=False).index + + if self.method_sell == "bottom": + sell = last[last.isin(get_last_n(comb, self.n_drop))] + elif self.method_sell == "random": + candi = filter_stock(last) + try: + sell = pd.Index(np.random.choice(candi, self.n_drop, replace=False) if len(last) else []) + except ValueError: + sell = candi + else: + raise NotImplementedError(f"This type of input is not supported") + + buy = today[: len(sell) + self.topk - len(last)] + for code in current_stock_list: + if not self.trade_exchange.is_stock_tradable( + stock_id=code, + start_time=trade_start_time, + end_time=trade_end_time, + direction=None if self.forbid_all_trade_at_limit else OrderDir.SELL, + ): + continue + if code in sell: + time_per_step = self.trade_calendar.get_freq() + if current_temp.get_stock_count(code, bar=time_per_step) < self.hold_thresh: + continue + sell_amount = current_temp.get_stock_amount(code=code) + sell_order = Order( + stock_id=code, + amount=sell_amount, + start_time=trade_start_time, + end_time=trade_end_time, + direction=Order.SELL, + ) + if self.trade_exchange.check_order(sell_order): + sell_order_list.append(sell_order) + trade_val, trade_cost, trade_price = self.trade_exchange.deal_order( + sell_order, position=current_temp + ) + cash += trade_val - trade_cost + + if len(buy) == 0: + return TradeDecisionWO(sell_order_list, self) + + value = cash * self.risk_degree / len(buy) + for code in buy: + if not self.trade_exchange.is_stock_tradable( + stock_id=code, + start_time=trade_start_time, + end_time=trade_end_time, + direction=None if self.forbid_all_trade_at_limit else OrderDir.BUY, + ): + continue + buy_price = self.trade_exchange.get_deal_price( + stock_id=code, start_time=trade_start_time, end_time=trade_end_time, direction=OrderDir.BUY + ) + buy_amount = value / buy_price + factor = self.trade_exchange.get_factor( + stock_id=code, start_time=trade_start_time, end_time=trade_end_time + ) + buy_amount = self.trade_exchange.round_amount_by_trade_unit(buy_amount, factor) + buy_order = Order( + stock_id=code, + amount=buy_amount, + start_time=trade_start_time, + end_time=trade_end_time, + direction=Order.BUY, + ) + buy_order_list.append(buy_order) + + return TradeDecisionWO(sell_order_list + buy_order_list, self) \ No newline at end of file diff --git a/code/tac-qlib/tac_qlib/data/__init__.py b/code/tac-qlib/tac_qlib/data/__init__.py new file mode 100644 index 0000000..92e6e90 --- /dev/null +++ b/code/tac-qlib/tac_qlib/data/__init__.py @@ -0,0 +1,25 @@ +from .config import ( + LakeConfig, + BAR_FIELD_MAP, + FREQ_TO_TIMEFRAME, + UNKNOWN_FIELD_NAMES, + timeframe_for_freq, + resolve_lake_root, +) +from .providers import ( + LakeCalendarProvider, + LakeInstrumentProvider, + LakeFeatureProvider, +) + +__all__ = [ + "LakeConfig", + "BAR_FIELD_MAP", + "FREQ_TO_TIMEFRAME", + "UNKNOWN_FIELD_NAMES", + "timeframe_for_freq", + "resolve_lake_root", + "LakeCalendarProvider", + "LakeInstrumentProvider", + "LakeFeatureProvider", +] diff --git a/code/tac-qlib/tac_qlib/data/__pycache__/__init__.cpython-312.pyc b/code/tac-qlib/tac_qlib/data/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..0f1eb6f41d44bd9064b1f16f529dac6bb9d7c3fd GIT binary patch literal 522 zcmaKozfQw25XS8!O`5cofFfo#$N+r-gecHT6*NF8Lh^Dl#6~NeR{=qvb2|6 zsLFP|1I<=re;3;ILp9&*G>)Kq0Nqx*(d^XQ4OKmf_M(H+Xj&C%?N}}3iC{fR1%qBD zp(ok3nwM;l@f!wQ+k?!qJhau~WgMxu;;29JRd;_yYS<;BYvU1NSpZmX0`TglQgFhC^1E8D(JfKhQh`v{9 zwBN7g^nC4_Cub&sNfOhX)&P<;$pO~;BURiGSv=%yQ!eN}v+lvN#+!{X9$Ox^#>imK S!PYy{_$UV@>(-NVy66X;=7>fB literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/data/__pycache__/config.cpython-312.pyc b/code/tac-qlib/tac_qlib/data/__pycache__/config.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..cf99dd8a481f6e2d1f323f8d8d754730eadb357a GIT binary patch literal 11507 zcmc&)drTZhn(vqix*d!R+fW2!Hr)lU02G1+q4VcN0 zwe}(%zBn?htphfROm0^SC^>~_ceTE~)5*Jvbe2wcnz`N;vn^H|WltyNKN51fTmN(S zRrU1DfHB_NUFpiE>YI9eRrOWX_xs|2}8kpwYM zTqZc8pWsMN7a)VWejSNueL&x@hm;CXK|{X*=NSUVpsC*!H20f>mVQgBp0(fRCpaT# zx=V5<9Gf`{&{iDVI6KfeTrR{8&IxfImk+UvD}dOIV-HsdbP-n!aS2xnaT$*HaC?C+ zM{&QMqxl@J;v@{6dyd=3?dJ|~2f0f4SM{E$KNpvJap|Fc2Y0yN$yN8~arIoyTU38O z*T6M$wQm{vT{_|{!PUJ*a7RAUsZ!`Kn9_TXj&pi~=o46u_qEai_Bu}wg+s&sk#T{Q z{NWHi$_FC6Al6!~eWScu!7tJ*J?fv}Ya*;5`B62Vqk(>XoR?@(3JW}4CGtEyFd(r* zHGF8qAL8r8>;6DMtdmezkc+w+81S}Otu%!FeZE$v%lC3S)7IHjhqJ!8_ud!xZo+>$ z$O_ka>D1+3t>C-R;){FtQpJ)#$PWu_kUv$=W!FG_GquF2;9xj#%ICY-RjUmZ)qt>% zzHqLq^$I=63StW_Rnr_R@zr#=njWpD1J(3UH9b*HhajHZHnL%!mBs~LtWz276z`+` zAgw6deYDF^l(@Rve3;i8}tXJXvxnDDZX1Lw&T1}ADLP0Hg~nACPcNfQxaRZ zqa^mBd5b}oIi@BJN2+N-^fUEGtLfU>+N|+ulgtbv#afK)#yY};$C*RGJi?3ygNVwo z)y6o)2KW%i3fmV*>k?8JCu|^pXap8%N~B?Hrnbknl^F_8@B%x+@5&0bu2Q93ot9*^ zDjzQ$2*Xb0V97#Ucv4L9*EXG6N!*1!43!l{I1oWrlwgIwM5oS6DmFG03QPEKiQVv< zy2whSUcExK`G+Kh>hg<{V!jYTr-}_IdYGo7tLGGb&>sRKs8IEYqK@I;$AXHX{svMeZer4F zQmAl*hhdNUM@AKWAPij$1;Qe)7$?Gk@gT2I6O(L2v9c2*3>yrOha@HEjIW33=xpz5 zW6t|7dMU*?%nrdZQj9~RY-og6tXP*Y;~~GK*u`-XI^dY$06XF}C?tpy9b?6)uCHQA zFR)@vt*~CvNvxt1BZ?u&4hi8QFbe5_j0k8(a`T5k%o9?&EJl2kqQ{_od-y1oF-^uu ziPB0ZNTE_PX}>r!M!c)zbdp8Ok!)!y-A3z0OW!5lr!@MzS-fLpF*NQay!z+DnBiobme*uYSczO;~{Z8 z5(x_uoP;!AI_f$>XYx^3pg6roMGsp-F~I8bg76%SO0iyH1LJ(VAfROvF&n7FOJ1Eq zp=Js>EFyD6(~l3-v5`m}G8m{IL@=0uRb=Z_n^YT_Qu5PdWk8QIh>6ev9MGTG1JO<5 z-_5p}wzn;J9Jd^@x%j3pX)^!PS&=YRtmhUbbN4UxEuD*3zwnvyGx3)%zbc=-CcpSv z`~@a%AK29E9M%m&XR&@`Bsd5LRPc;sI1iqW~c2DbU|wkibCFwJ^b2 zuUXNBMa7hwtfC*`K?nAL+5|9!kDGuLiEtVs#R{VU>5q%NpqK!ypoLx`6$|Q}L4y^M zLZgLtR5>qIfk~GN=u};VKd}#@o5V&gVa|)2iq;D}^DoW4G~Yefy>wbGXqd75+E%ch zTe6|k z)B<4J0;v^JE2K6^ZIIeI4`+uohb!cAAk76Jo(rjiE9M-KI;V8r62;mDQfkGJINF9x zwT6i4!}Jh4GzvRT-~lKIBEogBv8vL7t4=`^FF@+xpg)9der@+qYI1mnh=xCM^9r75 zmYGRwhZwsEA^d%YB1NZ}H6q!>V{gN|+j8bJ?{ z4-Df1oFk&faFEpvVP+UVSXhRMx-z1u)iQagDk4>XlUOf5vRW*cpO~R;=O)XZ`+oNl zw^mx8D6N-E8)v9@Z0ag!E64QYk@7}rbPUq%0a?qPR4Uj*T@>slr;CvqjnrB9V^az^ zKencA#ia&YQR1T&i5& z|H->48RXO-SoObMk>V8>&}r< zx^*mV5q4YmcR?8M>yc1;i4`P>6x4rg)^MsQ>C!B~5r0xC_@%^csp7h_)=hU8{-rKw;J%$S#PWW_1u{ zHi9+bIJ7*uD*g10Pv2M##%tOVjZXFG@22%5#59$G+ik?P z%LFlP0L768-C7%h$Ub2)gQhw_jp*V0=+*NxZHyVm(#R{*R_jqDZ#TrKF*s@4gOo;M zW4YVQpl_pw*tC=+Zhwm+Vw8C4nWNo0LQ~AN+Xw~sGe%<4M$k)S)pR~NW^K+;&!v^d z(pM3Vx(-@VFvKX%pk72YYT67G#TiQoSdV^K4@=BE=GHi5(wGIjXW%*e*!mHr=0Z#J zm?u+1quPi&)*;;pF{HZ&rPJ1!bqsa6EpRq%X0S@`m~~EfyNH4|(^eQq0d8j=E8W%z ze#WdH+qIU{=Gz{Mxcw@ujJX%ulDPdO;mwKui!TtMQ)dJWJ6v8vn(0^7Etq{?0GHPY zMur1v$4g9%#PB_Hz=~ip3TgkDNDm@MP>>fTHW;a;dwB3G0bjs^D)Pk+E_!m54*>{G zx7UU-3znB>~y#Zs8cjiUC*vKGi9NLNVNU1Q3HR3?7G3xW${r@!?_r4PF$GUIbKWN{Pgq zo6$0gv!mVDce$sX>F}NJ?CNaq#cD_~;t`9A5pI5PcUN=~PO&321tteO6BSDv-2OU{ za}S&b;6E@R(elDL8kQNFh?9QsPf(l?j(Ba`EE!xhr-I-#K&_ze7Mb*41x4jLp6(yZc~L6_p?E{hPhAyW@ks3%=F*pSFD3@?iK^jqTrKm-hbB-7!lh zi_7nqE|xA|S}2u^>*kDK=?TwSa>GD)=nXW7nVjE$4L{D3q8%jpnbXs5F@9D_LA>Se zjW#B>6*LgLg9U(ty6&iRyIWB}mzfn3ZsIw0VuC zt#77KE7N$Q9UpB*7Amz3=>{??4a2B#Nd` zgf8J&=L$ZcKwXQ{6kH#Xa{*r>6r&Ulfb9Wzc@^dT1*)avhPU1&>{wf-dW~jp+3bSv zLt8Nn(M{rWyK~LHCt=?s+skLjlq$R(O}gAOo$EP{JIpO+;gXzFKBG(8Tni2NTNYaq zwgYOh!#z8=P&5}<^LP^;@5lROPwkqgA>nCQ^Bhljj?12wHIFah@yVW6+0hm^x2c;h z8)TwwffY@F`U5tdKBnJo)6J5C)$*)Or`MJdz-$^uW}uevWsCyQxq~||<4Ob&xdVhl z5bY@~(*QfBL$#uCz@AbwJ-RFyn7*v;1njiUsqsZe<0&(XL7zJSK90g|3|WV@2$*aO zF@p*{HRC*wRqMjtZ1eFULE1Dm?puZr2R1 z=Jlz?ZawVp!Wtffd&M;D7et8xBpd*R&PXcS7tlk-m;VC+EQz(gx#>vFk$U(awKg1S zZmy|ssA;VCx`jJXYnz2oEbTXjc>HiFyn%{iC_*kBwfl^@p#0#uGxm$lSy|W{GY!vcV*$6 z>^`tu@JZ=PX~KOp2{?UXZeroOTzF{NpD1klN>93*KTqY3%7v95H!e?nbYkuB$;9E4 z^5N5o!qZ;Ep{!tmanet zmCKH=m7PkIos!F5m~CG#-E;rs;>l(G^0}29x%Bz9(i4f&6LRUP*|y}K=k7C$jJ&6I zmHKJ!r@8Wx7vp8Vc&ex?<@8FbBcExA4{`JiOE^#|pvd5Ma~wTk9MMYCLSLM}MDR^Uq%_~e4tnfBE6 zWJ&*_^#kj2;}0CGx&+<0MmHzuW|=-AyH2jTdNVkPvUVdZZD-&;ATBRQO4j4*p{yC?R|g99KKTQ$lcq9f5adq$R=5?4W^=Wk#jb zx~#FF+t;?B0o$}zPlMS^#XlCh?uHzY?E|L|m2&zNV*HQ+?mHlc3|a!WPce4;P`v~K z{3$061x_4&^n73TS0q(IKNo@z6W}dRQ%Rt5kzPfhIrV{vs#*sIw5y<|mIelDvHnsY zQ|!WDfk?J#Mio1w#}>sVg&BDI<55QdktM|vLSF*~{83)9p$BpXzj1=U0zPYi2YlgO zXyr9zaJ8z-42XxQkT*mm4XnCUt>wa4S{ozoN{sH!^5O1K{tI9su0r$_9d^5AJ(Ig@m_7xM3yES4l3`;xBG1>t^lF`96Blf|VAXBKXNf@%BY<&~EcWlf3V z*gp`=Wj6()*58H|W8*vnlDwo3Z@~as~eTX8o^MxNlw4!<*@IeyLTQ`mdfe;ajd5icp)%^{nGaTKtl7?P#(?}XF;dlUI zz2OSEY4sR9n@31PBPwe^QRAlBY-lGp3(5?&n&BkODdm0y-*! z2~{$pj{#Uwz1lyAavThk^-&sdJc4pvm%Cr?zVNN?)Hhh&@M%}C;_AKt87XgW>xHh% z=ev7_3BzpL@pZf`N-|q;^ zZ-}Db5|s&}@^^%5!q69eNLHY_HP(~f(Ivq6>?38HWLQ;-e<#tXsb%w z5N#(6_L<8YIf%|xX$PX6DxHVue6_v{(FG{KZ^Mmfk4hIJx=5vq5nZCzFGX~jn!g9p zdsVs|(KKPmnc+4nfQ}P+v!^$SJW{uDh{$m!?T(}^Hwi>;UefMN=6k+!m@9NQb7tP$ vB+O9BLlkzBaWrwFupJYCXRL2vOwKJxtp0PYlTRUwhX_HO}Lk{7K+^S8co&2Ls+llk1{k}ci z0fZnsNvB8-*bQDa@i?J|JnH6!3*^i^>h5tl36X#%l`_6s}x5CDURk0 zQ95P_8fYvVqsE{SeoawR%p5ewEI~`m8nni2L0hi9ENCbF96=}inxn3mJLryif}U7; zu$=U{6T-LGFTa_3RcCcgVi);pm;A=_A;$v6RhF9oc)q1Sj+i1 z2mG!wP)8}wd5Pj&y!&NiuGeZ(>xNnnU#{2Ik=k;ot>8U+?V2%Tz|Zc0q0dAky$r`n zEHfk|Mj{+9h)gdVKF@PZ?-(Nu@XTp}<@mOCW{4G@8RjJ>%AV(2?Do?GyvP>}8j240 zGm*HIAcI7SjdQHPF$8{$4-c^Mh#1?;^z>j&OC%nV8pCXq$6g_A>H{q;Elu`FT$F_2 z7$28bHZXlWD-8?$irPR=56p_wNfB1!MVvjxz~oGCB+lhNRix}aJ?(51`nLiW$6-va z*r|f;fX@p8J2ZqOXcT!~!xgF1LNbBxkBAa4gxF{_r1kCT3A7+J7=afhCef!Zxv{6G zD*+4afWaa!J&_|W%)yTiK~FJ~h%FP(Lq4Bt8}jL1&SV z;l74NpfI8u1?mGR?RbbQaPH~(192?Yn~3gZo@yirkf{Sr44kTS6FogS7zvgOY3Q#2 zz%mFQcaSXpbR-5_!p4RG5g;swICJcn*kb=B-nam*IM{eRA+fl8u{&T=OdXN1q?oz@ ziefuCg!ODxF`tP;Az)O>+jHk#I}C~^w|c$8i^ozu`&u2$Qvoyy#S+|bln+G{{r$Y~ zOJof&qw7f9lcCPzk9Kv0PP9Ft*t_6vwkP6!k$%N@ru*^klV9r2HFvk2c=VLwm0(eQ z0vqE)eF-7dC-BdVS4(U-L@p&c<3k!72b)+Mj)bM*verM7OR+}i3_|iM=K^oM?}cXZyH1&&UMHIQAT?mFB_>A5{u4ox1)IO=9x3y$UmQ}Z&kQDFn- zxnzT1vNi_hUa>Y>Tbs}!>8l$wX#?}WGSxv{F@z2M0QuY|ikdJK^ta^uX(id9K^`pA z>hixD)us)X8%zLkLhpTwDuDZWcvSx`c)(n>#dC#0gGRcRdwXMkhz#^kr%e zyy=^7yWzcAcdc%Qy}oX-wmDtf{BCXQZC|?PK*n=$!FKSj!*%76$w%hG8UMDq@NC65 zpS*f@`s~d57ece)w0~R1v3QbLdt==Zhb;`Bddh7O zz1##EQ+68Lh$aVWbVxd^y%YcNBi~uVqV|rNucxzZc+kht` z4d3)~2NYhQ;Xc5xnkJ1h^>t8@wDU<*Qchu_D4B*cox4>s71AM&QFAo%Mzc5!f>p7F zfPY7%P)KkCOxWAL{yAWrA|LHj$`JjKkVr_1CB_Qpc?m9eC>&))F%%LJFWC8uWa|ao z=J7S9Xapieonaev5$hoPHLwDkYtovw)!ebIy6an$ip*AJd=Dp0U$+UXpu>vw>)ha5 zu$t1`V8WQcg|f~gL~ouN94*t=4V*zXUZ*(YpPRMY{>QXT4VK^spTz|xOtMM*mj{f) z<5!eqqc+=wSvC*4ikjh5+dtvka0sa=D6bP1nu<^pR>`OJhLM$8IseP)E79iYpw9IG zpRvi7LA;|yNw#s8e)<)wb}}Z)=nD#a~zucY0Pm}U=NP&aWe|b1LIDnQ5We4#xDd=`2(GVf!j8e5G1(g zW4yGNXo`5d@_Zu^kB))L2z*=wttJ8=qNZ(9OXy&Xnq{%lRft_ z7y&Q-c)6;69oV8V9)*@bd$EQS(cxHJ{N;Z`5}eQzDuKC=i2#a_irbbip?D5%sOBB7khlad@ zrI(FHpL>Ziy4*{4s-p6$ZQ8b2-jptH%9OVx&F^{q_o`N19he?itlFHe+B`FssoI)6 zn)Os&`NPRSyyIE-etFZqhqp}`@7AuHvA*QIe!lwl74w6<-kExW4jX7jb?S$e*G zerKvVvue*&dv*l@x^dadcLwmaK8(`z~x>W-%UN0X0cJ^rMV@cTJKw_-SR zYR)8}VE&CuQ`Nz25zYzCP$&2y3s_NH#Zg%fDfW=2frk`lDD=!Q8_l&?Lm@5^28o6; zNq7u0#g!9VVIJ(n&>StaKrPX153Aw}FT8+RD>mRI`1K{~JyYltlfk_86I-3Rk|0XD zgrQYWCbdK_cS509r6{BDA7@s7N!e1MVpu`RFa^z=jWb;WHG?P{R^TRON);Zo`%RDW zujK2*A|*q;>!NnsK*zWckp_qbh5897ePBlB^b*2m(7MwVl9q68A!IpN28DrM4nmSm!1qVgE4s&I(VOl~(G4D}>@8B1jU zP}ji=Td0DpD@{3F3Bl1KLIfq`BPUhnE~SQCwk=txz{7v(n%@RGiF!s?SSVHiKnT=* zpH3p>bS{dL^rcF=hB1iROi_}E`wJ@vkqZ^2K9eS^g$9`}uxdaSgH_|vx))e56+S;?;C@!4KA#0yW=?WUemUWZ~I*5%roC?PtnhJ zUhSIhdZBx^Gws{H;M}eb4Bv5Vm=S0DZcSuccfKmm$+z3@v>wj*y3?(PXV<^g_JgBu zAN|MUQ*Bp|P9J^#_|=osC*SpNT4+6-_H{2fyR(jpsncjHrXB0EO%K1?HP>~!>YXMS z-Psy%d{~4Zx$#0Lji7@RI15 z89rwm@^*EcdBr#*m@f|52Ik0DfhjoqsRm~|7I7?Ayfr=CE z2~UAp@OhZIwKZrgZK7P#O(@y;IqUM=H@ZSDmn+VLeF`k>G1*npyn-Hmx^sx8E{Cr0ZIy#1 z3@q8Ft!ARqK*^PW3z)*P>E)b=ns8hOwhE`qA#K!lM%E1@FsU4yZV6#-*5h zAw&^ZjEJNmnotY_5g~%(T{(KGSSfU;V$0P?B2pC59!sOWshq=dD`wHl)1`vzCmvHEDg%Qw9HtcAxcCCXat?p~@QWxht+j zC!;f_WHjU6lr-M0tR^nGE$?jDHE*16d(D!m+?}-F^;KToKfQm(_71aswte0(cQoVM zk+j`gw?1k7S=E|kXVzZ@M!?z)HwUf_EUw+2Ub}t1^Y-zduI-$*PnoC0yOlMmRWCmI z(%D;Q=Qm|GJ(8*1GiAWanisddwEx!rd1Gedj!fmwDFdLE@;?6v*c>QDW;V?9&eXi? z-?CIjtvgD8Xs606vuig_8K*kG;aIAms@E?0sp<_^6Vr*ASf*;nJ5_sc8w;rIY@~pxyIb%A_|AY$IEm9df!P-@dlIuo%<$d` ze~Q^D%uYjAXaju_D=|BR4alm94HRCoQT2`4s#V$MtxIKg^R9b7X2}e{pd(lKm#kQ_ zQQn6>#Kr69|Fm@M3 zkBk;pZ9I9btNn`*3IvDDLlC6no5P9N5O`oXMq-;WO53N#DKQvk$E(^I;iDWQ02cs@ zYcHw7Y7CUsCxA`0)(=Gb2U_bPzS3GBjwVFDwSEMBG5C9d#kdwp%?EK-9g&>7W+gl} zmKcs>z_jrIGqPpm0z}nd7IvgRw4wzZPz-=8_AzU)SHgs-B}DxIH4^Y_g(#hdk&0Jg z4nV^(GE_2(n*4&}V+xmfg8KX}avm?%9_#y=z-dZeD5nz>^4Yp<`&E&oE&G4S+&H zS)4{_p%+?oLpI+&sM8@|2bu(}=^p3~9Yo`}C}~$#sFZX)bx=2Z^PNE>0Sqz4eFZ&B zn8g~z(kqaS$i5!bhI;C$ocF;hd=0)-CBufBH^fQ5WT<%?nwciK_b$M9D7x0J0VmIj!v9|>lcvWryKQ=8he zfTA|*a8Ge4xzdhJ*>b7a&0K<_J@m;ob<-kZD(t$*EH6u4q>kcBcNpJCpDAU4(?f{l48< ztO}&zzax-sc<80kTce8&2ht4(ZbyE&D${T}*|F$cmv*k3F{YjMum_8t4QbDYneAy$ z!-B0r-G#zyhwOqdwoxb)#=c5jP8obMD@-s$KK_-DiV0QVeJM0l0Wzjio+*r_fE+b+ z3lBpNIP&C5IHUye`4jaRCsYiaB%nnC!C?c}r>e{bsmP;XCNbgh+LGi&pZ;F~1QDGP z5Gi(ft~@>Ybm~Y-ymmC>s!tm3t^sHIjZjj!GCVnaWqfjc!M0<517s@E3i%Q)ju7oM zkP?nAHJIZmP_HN8fkvHcnPo%XZ3ph7lld>1Si=(*8Eo%T4-KIAZrLnbkPf9dpM-t@ zk~GdX=+J8Nzks7F|D}}xuI~KPMxUa_fO~U>^PtfQLC*Liz548JD) z9S8dT|6!kZHUjl%SWjsmK;aHqLCRPkNPlSkKcT<1H2nkS@h?3|9%iZf2~qptJq63C z@3Zigh59xNLr<-EPsn2orm;H_hZkafBqE59&LFa;X*p5QJp(P>V9^rIPB;eJq8LE2 zD-*r}MFG!;0FD3(dWdLEE9O2B7x1tNUbBT*h`9;Zu?sqcg%=>UT1SNixFZZ`bt`MLAy&3m#n zt5faQY%|+uq+5H^HBF1vJJQuVUO$tm-nUrYk*@A|>+}zTZwD8iIJI!*E15@wnd+w& z%AZo<-tjl29Su;vhYT4Ve>1`P+KdOZDV((fjmorxfetD2^u9&krnGO!VwjzUOp}GU_j!*+kt}rKTR;5?&-nV1 zoVu{fPXL;TjQ=9_zHRe;egThfs{OgK1zU|uj6$J<$Bppl9}^y3;1kuoXFwoo)ryB( zvi{fs_%T(IWceP}0)kmJ?m1!0Yn5Dq|1eKdT_%)CI!A^_6ZAM}t9g%MUQgBg+vPIZ ze!c_9Tb~cg43@(>i;{$PI!S{V)GSK*Er$RN?iNY16`=zb@U;rKwmz!^=zy%RC>79C z57T^BPy)Cs$T!^^=#-ro59U3sr5<|}P*?X87OV_o_eH6ol__e1Pgplz5J(fQc&qFz z_2eV(ksREl_ASgm7k*9)m~ege{>X*xn+s=9%lVzqb^*NXIi(FQq`8|ssUWu!{Geuz zcs`*(9^VH1Bu=CHL1zJ=VuUBxLV}dwsh%(dB~{I`sKP`*M~Px3yng#O@}y4q4tDCn zdQ;zUG^!ZmY&_Rx*AB%h@MJE<28vpII0nxr>`&rTVDh3b;LeK$;V%$|L-l*AW6Jc$ z!9@Z~TS`YwsOj%Yw7z7oIh#Ft`kb z*|u9F3u{~EO&QNlVEk_HmDptLM*EEY`pLz*t?9b0@FZgX%TTTb#>qLzLb>l)i3&*)4pc4alzM|@f}Kb z;5*2%$+46loSf%wotxh_FWhc>Z6vel;2qDQ1=}IQHHF84IKm0Yz>7GAWn>!awvf`F z!RK$B@HRYenp%Ww*ynqg5r&`dXoHUGF|Zdv!(&hkcjdrwpYEsBpP{3Na%HdIc!@_6 zdYsmDo^t`$gB_AYu|Zu$al#pdcjYl&8c1-eB0~;2UJ>E1F)Qh^OkxXW2t?cm8MrJz zFj&pDPrN4c{$;n>T)*r#nH!g!l*Rj@!C>CGY=qJhDcOBy|8kSbT>FW$!hD!k>5b6t zTbN?%O(deQ0}QkZe~a15l?#7^!`2dKD?Uw8pCQ2;6l7|Or8gq-@DkkuyZ~H-f(Zej z#1|}>p=4LgalGMQ<^G6mTOZpyxTVfBto+0KiWWeZKW xEZa=B@!G4}CF_ySEIrc lake timeframe partition name +FREQ_TO_TIMEFRAME: Dict[str, str] = { + "day": "1d", + "1d": "1d", + "min": "1m", + "1min": "1m", + "5min": "5m", + "10min": "10m", + "15min": "15m", + "30min": "30m", + "hour": "1h", + "1hour": "1h", + "2hour": "2h", + "4hour": "4h", + "week": "1w", + "1week": "1w", + "month": "1M", + "1month": "1M", +} + +#: bar-field map: qlib field name (without the leading ``$``) -> lake bar column +BAR_FIELD_MAP: Dict[str, str] = { + "open": "o", + "high": "h", + "low": "l", + "close": "c", + "volume": "v", + "vwap": "vw", + "avg_amount": "vw", # amount / volume +} + +#: fields that qlib core/backtest queries but the lake does not store -> all-NaN +UNKNOWN_FIELD_NAMES = ("factor", "change", "trade_unit", "suspend_flag") + +#: columns in the parquet files that are not features +NON_FEATURE_COLUMNS = ("t", "date", "market", "timeframe", "symbol") + +#: Feature-family partitions merged by ``LakeConfig.load_features`` and scanned +#: by the handler's field discovery. ``macro`` holds broadcast market-state +#: columns (see skills/tac-qlib-custom/examples/persist_macro_broadcast.py). +FEATURE_FAMILIES = ("ta", "sp", "macro") + + +def timeframe_for_freq(freq: str) -> str: + """Map a qlib frequency (e.g. ``day``, ``1min``) to a lake timeframe (e.g. ``1d``).""" + f = str(freq).lower() + if f not in FREQ_TO_TIMEFRAME: + raise ValueError( + f"unsupported qlib freq {freq!r}; supported freqs: {sorted(set(FREQ_TO_TIMEFRAME))}" + ) + return FREQ_TO_TIMEFRAME[f] + + +def resolve_lake_root(lake_root: Optional[str] = None) -> Path: + """Resolve the lake root: explicit arg > ``TAC_LAKE_DIR`` (no fallback). + + ``TAC_LAKE_DIR`` is **mandatory** — there is deliberately no default + A missing/empty value raises so a + misconfigured environment never silently points at a wrong directory. + """ + if lake_root is None: + lake_root = os.environ.get("TAC_LAKE_DIR") + if not lake_root: + raise RuntimeError( + "TAC_LAKE_DIR is not set. Point it at the TradeAC lake root, e.g. " + "export TAC_LAKE_DIR=/home/data/lake (docker) or set an absolute " + "path in your local .env." + ) + return Path(str(lake_root)).expanduser().resolve() + + +class LakeConfig: + """Path helpers + cached readers for a (lake_root, market) combination.""" + + def __init__(self, lake_root: Optional[str] = None, market: str = "US"): + self.lake_root: Path = resolve_lake_root(lake_root) + self.market: str = (market or "US").upper() + + # ---- paths -------------------------------------------------------------- + def bar_dir(self, timeframe: str) -> Path: + return self.lake_root / f"market={self.market}" / f"timeframe={timeframe}" + + def bar_path(self, timeframe: str, symbol: str) -> Path: + return self.bar_dir(timeframe) / f"symbol={str(symbol).upper()}.parquet" + + def features_dir(self, timeframe: str) -> Path: + return self.lake_root / "features" / f"market={self.market}" / f"timeframe={timeframe}" + + def features_path(self, timeframe: str, symbol: str) -> Path: + # Legacy flat path (no family tier). Prefer `load_features` which + # resolves the family=ta|sp partition layout. + return self.features_dir(timeframe) / f"symbol={str(symbol).upper()}.parquet" + + def load_features(self, timeframe: str, symbol: str) -> pd.DataFrame: + """All feature columns for a symbol, merging the `family=ta|sp|macro` + partitions by timestamp. Returns an empty frame when no + feature files exist (legacy flat layout falls back transparently).""" + sym = str(symbol).upper() + frames = [] + for family in FEATURE_FAMILIES: + p = self.features_dir(timeframe) / f"family={family}" / f"symbol={sym}.parquet" + if p.exists(): + frames.append(pd.read_parquet(p)) + if not frames: + flat = self.features_dir(timeframe) / f"symbol={sym}.parquet" + if flat.exists(): + return pd.read_parquet(flat) + return pd.DataFrame() + if len(frames) == 1: + return frames[0] + merged = frames[0] + for extra in frames[1:]: + merged = merged.merge(extra, on="t", how="outer", suffixes=("", "_dup")) + for c in [c for c in merged.columns if c.endswith("_dup")]: + merged = merged.drop(columns=c) + return merged + + def calendar_path(self) -> Path: + return self.lake_root / "calendar.parquet" + + def symbols_path(self) -> Path: + return self.lake_root / "symbols.parquet" + + def coverage_path(self) -> Path: + return self.lake_root / "coverage.parquet" + + # ---- metadata readers ---------------------------------------------------- + def load_symbols(self) -> List[str]: + """All symbols known to the lake (from ``symbols.parquet``).""" + p = self.symbols_path() + if not p.exists(): + return [] + df = pd.read_parquet(p) + if "symbol" not in df.columns: + return [] + return sorted(df["symbol"].astype(str).str.upper().tolist()) + + def symbol_spans(self, symbol: str, timeframe: str) -> List[tuple]: + """Listing span(s) ``[(start_iso, end_iso)]`` for a symbol from coverage.parquet.""" + p = self.coverage_path() + if p.exists(): + try: + df = pd.read_parquet(p) + except Exception: # pragma: no cover - defensive + df = pd.DataFrame() + if len(df): + df = df[ + (df.get("market") == self.market) + & (df.get("timeframe") == timeframe) + & (df.get("symbol") == str(symbol).upper()) + ] + if len(df): + row = df.iloc[0] + first = pd.Timestamp(row["first_t"]).date() + last = pd.Timestamp(row["last_t"]).date() + return [(first.isoformat(), last.isoformat())] + # fallback: derive from the bar file itself + p = self.bar_path(timeframe, symbol) + if p.exists(): + import pyarrow.parquet as pq + + tbl = pq.read_table(p, columns=["t"]) + first = pd.Timestamp(tbl.column("t")[0].as_py()).date() + last = pd.Timestamp(tbl.column("t")[-1].as_py()).date() + return [(first.isoformat(), last.isoformat())] + return [("1970-01-01", "2099-12-31")] + + def load_calendar_dates(self) -> List[pd.Timestamp]: + """Trading days (midnight timestamps) for the market, from ``calendar.parquet``.""" + p = self.calendar_path() + if p.exists(): + df = pd.read_parquet(p) + if "date" in df.columns: + if "market" in df.columns: + df = df[df["market"] == self.market] + dates = pd.to_datetime(df["date"]).dt.normalize().sort_values().unique() + return [pd.Timestamp(x) for x in dates] + return [] + + def __repr__(self) -> str: # pragma: no cover + return f"LakeConfig(lake_root={self.lake_root}, market={self.market})" diff --git a/code/tac-qlib/tac_qlib/data/providers.py b/code/tac-qlib/tac_qlib/data/providers.py new file mode 100644 index 0000000..8d0644f --- /dev/null +++ b/code/tac-qlib/tac_qlib/data/providers.py @@ -0,0 +1,230 @@ +"""qlib data providers backed by the TradeAC parquet lake. + +These providers plug into the standard qlib mechanism: ``qlib.init(calendar_provider=..., +instrument_provider=..., feature_provider=...)`` instantiates them and binds them to the +``Cal`` / ``Inst`` / ``FeatureD`` wrappers (see ``qlib.data.data.register_all_wrappers``). +The rest of qlib (``LocalDatasetProvider`` expression engine, backtest ``Exchange``) keeps +working unchanged because the interface contract is identical to the file-based providers: + +- ``feature()`` returns a ``pd.Series`` indexed by the **calendar position** range + ``[start_index, end_index]`` (matching ``FileFeatureStorage.__getitem__`` semantics). +- ``list_instruments()`` returns ``{symbol: [(start, end), ...]}``. +- ``load_calendar()`` returns a list of ``pd.Timestamp`` trading days. +""" + +from __future__ import annotations + +import bisect +from typing import Dict, List, Optional, Union + +import numpy as np +import pandas as pd + +from qlib.data.data import CalendarProvider, FeatureProvider, InstrumentProvider +from qlib.log import get_module_logger + +from .config import ( + BAR_FIELD_MAP, + LakeConfig, + UNKNOWN_FIELD_NAMES, + timeframe_for_freq, +) + +logger = get_module_logger("tac_qlib.data.providers") + + +def _day_freq(freq: str) -> bool: + return str(freq).lower() in ("day", "1d") + + +def _calendar_keys(cal: List[pd.Timestamp], freq: str) -> pd.Index: + """Convert calendar timestamps into the same key space as the lake parquet.""" + if _day_freq(freq): + return pd.Index([pd.Timestamp(x).date() for x in cal]) + return pd.Index([pd.Timestamp(x) for x in cal]) + + +class LakeCalendarProvider(CalendarProvider): + """Trading calendar read from ``/calendar.parquet`` (fallback: derived from bars).""" + + def __init__(self, lake_root: Optional[str] = None, market: str = "US"): + super().__init__() + self.cfg = LakeConfig(lake_root, market) + + def load_calendar(self, freq, future): + timeframe = timeframe_for_freq(freq) + if not _day_freq(freq): + raise NotImplementedError( + f"freq={freq!r} (timeframe={timeframe}) is not supported yet: the lake calendar " + f"only covers daily sessions; add a minute-level calendar to `calendar.parquet`" + ) + + dates = self.cfg.load_calendar_dates() + if not dates: + # Fallback: derive the trading-day set from the persisted bar files. + bar_dir = self.cfg.bar_dir(timeframe) + if bar_dir.exists(): + import pyarrow.parquet as pq + + cal: Dict[pd.Timestamp, None] = {} + for p in sorted(bar_dir.glob("symbol=*.parquet")): + tbl = pq.read_table(p, columns=["t"]) + for v in tbl.column("t"): + cal[pd.Timestamp(v.as_py()).normalize()] = None + dates = sorted(cal.keys()) + if not dates: + return [] + + if future: + # append the next calendar day so that "today" is a valid trade date + last = dates[-1] + dates = dates + [pd.Timestamp(last) + pd.Timedelta(days=1)] + return dates + + +class LakeInstrumentProvider(InstrumentProvider): + """Instruments from ``/symbols.parquet`` with listing spans from ``coverage.parquet``.""" + + def __init__( + self, + lake_root: Optional[str] = None, + market: str = "US", + markets: Optional[Dict[str, list]] = None, + ): + super().__init__() + self.cfg = LakeConfig(lake_root, market) + #: optional named pools, e.g. ``{"sp500": ["AAPL", "MSFT"], "etf": ["SPY"]}``. + #: ``all`` / any unregistered name resolves to every symbol in the lake. + self.markets: Dict[str, list] = markets or {} + + def _resolve_symbols(self, market: Union[str, list]) -> List[str]: + if isinstance(market, (list, tuple, pd.Index, np.ndarray)): + return [str(s).upper() for s in market] + if isinstance(market, str) and "," in market: + return [s.strip().upper() for s in market.split(",") if s.strip()] + if market in self.markets: + return [str(s).upper() for s in self.markets[market]] + return self.cfg.load_symbols() + + def list_instruments(self, instruments, start_time=None, end_time=None, freq="day", as_list=False): + market = instruments["market"] + timeframe = timeframe_for_freq(freq) + + symbols = self._resolve_symbols(market) + if not symbols: + if as_list: + return [] + return {} + + # clip listing spans to the queried window (mirror of LocalInstrumentProvider) + from qlib.data.data import Cal # pylint: disable=C0415 + + cal = Cal.calendar(freq=freq) + start_time = pd.Timestamp(start_time or cal[0]) + end_time = pd.Timestamp(end_time or cal[-1]) + + out: Dict[str, list] = {} + for symbol in symbols: + spans = [] + for begin, end in self.cfg.symbol_spans(symbol, timeframe): + lo = max(start_time, pd.Timestamp(begin)) + hi = min(end_time, pd.Timestamp(end)) + if lo <= hi: + spans.append((lo, hi)) + if spans: + out[symbol] = spans + + filter_pipe = instruments.get("filter_pipe") or [] + for filter_config in filter_pipe: + from qlib.data import filter as F # pylint: disable=C0415 + + filter_t = getattr(F, filter_config["filter_type"]).from_config(filter_config) + out = filter_t(out, start_time, end_time, freq) + + if as_list: + return list(out) + return out + + +class LakeFeatureProvider(FeatureProvider): + """Feature data from the lake parquet (OHLCV bars + pre-computed ta-lib features). + + Field routing: + - ``$open/$high/$low/$close/$volume/$vwap`` -> bar parquet columns + - ``$amount`` (= v*vw), ``$avg_amount`` (= vw) -> derived from bar parquet + - ``$factor/$change/...`` -> all-NaN (not stored) + - anything else -> a ta-lib column in the features parquet + """ + + def __init__(self, lake_root: Optional[str] = None, market: str = "US"): + super().__init__() + self.cfg = LakeConfig(lake_root, market) + self._bar_cache: Dict[tuple, pd.DataFrame] = {} + self._feature_cache: Dict[tuple, pd.DataFrame] = {} + + # ------------------------------------------------------------------ caches + def _load_bar_df(self, instrument: str, timeframe: str) -> pd.DataFrame: + key = (instrument, timeframe) + if key not in self._bar_cache: + p = self.cfg.bar_path(timeframe, instrument) + self._bar_cache[key] = pd.read_parquet(p) if p.exists() else pd.DataFrame() + return self._bar_cache[key] + + def _load_feature_df(self, instrument: str, timeframe: str) -> pd.DataFrame: + key = (instrument, timeframe) + if key not in self._feature_cache: + self._feature_cache[key] = self.cfg.load_features(timeframe, instrument) + return self._feature_cache[key] + + @staticmethod + def _keys(df: pd.DataFrame, freq: str) -> pd.Index: + ts = pd.to_datetime(df["t"]) + return ts.dt.date if _day_freq(freq) else ts + + # ------------------------------------------------------------------ fields + def _extract(self, instrument: str, field: str, timeframe: str, freq: str) -> Optional[pd.Series]: + """Return the field as a Series keyed by date/timestamp (None if not present in the lake).""" + bar = self._load_bar_df(instrument, timeframe) + + if field in BAR_FIELD_MAP: + col = BAR_FIELD_MAP[field] + if col in bar.columns: + return bar[col].astype(float).set_axis(self._keys(bar, freq)) + return None + if field == "amount": + if "v" in bar.columns and "vw" in bar.columns: + return (bar["v"] * bar["vw"]).astype(float).set_axis(self._keys(bar, freq)) + return None + if field in UNKNOWN_FIELD_NAMES: + return None + + feat = self._load_feature_df(instrument, timeframe) + if field in feat.columns: + return feat[field].astype(float).set_axis(self._keys(feat, freq)) + return None + + # ------------------------------------------------------------------ api + def _get_calendar(self, freq: str) -> List[pd.Timestamp]: + from qlib.data.data import Cal # pylint: disable=C0415 + + cal = Cal.calendar(freq=freq) + return list(cal) + + def feature(self, instrument, field, start_index, end_index, freq): + field = str(field)[1:] + timeframe = timeframe_for_freq(freq) + + cal = self._get_calendar(freq) + n = len(cal) + lo = max(0, int(start_index)) + hi = min(n - 1, int(end_index)) + if lo > hi: + return pd.Series(dtype=np.float32) + + keys = _calendar_keys(cal[lo : hi + 1], freq) + ser = self._extract(str(instrument).upper(), field, timeframe, freq) + if ser is None: + vals = np.full(len(keys), np.nan, dtype=np.float64) + else: + vals = ser.reindex(keys).to_numpy(dtype=np.float64) + return pd.Series(vals, index=pd.RangeIndex(lo, hi + 1))