Spaces:
Running
Running
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()
|