用于序列交易模型训练与推理的窗口数据集
代码 《交易机器学习》
总结
本文介绍用于准备时间序列面板以供模型训练和滚动推理的数据集工具。训练数据集返回由特征、前向收益、波动率尺度和掩码构成的滑动序列,并支持配置序列长度、日期范围和步长。日期边界会映射到面板中可用的日期;若指定范围无法提供完整序列,构造函数会拒绝该范围。
推理数据集则生成以选定时间索引为终点的窗口,并为每个预测步骤返回特征、掩码及对应索引。静态元数据工具会分配整数资产标识,并可选择将资产组和交易成本编码为张量。这些接口明确了下游模型使用的时间输入和资产背景。摘录说明了实现行为和输入形状,但未提供模型、交易结果、验证流程,也未讨论信息泄漏和重叠窗口依赖;这些属性取决于调用方如何构造和使用面板。
核心观点
- 训练样本是滑动窗口,包含特征、收益、波动率尺度和掩码序列。
- 日期范围和步长决定纳入哪些训练窗口。
- 推理窗口以指定时间索引为终点,并将这些索引与输入一同返回。
- 静态元数据可以编码资产身份、资产组和交易成本。
- 数据集工具定义了数据形状,但不能证明交易模型的表现或有效性。
标签
全文
# dataset.py
```py
"""Torch datasets for DeePM-style windowed training and inference."""
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
import numpy as np
import pandas as pd
import torch
from torch.utils.data import Dataset
from .features import FeaturePanel
@dataclass(frozen=True, slots=True)
class StaticAssetMetadata:
"""Static per-asset metadata used as context."""
assets: list[str]
asset_ids: torch.Tensor # (N,)
group_ids: torch.Tensor | None # (N,)
costs: torch.Tensor | None # (N, 1)
def build_static_metadata(
assets: Sequence[str],
*,
asset_to_group: dict[str, str] | None = None,
asset_to_cost_bps: dict[str, float] | None = None,
) -> StaticAssetMetadata:
"""Create tensors for asset id, group id, and costs."""
assets_list = [str(a) for a in assets]
n = len(assets_list)
asset_ids = torch.arange(n, dtype=torch.long)
group_ids: torch.Tensor | None = None
if asset_to_group is not None:
groups = [str(asset_to_group.get(a, "UNKNOWN")) for a in assets_list]
unique_groups = {g: i for i, g in enumerate(sorted(set(groups)))}
group_ids = torch.tensor([unique_groups[g] for g in groups], dtype=torch.long)
costs: torch.Tensor | None = None
if asset_to_cost_bps is not None:
cost_vals = [float(asset_to_cost_bps.get(a, 0.0)) / 10000.0 for a in assets_list]
costs = torch.tensor(cost_vals, dtype=torch.float32).unsqueeze(-1)
return StaticAssetMetadata(
assets=assets_list, asset_ids=asset_ids, group_ids=group_ids, costs=costs
)
class DeepmWindowDataset(Dataset[tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]):
"""Sliding-window dataset for training.
Each item returns (x_seq, y_seq, v_seq, m_seq) of shapes
(L, N, F), (L, N), (L, N), (L, N).
"""
def __init__(
self,
panel: FeaturePanel,
*,
seq_len: int,
start_date: pd.Timestamp | None = None,
end_date: pd.Timestamp | None = None,
stride: int = 1,
) -> None:
if seq_len <= 1:
raise ValueError("seq_len must be > 1")
self._panel = panel
self.seq_len = int(seq_len)
dates = panel.dates
start_idx = 0
end_idx_exclusive = len(dates)
if start_date is not None:
start_idx = int(dates.get_indexer([pd.Timestamp(start_date)], method="bfill")[0])
if end_date is not None:
end_idx_exclusive = (
int(dates.get_indexer([pd.Timestamp(end_date)], method="ffill")[0]) + 1
)
t_max_start = (end_idx_exclusive - 1) - self.seq_len
t_min_start = start_idx
if t_max_start < t_min_start:
raise ValueError("Date range too small for given seq_len")
self.start_indices = np.arange(t_min_start, t_max_start + 1, stride, dtype=np.int64)
def __len__(self) -> int:
return int(self.start_indices.shape[0])
def __getitem__(
self, idx: int
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
s = int(self.start_indices[idx])
e = s + self.seq_len
return (
torch.from_numpy(self._panel.x[s:e]),
torch.from_numpy(self._panel.y_fwd1[s:e]),
torch.from_numpy(self._panel.vol_scale[s:e]),
torch.from_numpy(self._panel.mask[s:e]),
)
class RollingWindowInferenceDataset(Dataset[tuple[torch.Tensor, torch.Tensor, int]]):
"""Dataset for batched rolling-window inference."""
def __init__(self, panel: FeaturePanel, *, seq_len: int, start_t: int) -> None:
if seq_len <= 1:
raise ValueError("seq_len must be > 1")
if start_t < seq_len - 1:
raise ValueError("start_t must be >= seq_len - 1")
self._panel = panel
self.seq_len = int(seq_len)
self.times = np.arange(start_t, len(panel.dates) - 1, dtype=np.int64)
def __len__(self) -> int:
return int(self.times.shape[0])
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor, int]:
t = int(self.times[idx])
s = t - self.seq_len + 1
e = t + 1
return (
torch.from_numpy(self._panel.x[s:e]),
torch.from_numpy(self._panel.mask[s:e]),
t,
)
```在遵守原作品许可的前提下,附作者信息全文展示。 许可协议: MIT
此摘要由 Stratmill 研究智能体根据原文撰写,并非原文副本。