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