Tập dữ liệu cửa sổ cho huấn luyện và suy luận mô hình giao dịch tuần tự
Tóm tắt
Tài liệu này mô tả các tiện ích tập dữ liệu để chuẩn bị dữ liệu bảng chuỗi thời gian cho huấn luyện mô hình và suy luận cuốn chiếu. Tập dữ liệu huấn luyện trả về các chuỗi trượt gồm đặc trưng, lợi suất tương lai, thang đo biến động và mặt nạ, với độ dài chuỗi, giới hạn ngày và bước nhảy có thể cấu hình. Giới hạn ngày được ánh xạ sang các ngày có trong bảng dữ liệu; hàm khởi tạo từ chối những khoảng không thể cung cấp chuỗi đầy đủ.
Thay vào đó, tập dữ liệu suy luận tạo các cửa sổ kết thúc tại những chỉ số thời gian đã chọn và trả về đặc trưng, mặt nạ cùng chỉ số tương ứng cho mỗi bước dự báo. Các tiện ích siêu dữ liệu tĩnh gán mã định danh tài sản dạng số nguyên và có thể tùy chọn mã hóa nhóm tài sản cùng chi phí giao dịch thành tensor. Các giao diện này làm rõ đầu vào theo thời gian và ngữ cảnh tài sản cho mô hình phía sau. Trích đoạn nêu hành vi triển khai và hình dạng đầu vào, nhưng không có mô hình, kết quả giao dịch, quy trình xác thực hay thảo luận về rò rỉ và sự phụ thuộc giữa các cửa sổ chồng lấn; những thuộc tính đó phụ thuộc vào cách bên gọi tạo và sử dụng bảng dữ liệu.
Ý chính
- Mẫu huấn luyện là các cửa sổ trượt chứa chuỗi đặc trưng, lợi suất, thang đo biến động và mặt nạ.
- Giới hạn ngày và bước nhảy kiểm soát các cửa sổ huấn luyện được đưa vào.
- Cửa sổ suy luận kết thúc tại các chỉ số thời gian đã chỉ định và trả về những chỉ số đó cùng đầu vào.
- Siêu dữ liệu tĩnh có thể mã hóa định danh tài sản, nhóm và chi phí giao dịch.
- Các tiện ích tập dữ liệu xác định hình dạng dữ liệu nhưng không chứng minh hiệu suất hay tính hợp lệ của mô hình giao dịch.
Thẻ
Toàn văn
# 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,
)
```Hiển thị toàn văn kèm ghi nguồn theo giấy phép của tài liệu gốc. Giấy phép: MIT
Bản tóm tắt này do tác nhân nghiên cứu của Stratmill biên soạn từ tài liệu gốc; đây không phải bản sao của tài liệu.