Skip to content
Open
Show file tree
Hide file tree
Changes from 5 commits
Commits
Show all changes
16 commits
Select commit Hold shift + click to select a range
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
8 changes: 8 additions & 0 deletions pytorch_forecasting/models/autoformer/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
"""
Autoformer model for time series forecasting.
"""

from pytorch_forecasting.models.autoformer._autoformer_pkg_v2 import Autoformer_pkg_v2
from pytorch_forecasting.models.autoformer._autoformer_v2 import Autoformer

__all__ = ["Autoformer", "Autoformer_pkg_v2"]
155 changes: 155 additions & 0 deletions pytorch_forecasting/models/autoformer/_autoformer_pkg_v2.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,155 @@
"""
Packages container for Autoformer model.
"""

from pytorch_forecasting.base._base_pkg import Base_pkg


class Autoformer_pkg_v2(Base_pkg):
"""Autoformer package container."""

_tags = {
"info:name": "Autoformer",
"info:compute": 2,
"authors": ["harshsomankar123-tech"],
Comment thread
harshsomankar123-tech marked this conversation as resolved.
"capability:exogenous": True,
"capability:multivariate": True,
"capability:pred_int": True,
"capability:flexible_history_length": True,
"capability:cold_start": False,
}

@classmethod
def get_cls(cls):
"""Get model class."""
from pytorch_forecasting.models.autoformer._autoformer_v2 import Autoformer

return Autoformer

@classmethod
def get_datamodule_cls(cls):
"""Get the underlying DataModule class."""
from pytorch_forecasting.data.data_module import TslibDataModule

return TslibDataModule

@classmethod
def _get_test_datamodule_from(cls, trainer_kwargs):
Comment thread
harshsomankar123-tech marked this conversation as resolved.
Outdated
"""Create test dataloaders from trainer_kwargs - following v1/v2 pattern."""
from pytorch_forecasting.data.data_module import TslibDataModule
from pytorch_forecasting.tests._data_scenarios import (
data_with_covariates_v2,
make_datasets_v2,
)

data_with_covariates = data_with_covariates_v2()
data_loader_default_kwargs = dict(
target="target",
group_ids=["agency_encoded", "sku_encoded"],
add_relative_time_idx=True,
)

data_loader_kwargs = trainer_kwargs.get("data_loader_kwargs", {})
data_loader_default_kwargs.update(data_loader_kwargs)

datasets_info = make_datasets_v2(
data_with_covariates, **data_loader_default_kwargs
)

training_dataset = datasets_info["training_dataset"]
validation_dataset = datasets_info["validation_dataset"]

context_length = data_loader_kwargs.get("context_length", 8)
prediction_length = data_loader_kwargs.get("prediction_length", 2)
batch_size = data_loader_kwargs.get("batch_size", 2)

train_datamodule = TslibDataModule(
time_series_dataset=training_dataset,
context_length=context_length,
prediction_length=prediction_length,
add_relative_time_idx=data_loader_kwargs.get("add_relative_time_idx", True),
batch_size=batch_size,
train_val_test_split=(0.8, 0.2, 0.0),
)

val_datamodule = TslibDataModule(
time_series_dataset=validation_dataset,
context_length=context_length,
prediction_length=prediction_length,
add_relative_time_idx=data_loader_kwargs.get("add_relative_time_idx", True),
batch_size=batch_size,
train_val_test_split=(0.0, 1.0, 0.0),
)

test_datamodule = TslibDataModule(
time_series_dataset=validation_dataset,
context_length=context_length,
prediction_length=prediction_length,
add_relative_time_idx=data_loader_kwargs.get("add_relative_time_idx", True),
batch_size=batch_size,
train_val_test_split=(0.0, 0.0, 1.0),
)

train_datamodule.setup("fit")
val_datamodule.setup("fit")
test_datamodule.setup("test")

train_dataloader = train_datamodule.train_dataloader()
val_dataloader = val_datamodule.val_dataloader()
test_dataloader = test_datamodule.test_dataloader()

return {
"train": train_dataloader,
"val": val_dataloader,
"test": test_dataloader,
"data_module": train_datamodule,
}

@classmethod
def get_test_train_params(cls):
"""
Return testing parameter settings for the trainer.
"""
from pytorch_forecasting.metrics import SMAPE

params = [
# First set: smaller network params for fast testing
dict(
hidden_size=16,
n_heads=2,
e_layers=1,
d_layers=1,
d_ff=32,
Comment thread
harshsomankar123-tech marked this conversation as resolved.
),
# Second set: custom moving_avg and logging metrics
dict(
hidden_size=8,
n_heads=2,
e_layers=1,
d_layers=1,
d_ff=16,
moving_avg=5,
logging_metrics=[SMAPE()],
),
# Third set: custom scheduler
dict(
hidden_size=8,
n_heads=2,
e_layers=1,
d_layers=1,
d_ff=16,
optimizer="adamw",
lr_scheduler="cosine_annealing",
lr_scheduler_params={"T_max": 5},
),
]

default_dm_cfg = {"context_length": 8, "prediction_length": 2}

for param in params:
current_dm_cfg = param.get("datamodule_cfg", {})
default_dm_cfg.update(current_dm_cfg)

param["datamodule_cfg"] = default_dm_cfg

return params
Loading
Loading