"""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 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 ("ta", "sp"): 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 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