quant_test / factor_engine /gp /qlib_engine.py
lucky-loster's picture
Upload folder using huggingface_hub
590a501 verified
Raw
History Blame Contribute Delete
6.54 kB
"""
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"]),
)