From 2c2684b103198ce529e71c078a9a93e04f6b113d Mon Sep 17 00:00:00 2001 From: zhaoli Date: Sat, 15 Aug 2026 05:14:48 +0000 Subject: [PATCH] start experiment 16 (16-scheduled-algo-retrain-on-20260814-tacrd) --- code/MANIFEST.txt | 22 ++ code/tac-qlib/tac_qlib/contrib/__init__.py | 11 + .../tac_qlib/contrib/data/__init__.py | 3 + .../tac-qlib/tac_qlib/contrib/data/handler.py | 236 ++++++++++++++++++ .../tac_qlib/contrib/model/__init__.py | 4 + .../tac_qlib/contrib/model/rank_ensemble.py | 189 ++++++++++++++ .../tac_qlib/contrib/model/rank_gbdt.py | 200 +++++++++++++++ .../tac_qlib/contrib/strategy/__init__.py | 3 + .../tac_qlib/contrib/strategy/optimal_stop.py | 217 ++++++++++++++++ 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 -> 9752 bytes .../__pycache__/providers.cpython-312.pyc | Bin 0 -> 13488 bytes code/tac-qlib/tac_qlib/data/config.py | 175 +++++++++++++ code/tac-qlib/tac_qlib/data/providers.py | 231 +++++++++++++++++ 15 files changed, 1316 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/data/__init__.py 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/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/optimal_stop.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..bbcf811 --- /dev/null +++ b/code/MANIFEST.txt @@ -0,0 +1,22 @@ +# TradeAC custom-qlib-code snapshot (auto-generated) +# parent repo HEAD : HEAD +unknown +# parent repo date : unknown +# copied paths: +# tac-qlib/tac_qlib/contrib +# tac-qlib/tac_qlib/data +# per-file hashes (git hash-object): + 1b6298c4a5652f2e863cbdc385a1014a570fcd59 tac-qlib/tac_qlib/contrib/__init__.py + c76a9f17f680e74eea766eff27f7624359749ed6 tac-qlib/tac_qlib/contrib/data/__init__.py + 871ff1e163c29261f140c3f53d42a41e6504c779 tac-qlib/tac_qlib/contrib/data/handler.py + b151d139a0dcde87d74b21e7c4b729176ba5c39b tac-qlib/tac_qlib/contrib/model/__init__.py + d3f051f3a8650c42fedc7b367b966f7c74fb5789 tac-qlib/tac_qlib/contrib/model/rank_ensemble.py + ccfe7d554989aa7f3e5a2128ae663e51b2207149 tac-qlib/tac_qlib/contrib/model/rank_gbdt.py + 4afcf9058231111c412925f4c4b84e81d656db87 tac-qlib/tac_qlib/contrib/strategy/__init__.py + 79aaad9e39fcc740a773f4f63c512ce1086cfde0 tac-qlib/tac_qlib/contrib/strategy/optimal_stop.py + 92e6e90eb0cd0a25142034560f27adb6b705b1a8 tac-qlib/tac_qlib/data/__init__.py + 6d5b9cca6970261fed1f633f6b588ec3e2b399bb tac-qlib/tac_qlib/data/__pycache__/__init__.cpython-312.pyc + ba4f5d1741728ad1e364e537c7b50488e9112b9f tac-qlib/tac_qlib/data/__pycache__/config.cpython-312.pyc + 608c88f0f45378b170bddb5b301cb32fad9abb35 tac-qlib/tac_qlib/data/__pycache__/providers.cpython-312.pyc + 686d36f6d101c547491ca866aa143aa542e17518 tac-qlib/tac_qlib/data/config.py + d9f839be30026f337754a3f015425a8efdbe8e2a 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/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/handler.py b/code/tac-qlib/tac_qlib/contrib/data/handler.py new file mode 100644 index 0000000..871ff1e --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/data/handler.py @@ -0,0 +1,236 @@ +"""TACHandler: a qlib DataHandlerLP that builds datasets from the TradeAC lake. + +This is the "custom DataHandler" entry point (Option B): the handler is referenced from the +workflow yaml's ``dataset.handler`` and reads OHLCV + pre-computed ta-lib features straight +from the lake parquet files through ``QLibDataLoader`` + the tac_qlib feature provider. + +The standard qlib processor pipeline (``infer_processors`` / ``learn_processors``) still runs +on top, so existing recipes such as ``DropnaLabel``, ``CSZScoreNorm`` or ``RobustZScoreNorm`` +keep working unchanged. +""" + +from __future__ import annotations + +import os +from inspect import getfullargspec +from typing import List, Optional, Tuple, Union + +from qlib.data.dataset import processor as processor_module +from qlib.data.dataset.handler import DataHandlerLP +from qlib.utils import get_callable_kwargs + +from ...data.config import ( + LakeConfig, + timeframe_for_freq, + NON_FEATURE_COLUMNS, +) + +DEFAULT_INFER_PROCESSORS = [ + {"class": "DropAllNaN", "kwargs": {}}, + {"class": "ProcessInf", "kwargs": {}}, + {"class": "ZScoreNorm", "kwargs": {}}, + {"class": "Fillna", "kwargs": {}}, +] +DEFAULT_LEARN_PROCESSORS = [ + {"class": "DropnaLabel"}, + {"class": "CSZScoreNorm", "kwargs": {"fields_group": "label"}}, +] + +#: always include raw OHLCV; ta-lib columns are discovered from the lake and appended. +RAW_FEATURE_FIELDS = ("$open", "$high", "$low", "$close", "$vwap", "$volume") + +DEFAULT_LABEL = "Ref($close,-2)/Ref($close,-1)-1" + + +def check_transform_proc(proc_l, fit_start_time, fit_end_time): + """Port of ``qlib.contrib.data.handler.check_transform_proc`` (inject fit window into procs).""" + new_l = [] + for p in proc_l: + if not isinstance(p, processor_module.Processor): + klass, pkwargs = get_callable_kwargs(p, processor_module) + args = getfullargspec(klass).args + if "fit_start_time" in args and "fit_end_time" in args: + assert fit_start_time is not None and fit_end_time is not None, ( + "Make sure `fit_start_time` and `fit_end_time` are not None." + ) + pkwargs.update({"fit_start_time": fit_start_time, "fit_end_time": fit_end_time}) + proc_config = {"class": klass.__name__, "kwargs": pkwargs} + if isinstance(p, dict) and "module_path" in p: + proc_config["module_path"] = p["module_path"] + new_l.append(proc_config) + else: + new_l.append(p) + return new_l + + +def get_common_feature_fields(lake_root=None, market="US", timeframe="1d") -> List[str]: + """Discover ta-lib columns present in *every* features parquet file of the lake. + + 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 + + common = None + for p in sorted(feat_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 sorted(common) if common else [] + + +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 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", "get_common_feature_fields"] + + +# Make `DropAllNaN` resolvable by bare name from processor configs (e.g. the default +# ``infer_processors`` and workflow yamls that reference it without a ``module_path``), +# mirroring how qlib registers its own processors in ``qlib.data.dataset.processor``. +processor_module.DropAllNaN = DropAllNaN 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/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..ccfe7d5 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/model/rank_gbdt.py @@ -0,0 +1,200 @@ +"""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 _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. + """ + if group is None or len(group) == 0: + return 0.0 + offs = np.concatenate([[0], np.cumsum(group.astype(int))]) + vals = [] + for i in range(len(group)): + s = slice(offs[i], offs[i + 1]) + p, l = preds[s], labels[s] + if len(p) < 3 or np.std(p) == 0 or np.std(l) == 0: + continue + vals.append(np.corrcoef(pd.Series(p).rank(), pd.Series(l).rank())[0, 1]) + return float(np.mean(vals)) if vals 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..4afcf90 --- /dev/null +++ b/code/tac-qlib/tac_qlib/contrib/strategy/__init__.py @@ -0,0 +1,3 @@ +from .optimal_stop import OptimalStopControl # noqa: F401 + +__all__ = ["OptimalStopControl"] 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/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..6d5b9cca6970261fed1f633f6b588ec3e2b399bb GIT binary patch literal 522 zcmaKozfQw25XS8!O`4QSKoK(=$^d-`@7Jp9jb*!yM6?n1L!nc^+vnas;kOLv=<$mM$@WDYsYFaPXy~pDj4)K z2|dZK)PiJ#j9)Y0+8$(<<)N*XCT~&B(wNFanO!F_lWN(h&2*5w S2V3t*<6|x;S+|~?*F`@H0Em15 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..ba4f5d1741728ad1e364e537c7b50488e9112b9f GIT binary patch literal 9752 zcmc&aZA={3b~CfH-wO*YA8S6$_!HLV!+>o*{E=Y5v13?&VEZ*B&M@8?VDWzC&H!fL zI(47&1FMb$x*K zvkS(~Q(vTx&An&tJ@?#u&pG#;bI>yTAzvqo%D`;a~A7;;3NL(VLoYseiYcpGnj zi{$NE>EN9JyR_2HdjKxr3!(J#MNk&=B~beKQYg!`vYf8~*w62RvXZZYvRW%^_}u{4 z;&RBtGeQBs=Oiqh-^=gg_wxt%gM1zQtAEEnRH*d^wBAEQUjFb<5r2elc!M4)=3Dqy zzVQv~Pzgo6K=4g(5Pb88l+J-6-?SyrHt~CSy^1(q=GC8kp?Ezz)UtUF({|Dt!zZ#qzNR5%{rl-%6*K)HAccW znQ)8`hd5x4*uYGMd4U;;qbMh3DQ?};jyuNGbNhMLuD4=7W#_Q421bl>>@mG+J<`A+ zG0e6cZD5+3nzrpU#6^S{&!O$S%iAlfjTqrksukpZpAN^U!;_@sciQQ?`S@aTG- z4Grd28}zLNm`EIyhzB;t`1q8VrElnijwC*VJV=}qMbOnKCrScxeTxr7fK=UFEEbnE zy%&SZ2?V3Z07C<esJqhl&PB?z|wZ3@sd&>D@$ zq%mk2P;D(5oGir4d|8wb2vkBh)Z(*O`Yy9qFS1w95A>bA(ml|pl5v$BQ^|-*hE#G= zr6#8Wc9o7N1X%W1cyvs)MB*?-C=wS1)ixQAOhg5ho}A(ms*9T(Ww~g4A||N?XS%Pj zXV3Tb_p$@smjbkE8{tBr+Nv!y#>GYj)ukPCb|My*RF60z!UQ}!65&PzR+R)+;uEIY zbor`IQ^2ZC7hVWVv-x0CZXpPA^REGzCnU2kMSPfr$Dx1w^I^8(4w)h)+UT7mQJrKr z-gq)ayiM_xCss%JawSq7Og!aO!RM&MbG9mOKgry{EWB^H!0+dSQVJbNyV#mNRj}J6+ zi9|Ct-FO}huo?hDxe8-yQCm-amIf@ZR$( zj$L=VGj_*si}uL&J?n-3OyT~et4rt74bT1B_G|IC7hX|bxS>4%YWg`g?HS&*P+r#t zK{;JN*a%O}!e!ZAo2F|wfhw^@IjqnC!Y4QX8Gv~rMcl2GFp~C+;|*~On+Y4}^A$miM^`p_Zf z>$?x}MBuQYix)uLV45W%GR*)vUKrsfBD#m_W};zH1moT;L=)09Gl{i{L0niAL`DSl z)g}ZSBYFTW6=IWN6puxPn8d^cF#3!b1}`E3^+Y@zlOT+M?J!e-13n{+1Z+a4bFeL`_MbcQdC6vLeYhbif2mAx2XY%QsY95P6g=S5ywZm7> z$3wRS6ky`8RB#VmjNxvAG$teg!W(G}d^j@=<{JDD_(P@%HrNz!s8n23?b*$$mQevr z-W9Nt5SD1i4PirqPC%i$U=_gq2~j|*9pVNo=oM0R;@MfeSP>gE2A#pk#aRT_>_r!$ z{YCf_uR?K`*eE0%#c8{Ly|jG(#km*fgLA>9(@JUUjPpx(>3U)121OOrZCYr2lxR3g zF2Ymv&@y}etwP0JyXHP5yALVuhBfz5*?m-TA5YWAQ7;eieolv`=c0ujtOL;fc z9ted!P#5s!d;!#jd<9y9q zLe-lLYsf|*0SIcNH^VW_wKoMr*-dFX!!Z1jo7c4s=D8QsDx?TD=X9vsq^R8zcCsYK zNrw^U9Mi-K6;!Efg8^C-APT}nsuhem*yeVCafBfXkrAy!YZ38aIKa9hI9N6ej+AAS zz8ou>j7Anv#l2A6CDv~puQz_Dbu^9_Kwr2y_$1c-HoLZrB`vVwr8+AaD^*Z`i z?GS^rU<#5Yc{|Ov@ff)tin~O{>wEC}z1NrOmiK>rc;)bNO`2(!z3rdA_}TPlx6)l# zWbc)<P)Elqq+fl+4j7(-S;P?r#G#-nC$+3=%7rVA7rq zzLO3k91U2|4VXx?2k06=E0!EDP$x8O6F%NOHa=T~Z#0f#{6;8ro}})j7S1c)gUjuI z*S0$OwX4)O4UCPY3clu#Sq&P~;TBC3_DoPpwM7`uQh6UyY_4du4C$ z7nSL&H`0Tzre9!XFPnC-|Ihk>@JR*nV?EA^9}UCnR*@H&o7Lp*--!*lCyGa!d|2}` zvX@!frFajo+W*P&>DhmDuXUW4JI*T|7t-zhUz|%fU6Z}n(vE97AM!OHn*`8wa+3kz z+2V%-P6d-i+wX&D07kBmnkwfzu>u28*WvuVZ0SkB@XZn^Oqr&cR@vK{cC`Kk%k^%% z6KbnPV=!I*T0L1TeX04{m){}cQQKD9B70lXj+UQbrA47_x8t3cT412Wk%m-i_rrki*v968T%Ud87;Zm=rd?yjZb!VM+hBh7%s-4xJ!KH zDO&T?$etR-Q#(UuyhRUg-@AQ3nemm)oL?{SK49;$3zwCG+8HY2_ARtN>RRlQ-3Rn; zZ`th41^--Rtvn!?2R_=TlsB!Fx60+MYvmnsd52QowN~COmv<}WJ&LzC?da7t(DbLU z0IVaL)rru5fCgGpmS<_;EQwssAJagKA=5jwWZu(K=4Bpzo1#(4Jis^CE>I^1=C!eWNKyZ%#4?9n-X3PO*QD^bMlm(rC`qTCn@l>HpZHF za{W(YrL4Ra;~Cr=uX+;4Sj!#TkJ7UJ!L;lt+jz~>=Stb1)jL9sl7l#;t$mV03PM&R z4}RxIGz)#s)4V-W6fM5vAc>S?{J;+MEWBfs{K#q2bp8kzoa1#*QsEs2hjqD(2cE<+ z_-tVQQL^~4Rk*OE{^(Ww1XASPSeE+;$@Cff3FIg9&%5jCyt+fbM9h-+A5RgtEiVyM zWWbXQ;$$w)Rf}*TuNMPgyI-ysCHS7fhDFHy<9sCmfg6xo(QlQyU4?8O7<$9|8bN)N zc5BnFGn%k&XfJ8Wi`1bekfMapBbhaO9JkkRHldjey zot=#>t&MFhfie^WEQI4C^cpVk{abbR-3|#_0$O$ULOSj&{_X)uTOo$dVXwBRJ-AnT$+oLxSJ8k*6I#QwQ?ex zm=OYMWPSymfGQuc;R4W<W0x)~Wp%R@BoB*x zv#s~1GG!I;fByR1^@Vdv*@5NKkE>RyPK~pbxV5| z56$+hSMPq*zu3R*TYhO}w^H4)R((pYKBZJYH`}*fRrBcN;>l&p^0}1)rRw-v)d{)k zgi>{Cwl`C=_Yu3uDm6{3^sfp(DO8R;pRVpsSDi_ho&i_=VBp@sLSX5-;%P8 zp8Bs!O6MEr8kflBp5(&Ob<>8_0RP$`97$8uX4HP{o2)Hc~9Hgo=$mBr?TgS zQhIW&v|BFiR!Vzj`qq7A^C#y{E?GWsz2{nPd*8cC$xPcC(K z*E7?bDXYXBR>}@8(;pYE6e{%{>C(=$uPg28$`m!Ey^U!{{Egopt_88HwY2AOqc=lNl1onZw!~;46(RYZt zQK7{VxoCce$RQUs%~AJv(J<-`JxJhNKnSv9_-zU(p<$MWhnqC}rRU5%C=9&YW*;@ns`aaGDb7M>SHKfNs4S_@ z82;WG2`2^BjlX}eS_)fKZSau_zR#dP19AbYZq;(O%!+A86)iHZ!-Cy~raItU+EgyU@srXyh{KsVfu{Hm5 zvi~{7-@WD^ko^OS|Kghey6nHM`2S?q`en;;rK0_lvgPh~s7Ibf&(h_0y{lzl2h&CE z8E^57^MAg#5dIDkADe$^ukP{FTh1xwi59`*-yY1vFDF3z0owlHtQhp?{l4J%EKj!*grN zNgNjFd8!c%X(tWeU-g*eHiq8B#b4rL5sDQ(!Jr)tbQ!neLd);H0Kgl>mvrz4nzA0< zbdlCUa??gyFKgwnR$kZ2Hr(8bi?&UN!`eq~mR4JvHcz^&B|rE{>qSx*5T*?K_+?!){DS8AAm&Xiq1=U?HsA z2u!pdEX0B}%V7O3+y%;Xi;bIJz==eSim*V!T8)X7CaH}>L_tx;Eu7zw4K=hwBNUf zJAe>mPugk!>PlSf?ep7jcYohw_g$CEPC@w2zh4}DtDd5Mfgf5hs|kAPpCNIL;;10S z(VQVl#|%LOjcH@l7&OAKDQb$DgXWkeXo*>a)|f45%axY}?WCR~=!9Q$)D?3F-7!zl z6DtpvlRB1YMa&!YLfXpNqQ00v=#NzfD`Qo`s#tZfnx+gC@8!x~qE%{wHJq2TUp576 zIUnbM-xUVxFvU49Q=E%;zhunST1j%g_1rp=TMoGuyhqPnHD(O>*=^ADd^FO_ zaIC~KLqcLC!tsL0^s?b|Jje8oG136foDx`$Z)<0USmFFIFELT}9N%KMpBmsrzM#=i zbhw|1#H9pjBuZ?YV+D>O{m1z502_~pvE58h59YK);t{Da%tm>v71D;@*V59`WRJu} zNf?gtacOx0)5o*Yu)r_N4fOQDs5qPyVI^L~(PIn@&h$p&T<%jP%HGq{&PJhrD^PJ1 z+T@a*YPTKmc|l-@hL8n~BF}5IB6nIyB=G$aQR0OV8;ypvx;;IC7UTva@S?;d`qU{m z_VjcmV1gacSmdQAa?FJ>_|YM#DMk`;hL870;(QZfIM6P^h>wmUtNMB9l{v@rL!$jc zLO2(R_cOzBwFG9yhuL9ZKCTXMoEQ4oFwcY&aYdY9s`Zp#x0}?5eX9Jw3TT5-bWG9T z#nc78D7NE6n9oKP^XWJw0!F30J-6Sr&7gR4v)41cxGmMaujRow6~L2FEWr&&`A{^` z-_HxbLTG>yU5DCEggTEt+SL&{*7k&A?}D@0o{0BF`W4^l?#H{2f1^8B+}(EU(UXc- zf=TrWY>W@}C4^9)z@Hzlme_EJ97?jshcp-n=irCH3xW*%(Myj&;uLj)djVWaz=#R`qsD-A0=#*Fk(Ix7R!=03#*>&~6T+%2X)mDHuOW!XE#vPgrT6mC0|cV$p)?G z!7?o`|EqCr+Hj@81U*jZwNFw7{eBf4)xQlLFjj5xSYZ%(IQj*MwN!u}KLR^!gcl@6 zJ4FocDsof+d2|T?n+yP!k^d&(Ix7a*77O;<%@Sb6<0?mN8j{i zYxceAn{T_}y;*m?Zian+%|dN+rndS0+SWIGnVNlB&;GP+{~d?x>LZhn%!RZ5Epy@7 zitnDdc4qp_%-ZKdv*C6FH=GI9;Jg8I4BmyPmXdCV2qze zVT?eevtHm?j_DH;F#xT-D4O@TYGo}siH)e+2NDa#`fdg!d}IWmoV4j>1;Ey@zPEu+ zBpSZyr4C3uN5gr5Uo}h`W$HViB5C`Rrlg$0L{TygdpdWjWGbZhI7ZFUh>d1(7zC?g z3juydq)*xZR1xbS`VByV%{6Jw*lKRuR^0KeN=0U?vc89trtjE<6;NT>{B<<= z9?YgR8cZ1Tr%={0LiFaj!O=4Pyn!>w#^)){_-AJAwEsRWQ-dY&;LA{8!X%r-e|kVW z+ztEd=0we=Ie51WvRg7P|Hp{WQpVU>JZZD?7krSrd>J`-(>4(etQqbQH=+DkkvQ=8CRU5>0EDExHa7|Gb zd}?#k`WsyP;5vJu=*lA|>PoMPDww%-aQ%bZX#J8^0lYs+$p&G&Y~pMhB+qD|^|l)f|W7ei~yO72v_qJ#I!}xo_ObH0mOK-}r?9Du1ApFmT$25`qNh ze2kZN6HO6MR~|PK@#q++jDX`JXf+Y|5LKmA+{3UO$FNK+5+9cM<|seHNAvAKpX|90 zU<5q)DMc^es{G42`kS=7K)BH3xr+MI@Y1 zO#RVBuVP2U3`uNnlvgY&X#(ypVKt6rX2lRH74d`+W22FAUNNC)QRNcHMx<4?&UEvEEOYQ@wQW;vi!17GHeYX^rRUq{ zx2Kx3D|St_FRo&42CoNajo%NY=*9NKDdUYpH@mKP%~rfpJy-owZMJUfLfyVh-M(zy zf%j?-eq^LpH{UZ;D-O}0*r?hyDe>m$_0i|Yvo)Kh%pdq_7OQKg%Kp_+ao1Oq@=otc zb}ag9Z=AS!=K7i0`s}LK*Vet>^lH=FZMU~{W>$5k>keo9hm(&kdi+Tz!TUKww_-Sb za?T{6VE)3Tsp??12xo<6$P@e!1*|Bp;;1Z#6njY1z(a~N6gof5Msp?BP>4%}L875d z5*~v{aplBTm&L}S_FkPLV+Tpk26LI8=-s9IUC?dT`bf=RE)8z3Wo^AbM4_V4XVPS~&>+(VRt?BvuxdP7^#Ti~Ld(nl zYEl@iidCvr0hS9|b|lEOgq$ZqAMFX$nMNp~OQr{PEkvtDf$L3^h6OZjMw(4*fD~TiW$@l z(CW?NP&5J!06!x`0y!~6ayyBv%0WfIs2(lESA`arn{wqL@QEd&BYY@Vd%WR6C#$F; z2i=M9Lxcu@4noTAA04;+>zIMUmtvRFul6%w&TzTOrOkn7abK-r_ff+IMyyUJ^V`7T-O^_?=?aD zwoKE(<>RKCpix`e*|zBMp>4Y0sn2-ov!0D<+s4m7c2NFppZ}Y$8~)lvR5RXWI%-Tm zG*))3GX1cYf%xxMG9BwoVf~mx9o=mCdjb(M&5_MFVJSFvChZ7KSYaappcDsLSHnS} zw}}?nED#kp0%|U20N=MEW}GOaDNy1|xGiMU%lQ%Pl0)mg6wJs5z-Tk{WakWcNOZ~! zU(pV+U2SJx*3Jm#i$k`7Ir1$41!q4s;et8pFr(1|pjYC!%T8H0QJ{x26m9L6>_;;aSrTLfUw78S4r~&Iv&%Cl1=8v%IgY~xJ-Nl3-%zWL@lKrwzHFdOE4wf*m zWS=&hiAn<{R{||y3d^RKaw2NNaRp!%c9%n1udR%%8%AJKIZ!R>Rs}E$IPeY{+rbTl z@GGEUuQojEQ0Jl&JX#!Ju&|=3cm> zh$}`U(hyB3hJlC>!S=2kzf`Iex>K>`awHMCii*wsK{ZRtk(NTg`vbwbxMB=?7M~{O z$znxa>d8#S#-t^0sND5dEqEI;-iBFA*4vu2e&DHs|3tfA^i?L0erlo08t%F)u0|)L zGp1xT>)wzw-l?o6F1bzbt=lngoNs&8lC9jCwBPYnUfVmpcgFS}vvsz8-Y|DK>)V#J z-CeUbY5aNBs$}P)zY2_i)$48!Tpw6iy*0CX>wM=MM}M}ubJ{*-o)Yg=)}&UvaN@-? zx6aIO$ZmKfTe)k>fSENfYRf~o9UgYdEdY3 zUKzFKF#WNeDz98zy?)9#)%homdlgjm>U(~wdfm0ebYdo!t=jfp)$TV8NypDuH71XJ z=xqLF_3G65ndW&{X5Id5^?`-z$1~NB|KqW2bufA41MiB|)*1Wkk$Eob-80z^BdiX5 zD16UG)i*9ytypZ{e6P%I-f`E*+%v;3=*SiRdsa-@ zDDOj`V%buS)!guF?*_9^wU@BFF#PEu2RQW7(~wncFL7m@y}&+lMTWfI_x{J{d0NhMJut2@a`^( z9vSUgwek3ouJ*6PRUkNI9)e3czB!zT4S@%SVJtD{t@Q(u{(;tdxL;|l4@VOs-&#L{z8L(yz~a3Y3C;U)RGpBVyJmTR zY%DPx#|zWOeay(FkqdCE2BWYe{h?(g;DBPFyJ8=+*7x#0;nosv{eU!b;n@nebQ(n} zU4=1p3Lfm{Znm3gB>pXE$#a2%4%zsaMD7*&C|r$xWxx(K#;P-hs|(^M)7ZR>=w&Wq z+*C%iS9>5)WUu}=T;$8`RPu$UCJCC=YAI>8+C%6C3$x(?9xg+2gsb1eCTK4rrND|p zi?y!A9@Wj&o7m_fh<@`xpZOKr;`W_M(-qI+u02W9`?i(jgBgJ>KV1?|WRm6EQf4(euazA|VeKp}>> zub^fLvsmk~^bBMp!q{5GhXLozy&+WgVcZ zbmDdGoa&bVD^1rR1fmxh!96LUGe|^Uk!`OEZ6{n6gQFA2;QGdylB-@&<1era4x&rc zqNg^sD~+Od(czxrP;zA)8y3s`s+VwWrhM&8U^bX7e-^~WZ=Iepzn%}nR0Lw z&u_|>?}J^jZ~qUTeDle-f-eX!Uc7a2{`~ERT|YVT(=$Il^X}8hjs@qMjC0M5G2^U< zeX-zKm+`Ed*_!b*q-_oA&M1TegfIfOqmU@Pk1GWjGWdy>nT3Y@om~E=PN)KpP@(Y* zJ!2^4EyMdAknM(U;bEu&?m==)9a4haLlBi62UHB4B%qZ97uE)@PgUa$QjwR#Ozx@2 zYfCa8?-^c$9z=AI!0olmbM@)Tr&EVg;`PH>SAEiOXB9a3Z-kP<)#1tEtK*a7Y1_8> zbr7jcE5u5i9AVliU?m_hHJHOW$kzklvQbA;vuwz_62YZ)JpUyV3w^>OgFm3u;|Hj{ zQ#Q*Myy+NNEFU+4zx@m$Qd6<)#IG$ zfo1z-GiOAK;n#$}ZeP~aX7^C|6@C|n^cNIvTW`46T4EBv>X=6}FE{*4pl zv6!ke5mgl4c(9E6W(;3`sPDz_I;<7X33(vGG7Qo#KofOCR^8X{Wwin$L&7(BLu7j_{Q?g51xSOuNb!uKEo z>*OI!PeTNCkw_vn)#GUd*$6KX5>a@01gf=&+mr}Uh*TMB3-P1mP2G5DttnUf*U(UO z!wI`Y-F15x+-oxKHDISI6I8|F)Ff*_Ew6mM%ZG=&N3IlusRb&w(bN z-hPIg4Zm`3?p$W$uEm;_srKu(nXNO@tzDU#riJQlnd)t?oz7P8S*Y&FRCm06>PNwM zg6SttrcZw}`)Dv*{ZzXADYf6*{)UXB0n&F7kkJ`A6P&NjdO*X&Q9IDk%{UmSkTOs2 zS@3Ph_%^)2XMIfzzU>*`_IVCHt#A8*-W^?8-?6mw7$E!oSPqb+!V<15Jg>#4wN#9r zfXf;{Yu0e-`?o0r3!lDI3M~07fy4&OijYWHX>g zHSa~yROH->yYT&?{;X(xZAqjqtdfS_L_xUP3zNz+S#?rPLl^KOX2e%u+ z=${abF7VUp9z9?XwPM+0GFg9s0eDQ6WLdt3wSZw(4Lv7pc}_5jNk?RO z$U%>T2A%g}=5=SizFjVp?dLjxz4iHvqQP>QXHk&QcR<1*UbGg4{F1|^5v~>qvK4y= zD&T7cpte4$1NeZ@R}>2Psb3v^*+mJUE66w9L+O;Am-gqqzoj0a6mVDfT^7s?-xCyt zf>Nfa2tHxnct9XbxZuJyTEK+s%hyLPEZCHNJ752J$+!3$Qn{xXfQQBlOZ22FgZ zzp4mK8!fdr>2GGvqE z5fR?PsH7V;i6t042Y=!oh`^2dk-=)VedaZp_b$23=6Z}8@3|?9_hW;>ynV?Cse2@4 z_nG}mO(t{gXU+=qL0aWE_IBUI6jN^^5e3FFP%8WYqvbOf{uZ09CJtbHE~7q;f_E_p zYKo;dBJ%J$-U3(v?n=Rg9iYS)P$mi}*%dQ6sp@IBz;y=4iUbE#X^!$zKzELSo*czS zZeQRzynx08KIJ0UEksAeqo`_uD0~!)`e>fKGgHx{9xHquBb@6TrOq3#))m`c^`-rO z;a&Jdz!W8j?ip#C{>b2@O`q0L^!i^=u79CgvQ*12DDTfH-+xd~WT+=Tq&EKAWTp*E z6ht4nDf`Ow%5B-Q?UyV|7MgBZvYF^@pSai4mFblaeL}(SCtdVD+LvC{`UwTUpByxt Qpy`8js`XQfBnVai1x>3EVgLXD literal 0 HcmV?d00001 diff --git a/code/tac-qlib/tac_qlib/data/config.py b/code/tac-qlib/tac_qlib/data/config.py new file mode 100644 index 0000000..686d36f --- /dev/null +++ b/code/tac-qlib/tac_qlib/data/config.py @@ -0,0 +1,175 @@ +"""TradeAC lake configuration helpers. + +The lake is a hive-partitioned parquet store (see ``tac-engine/skills/tradeac-lake``): + + $TAC_LAKE_DIR/ + ├── market=US/ + │ └── timeframe=1d/ + │ └── symbol=AAPL.parquet # OHLCV bars: t, date, o, h, l, c, v, n, vw + ├── features/ # ta-lib indicators, wide format + │ └── market=US/ + │ └── timeframe=1d/ + │ └── symbol=AAPL.parquet # t, sma_5, sma_20, rsi_14, ... + ├── calendar.parquet # trading days per market + ├── coverage.parquet # per (market,timeframe,symbol) loaded windows + └── symbols.parquet # asset master +""" + +from __future__ import annotations + +import os +from pathlib import Path +from typing import Dict, List, Optional + +import pandas as pd + +#: qlib freq string (Freq.__str__) -> 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") + + +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: + return self.features_dir(timeframe) / f"symbol={str(symbol).upper()}.parquet" + + 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..d9f839b --- /dev/null +++ b/code/tac-qlib/tac_qlib/data/providers.py @@ -0,0 +1,231 @@ +"""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: + p = self.cfg.features_path(timeframe, instrument) + self._feature_cache[key] = pd.read_parquet(p) if p.exists() else pd.DataFrame() + 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))