237 lines
9.1 KiB
Python
237 lines
9.1 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 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
|