atelier / src /forecasting /config.py
adam2504's picture
Initial upload — Atelier luxury demand intelligence demo
89ea727 verified
Raw
History Blame Contribute Delete
1.42 kB
"""
TFT model configuration.
Centralises all hyperparameters so the Streamlit app and training script share one source of truth.
"""
from dataclasses import dataclass, field
from typing import List
@dataclass
class TFTConfig:
# Temporal structure
input_size: int = 60 # lookback window (days)
h: int = 28 # forecast horizon (days)
# Architecture
hidden_size: int = 64
n_head: int = 4
attn_dropout: float = 0.1
dropout: float = 0.1
# Training
max_steps: int = 500
learning_rate: float = 1e-3
batch_size: int = 32
val_check_steps: int = 50
# Quantiles for probabilistic output (confidence bands)
loss_quantiles: List[float] = field(
default_factory=lambda: [0.1, 0.5, 0.9]
)
# Covariates
# Static: features that don't change over time per series
stat_exog_list: List[str] = field(
default_factory=lambda: ["category_enc", "store_enc"]
)
# Historic: known past values
hist_exog_list: List[str] = field(
default_factory=lambda: ["price", "event_flag"]
)
# Future: known future values (calendar events we schedule in advance)
futr_exog_list: List[str] = field(
default_factory=lambda: ["event_flag", "day_of_week", "month", "is_fashion_week"]
)
# Persistence
model_dir: str = "models"
model_name: str = "tft_luxury"
DEFAULT_CONFIG = TFTConfig()