File size: 6,543 Bytes
590a501
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
"""
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"]),
        )