Skip to content
Open
Show file tree
Hide file tree
Changes from 7 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 80 additions & 3 deletions pytorch_forecasting/data/data_module/_tslib_data_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import torch
from torch.utils.data import DataLoader, Dataset

from pytorch_forecasting.adapters import ScalerAdapter
from pytorch_forecasting.data.encoders import (
EncoderNormalizer,
NaNLabelEncoder,
Expand Down Expand Up @@ -341,7 +342,19 @@ def __init__(

self._metadata = None

self.scalers = scalers or {}
# wrap in the unified adapter so we speak one fit/transform interface
self._target_normalizer = (
ScalerAdapter(self._target_normalizer)
if self._target_normalizer is not None
else None
)
self._scalers = {
name: ScalerAdapter(scaler) for name, scaler in (scalers or {}).items()
}
self._target_normalizer_fitted = False
self._feature_scalers_fitted = False
self._preprocess_cache = {}

self.shuffle = shuffle

self.continuous_indices = []
Expand Down Expand Up @@ -537,6 +550,55 @@ def metadata(self) -> dict[str, Any]:
self._metadata = self._prepare_metadata()
return self._metadata

def _fit_target_normalizer(self, train_indices):
"""Fit the target normalizer on the training targets only."""
if self._target_normalizer is None or self._target_normalizer_fitted:
return
targets = [self.time_series_dataset[idx.item()]["y"] for idx in train_indices]
if not targets:
return
self._target_normalizer.fit(torch.cat(targets, dim=0))
self._target_normalizer_fitted = True

def _fit_scalers(self, train_indices):
"""Fit each named continuous-feature scaler on the training data only."""
if not self._scalers or not self.continuous_indices:
return
names = self.time_series_metadata["cols"]["x"]
for orig_idx in self.continuous_indices:
name = names[orig_idx]
if name not in self._scalers:
continue
column = [
self.time_series_dataset[idx.item()]["x"][:, orig_idx]
for idx in train_indices
]
if not column:
continue
self._scalers[name].fit(torch.cat(column, dim=0))
self._feature_scalers_fitted = True

def _normalize_target(self, target):
"""Apply the fitted target normalizer (no-op until fitted)."""
if self._target_normalizer is None or not self._target_normalizer_fitted:
return target
return self._target_normalizer.transform(target)

def _normalize_features(self, continuous):
"""Apply fitted continuous-feature scalers (no-op until fitted).

``continuous`` columns are ordered by ``self.continuous_indices``.
"""
if not self._feature_scalers_fitted or not self.continuous_indices:
return continuous
names = self.time_series_metadata["cols"]["x"]
out = continuous.clone()
for pos, orig_idx in enumerate(self.continuous_indices):
name = names[orig_idx]
if name in self._scalers:
out[:, pos] = self._scalers[name].transform(continuous[:, pos])
return out

def _preprocess_data(self, idx: torch.Tensor) -> list[dict[str, Any]]:
"""
Process the the time series data at the given index, before feeding it
Expand All @@ -559,9 +621,14 @@ def _preprocess_data(self, idx: torch.Tensor) -> list[dict[str, Any]]:
- Splits data into categorical and continuous features, which are grouped based on the indices.
""" # noqa: E501

series = self.time_series_dataset[idx]
# cache: a series is transformed once, not per window
i = idx.item() if isinstance(idx, torch.Tensor) else idx
if i in self._preprocess_cache:
return self._preprocess_cache[i]

series = self.time_series_dataset[i]
if series is None:
raise ValueError(f"series at index {idx} is None. Check the dataset.")
raise ValueError(f"series at index {i} is None. Check the dataset.")
target = series["y"]
features = series["x"]
timestep = series["t"]
Expand Down Expand Up @@ -594,6 +661,10 @@ def _preprocess_data(self, idx: torch.Tensor) -> list[dict[str, Any]]:
else torch.zeros((features.shape[0], 0))
)

# apply fitted scalers / normalizers (no-op until fitted)
continuous_features = self._normalize_features(continuous_features)
target = self._normalize_target(target)

res = {
"features": {
"categorical": categorical_features,
Expand All @@ -611,6 +682,7 @@ def _preprocess_data(self, idx: torch.Tensor) -> list[dict[str, Any]]:
if target_scale:
res["target_scale"] = target_scale

self._preprocess_cache[i] = res
return res

def _create_windows(self, indices: torch.Tensor) -> list[tuple[int, int, int, int]]:
Expand Down Expand Up @@ -719,6 +791,11 @@ def setup(self, stage: str | None = None) -> None:
self._train_size + self._val_size : total_series
]

self._preprocess_cache = {}
if stage is None or stage == "fit":
self._fit_target_normalizer(self._train_indices)
self._fit_scalers(self._train_indices)

if stage == "fit" or stage is None:
if not hasattr(self, "_train_dataset") or not hasattr(self, "_val_dataset"):
self._train_windows = self._create_windows(self._train_indices)
Expand Down
172 changes: 172 additions & 0 deletions pytorch_forecasting/data/tests/test_tslib_data_module.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,39 @@
import numpy as np
import pandas as pd
import pytest
from sklearn.preprocessing import StandardScaler
import torch

from pytorch_forecasting.adapters import ScalerAdapter
from pytorch_forecasting.data.data_module import TslibDataModule
from pytorch_forecasting.data.encoders import TorchNormalizer
from pytorch_forecasting.data.timeseries import TimeSeries


def _make_ts(n_series: int = 20, length: int = 40, offset: float = 100.0) -> TimeSeries:
"""合成数据集:连续特征 ``x`` 远离 0(~offset),便于看出标准化;目标 ``y`` 是正弦。"""

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

please use english docstrings

rows = []
for i in range(n_series):
for t in range(length):
rows.append(
{
"series_id": i,
"time_idx": t,
"x": offset + 10.0 * np.sin(t / 5.0) + i,
"y": np.sin(t / 5.0),
}
)
df = pd.DataFrame(rows)
return TimeSeries(
data=df,
time="time_idx",
target="y",
group=["series_id"],
num=["x"],
unknown=["x"],
)


@pytest.fixture(scope="session")
def sample_timeseries_data():
"""Fixture to generate a sample TimeSeries."""
Expand Down Expand Up @@ -530,3 +557,148 @@ def test_multivariate_target():
assert (
y.shape[-1] == 2
), "Target should have two dimensions for n_features for multivariate target."


def test_init_wraps_scalers_in_adapter_and_sets_flags():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
scalers={"x": StandardScaler()},
target_normalizer=TorchNormalizer(),
batch_size=8,
)
assert isinstance(dm._scalers["x"], ScalerAdapter)
assert isinstance(dm._target_normalizer, ScalerAdapter)
assert dm._feature_scalers_fitted is False
assert dm._target_normalizer_fitted is False
assert dm._preprocess_cache == {}


def test_fit_scalers_standardizes_train_feature():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
scalers={"x": StandardScaler()},
batch_size=8,
)
train_idx = torch.arange(len(ds)) # 本单测在全部序列上 fit
dm._fit_scalers(train_idx)

assert dm._feature_scalers_fitted is True
# transform 原始 train 列,检查已标准化(~0 均值,~1 标准差)
names = dm.time_series_metadata["cols"]["x"]
oi = dm.continuous_indices[names.index("x") if "x" in names else 0]
raw = torch.cat([ds[i.item()]["x"][:, oi] for i in train_idx], dim=0)
scaled = dm._scalers["x"].transform(raw)
assert abs(float(scaled.mean())) < 1e-3
assert abs(float(scaled.std()) - 1.0) < 1e-2


def test_fit_target_normalizer_sets_flag():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
target_normalizer=TorchNormalizer(),
batch_size=8,
)
dm._fit_target_normalizer(torch.arange(len(ds)))
assert dm._target_normalizer_fitted is True


def test_normalize_features_scales_only_configured_columns():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
scalers={"x": StandardScaler()},
batch_size=8,
)
# before fit: no-op
raw_cont = ds[0]["x"][:, dm.continuous_indices]
assert torch.equal(dm._normalize_features(raw_cont), raw_cont)

# after fit: scaled
dm._fit_scalers(torch.arange(len(ds)))
scaled = dm._normalize_features(raw_cont)
assert not torch.equal(scaled, raw_cont)
# magnitude drops from ~100 toward ~0
assert abs(float(scaled.mean())) < abs(float(raw_cont.mean()))


def test_normalize_target_is_noop_until_fitted():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
target_normalizer=TorchNormalizer(),
batch_size=8,
)
tgt = ds[0]["y"].float()
assert torch.equal(dm._normalize_target(tgt), tgt) # not yet fitted
dm._fit_target_normalizer(torch.arange(len(ds)))
assert dm._normalize_target(tgt).shape == tgt.shape


def test_preprocess_data_scales_and_caches():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
scalers={"x": StandardScaler()},
batch_size=8,
)
dm._fit_scalers(torch.arange(len(ds)))

out1 = dm._preprocess_data(0)
cont = out1["features"]["continuous"]
# x was originally ~100; after scaling it should be near 0
assert abs(float(cont[:, 0].mean())) < 5.0
# cache: second call must return the exact same object
out2 = dm._preprocess_data(0)
assert out1 is out2


def test_setup_fit_produces_scaled_samples():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
scalers={"x": StandardScaler()},
target_normalizer=TorchNormalizer(),
batch_size=8,
)
dm.setup(stage="fit")
assert dm._feature_scalers_fitted is True
assert dm._target_normalizer_fitted is True

# retrieve encoder continuous features; x should be standardized
x, _ = dm.train_dataset[0]
hist = x["history_cont"]
assert abs(float(hist[:, 0].mean())) < 5.0 # unscaled value was ~100


def test_no_scalers_leaves_data_untouched():
ds = _make_ts()
dm = TslibDataModule(
time_series_dataset=ds,
context_length=16,
prediction_length=4,
batch_size=8,
)
dm.setup(stage="fit")
out = dm._preprocess_data(0)
raw_cont = ds[0]["x"][:, dm.continuous_indices].float()
# when no scaler is configured, continuous features must be byte-identical to raw
assert torch.allclose(out["features"]["continuous"], raw_cont)
# target_scale is not produced (out of scope for this PR)
assert "target_scale" not in out
Loading