"""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: self._feature_cache[key] = self.cfg.load_features(timeframe, instrument) 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))