File size: 6,661 Bytes
878bbb6
81f5bb2
878bbb6
81f5bb2
 
878bbb6
 
81f5bb2
878bbb6
 
 
81f5bb2
878bbb6
81f5bb2
878bbb6
81f5bb2
 
 
 
 
878bbb6
 
 
 
 
 
 
 
 
 
 
81f5bb2
 
 
878bbb6
 
81f5bb2
 
 
878bbb6
 
 
 
81f5bb2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
878bbb6
 
 
 
81f5bb2
 
878bbb6
 
 
 
81f5bb2
 
878bbb6
81f5bb2
878bbb6
 
 
 
 
 
81f5bb2
878bbb6
 
 
81f5bb2
878bbb6
 
 
 
 
 
81f5bb2
 
 
 
 
 
878bbb6
 
 
 
 
 
81f5bb2
 
 
 
 
 
 
 
878bbb6
 
 
 
 
81f5bb2
878bbb6
 
 
 
81f5bb2
878bbb6
81f5bb2
 
878bbb6
 
 
81f5bb2
878bbb6
 
 
 
 
 
81f5bb2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
878bbb6
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
import pandas as pd
import matplotlib.pyplot as plt
import lightning.pytorch as pl
from lightning.pytorch.callbacks import EarlyStopping, ModelCheckpoint
from lightning.pytorch.callbacks import LearningRateMonitor
from pytorch_forecasting import TimeSeriesDataSet, TemporalFusionTransformer
from pytorch_forecasting.data import GroupNormalizer
from pytorch_forecasting.metrics import QuantileLoss, MAE, SMAPE, RMSE
import torch

def train_multivariate():
    print("=== Advanced Deep Learning Pipeline (TFT Edition) ===")
    
    # 1. Load the advanced feature dataset
    print("Loading prepared dataset...")
    # Adjusting path to match typical execution from project root
    try:
        df = pd.read_csv("data/processed/dl_advanced_features_data.csv")
    except:
        df = pd.read_csv("../data/processed/dl_advanced_features_data.csv") # Fallback
    
    # Add time index required by PyTorch Forecasting
    df["date"] = pd.to_datetime(df["date"])
    df = df.sort_values(["Mandi", "Commodity", "date"])
    
    # Create purely sequential integer time index per group
    df["time_idx"] = df.groupby(["Mandi", "Commodity"]).cumcount()
    
    # Ensure all targets are positive for relative scaling
    df["target_price"] = df["target_price"].clip(lower=1.0)
    
    # Create string format for groupings
    df["Mandi"] = df["Mandi"].astype(str)
    df["Commodity"] = df["Commodity"].astype(str)

    max_prediction_length = 14 # Predict next 14 days
    max_encoder_length = 60    # Elevated context lookback (60 days)
    
    # Time-based exact split (NO random)
    training_cutoff = df["time_idx"].max() - max_prediction_length

    print("Configuring TimeSeriesDataSet...")
    
    # Unknown reals that depend strictly on historic values
    unknown_reals = [
        "Arrivals", "temp_avg", "temp_max", "temp_min", "rainfall", "humidity", "solar_radiation", "wind_speed",
        "modal_lag_1", "modal_lag_3", "modal_lag_7", "modal_lag_14", "modal_lag_30", "modal_lag_60",
        "arrivals_lag_1", "arrivals_lag_7", "arrivals_lag_30", "arrivals_lag_60",
        "rolling_mean_7", "rolling_mean_30", "rolling_mean_60",
        "rolling_std_7", "rolling_std_30", "rolling_std_60",
        "volatility_7", "volatility_30", "momentum_7", "momentum_14",
        "arrival_change_7", "price_change_1",
        "temp_anomaly", "rain_anomaly",
        "price_x_arrivals", "roll_mean_x_volatility", "rainfall_x_arrivals", "temp_x_price",
        "high_price_regime", "high_arrival_regime", "volatility_regime",
        "sudden_spike_flag", "arrival_shock_flag",
        "price_residual_30", "normalized_dev_30"
    ]
    # Filter to only columns that actually exist to avoid crash if some are all-NaN/dropped
    unknown_reals = [f for f in unknown_reals if f in df.columns]

    # Define exact TimeSeries boundaries
    training = TimeSeriesDataSet(
        df[lambda x: x.time_idx <= training_cutoff],
        time_idx="time_idx",
        target="target_price",
        group_ids=["Mandi", "Commodity"],
        min_encoder_length=max_encoder_length // 2,
        max_encoder_length=max_encoder_length,
        min_prediction_length=max_prediction_length,
        max_prediction_length=max_prediction_length,
        static_categoricals=["Mandi", "Commodity"],
        time_varying_known_reals=["time_idx", "day_of_year", "day_of_week", "month", "sin1", "cos1", "sin2", "cos2"],
        time_varying_unknown_reals=["target_price"] + unknown_reals,
        target_normalizer=GroupNormalizer(
            groups=["Mandi", "Commodity"], transformation="softplus"
        ),  
        add_relative_time_idx=True,
        add_target_scales=True,
        add_encoder_length=True,
    )

    # 2. Creating dataloaders for the GPU/MPS
    validation = TimeSeriesDataSet.from_dataset(training, df, predict=True, stop_randomization=True)
    batch_size = 128  
    train_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=4)
    val_dataloader = validation.to_dataloader(train=False, batch_size=batch_size * 2, num_workers=4)

    # 3. Architect TFT (Temporal Fusion Transformer)
    print("Instantiating Multi-variate TFT Model for exact external shock modeling...")
    pl.seed_everything(42)
    net = TemporalFusionTransformer.from_dataset(
        training,
        learning_rate=1e-3,             # Tuned learning rate
        hidden_size=64,                 # Strengthened hidden representation
        attention_head_size=4,          # Expanded attention mapping 
        dropout=0.2,                    # Prevent Overfitting
        hidden_continuous_size=16,
        output_size=7,                  # Quantiles for confidence bands!
        loss=QuantileLoss(),
        log_interval=10,
        reduce_on_plateau_patience=4,
    )

    # 4. Spin up the Lightning Trainer
    early_stop_callback = EarlyStopping(
        monitor="val_loss", 
        min_delta=1e-4, 
        patience=10, 
        verbose=True, 
        mode="min"
    )
    lr_logger = LearningRateMonitor()  # Logs the learning rate
    
    checkpoint_callback = ModelCheckpoint(
        monitor="val_loss",
        mode="min",
        save_top_k=1,
        filename="best-advanced-tft-{epoch:02d}-{val_loss:.2f}"
    )
    
    trainer = pl.Trainer(
        max_epochs=50,
        accelerator="auto", # auto detects MPS/GPU/CPU
        devices=1,
        gradient_clip_val=0.1, # Gradient clipping
        callbacks=[early_stop_callback, checkpoint_callback, lr_logger],
    )

    # 5. Execute Training!
    print(f"Executing High-Performance Multi-variate Training...")
    trainer.fit(
        net,
        train_dataloaders=train_dataloader,
        val_dataloaders=val_dataloader,
    )
    print("Optimization Complete! Best multi-variate model saved.")
    
    # --- INTERPRETABILITY TOOLS ---
    print("Generating Global Feature Importance and Attention...")
    best_model_path = trainer.checkpoint_callback.best_model_path
    if best_model_path:
        best_tft = TemporalFusionTransformer.load_from_checkpoint(best_model_path)
    else:
        best_tft = net

    # Run evaluation across sample
    raw_predictions = best_tft.predict(val_dataloader, mode="raw", return_x=True)
    
    # 1. Feature Importance plotting
    interpretation = best_tft.interpret_output(raw_predictions.output, reduction="sum")
    figs = best_tft.plot_interpretation(interpretation)
    for key, fig in figs.items():
        fig.savefig(f"tft_feature_importance_{key}.png")
    
    print("✅ Feature importance successfully mapped out to PNGs. Model represents State-of-the-Art config setup.")

if __name__ == "__main__":
    train_multivariate()