Files
tac-exp-dev/code/tac-qlib/tac_qlib/contrib/data/handler.py

255 lines
9.9 KiB
Python

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