| """ |
| 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"]), |
| ) |
|
|