""" Qlib-backed tensor data engine for GP factor mining. Replaces the legacy parquet TensorDataEngine: all OHLCV data is loaded via qlib D.features(), then converted to GPU tensors for fast GP evaluation. """ from __future__ import annotations from typing import Any import numpy as np import pandas as pd import torch from data_pipeline.load_data import load_market_features, pivot_to_symbol_time from factor_engine.gp.operators import safe_div, safe_log, ts_delay_raw, ts_future_raw class QlibTensorDataEngine: """Load qlib market data and expose tensors compatible with GP operators.""" PRICE_FIELDS = ["开盘价", "收盘价", "最高价", "最低价"] VOLUME_FIELDS = ["成交量", "成交额"] def __init__( self, instruments, start_time: str, end_time: str, freq: str = "day", forward_steps: int = 5, splits: dict[str, dict[str, str]] | None = None, device: str = "cpu", qlib_fields: list[str] | None = None, ): self.device = device self.forward_steps = forward_steps self.splits = splits or {} if qlib_fields is None: qlib_fields = ["$open", "$close", "$high", "$low", "$volume", "$vwap", "$amount"] print(f"[QlibTensorDataEngine] Loading {freq} data {start_time} ~ {end_time} via qlib...") raw_df = load_market_features( instruments, qlib_fields, start_time=start_time, end_time=end_time, freq=freq, rename_for_gp=True, ) self.times = sorted(raw_df.index.get_level_values("datetime").unique()) self.symbols = sorted(raw_df.index.get_level_values("instrument").unique()) self.tensors: dict[str, torch.Tensor] = {} for field in self.PRICE_FIELDS + self.VOLUME_FIELDS + ["vwap"]: if field in raw_df.columns: pivot = pivot_to_symbol_time(raw_df, field) pivot = pivot.reindex(index=self.symbols, columns=self.times) self.tensors[field] = torch.tensor(pivot.values, dtype=torch.float32, device=device) self._build_derived_features() self._build_target() self._build_masks() self._print_summary() def _build_derived_features(self): close = self.tensors["收盘价"] open_ = self.tensors["开盘价"] high = self.tensors["最高价"] low = self.tensors["最低价"] volume = self.tensors["成交量"] if "vwap" not in self.tensors or torch.isnan(self.tensors["vwap"]).all(): self.tensors["vwap"] = close.clone() vwap = self.tensors["vwap"] safe_volume = torch.where(volume > 1e-8, volume, torch.ones_like(volume)) self.tensors["return_1"] = safe_div(close, ts_delay_raw(close, 1)) - 1.0 self.tensors["return_2"] = safe_div(close, ts_delay_raw(close, 2)) - 1.0 self.tensors["return_4"] = safe_div(close, ts_delay_raw(close, 4)) - 1.0 self.tensors["oc_return"] = safe_div(close, open_) - 1.0 self.tensors["co_return"] = safe_div(open_, close) - 1.0 self.tensors["hl_spread"] = safe_div(high - low, close) self.tensors["ho_gap"] = safe_div(high - open_, open_) self.tensors["lo_gap"] = safe_div(low - open_, open_) self.tensors["vwap_close_gap"] = safe_div(vwap, close) - 1.0 self.tensors["log_volume"] = safe_log(volume) if "成交额" in self.tensors: amount = self.tensors["成交额"] self.tensors["amount_per_volume"] = safe_div(amount, safe_volume) self.tensors["log_amount"] = safe_log(amount) def _build_target(self): close = self.tensors["收盘价"] entry = ts_future_raw(close, 1) exit_ = ts_future_raw(close, self.forward_steps) target_valid = (~torch.isnan(entry)) & (~torch.isnan(exit_)) & (entry > 1e-8) & (exit_ > 1e-8) target = torch.full_like(close, float("nan")) target[target_valid] = exit_[target_valid] / entry[target_valid] - 1.0 target[:, -self.forward_steps:] = float("nan") target = torch.where( (target > -0.30) & (target < 0.30), target, torch.full_like(target, float("nan")), ) self.tensors["target_return"] = target def _build_time_mask(self, start: str, end: str) -> torch.Tensor: times_series = pd.Series(self.times) entry_times = times_series.shift(-1) exit_times = times_series.shift(-self.forward_steps) start_ts, end_ts = pd.Timestamp(start), pd.Timestamp(end) mask = ( (times_series >= start_ts) & (times_series <= end_ts) & (entry_times >= start_ts) & (entry_times <= end_ts) & (exit_times >= start_ts) & (exit_times <= end_ts) ).fillna(False).values return torch.tensor(mask, dtype=torch.bool, device=self.device).unsqueeze(0) def _build_masks(self): default_splits = { "train": {"start": "2010-01-01", "end": "2021-12-31"}, "valid": {"start": "2022-01-01", "end": "2023-12-31"}, "test": {"start": "2024-01-01", "end": "2025-12-31"}, "holdout": {"start": "2026-01-01", "end": "2026-12-31"}, } merged = {**default_splits, **self.splits} for name, span in merged.items(): tensor = self._build_time_mask(span["start"], span["end"]) self.tensors[f"{name}_mask"] = tensor setattr(self, f"{name}_mask", tensor) def _print_summary(self): print(f" Stocks={len(self.symbols)}, TimeSteps={len(self.times)}, Device={self.device}") for split in ("train", "valid", "test", "holdout"): key = f"{split}_mask" if key in self.tensors: print(f" {split}: {self.tensors[key].sum().item()} bars") def get_data(self, name: str) -> torch.Tensor: return self.tensors[name] def export_factor_panel(self, tree) -> pd.DataFrame: """Export a single GP factor tree as a qlib-compatible MultiIndex DataFrame.""" factor = tree.evaluate(self).detach().cpu().numpy() dates_col = np.repeat(self.times, len(self.symbols)) symbols_col = np.tile(self.symbols, len(self.times)) return pd.DataFrame( {"factor": factor.flatten(order="F")}, index=pd.MultiIndex.from_arrays([symbols_col, dates_col], names=["instrument", "datetime"]), )