agrisense-backend / Modules /Price_TFT /scripts /tft_multivariate_training.py
Karthik Reddy
System Core: Production Engine Refactoring and Backend Hub
9f482ab
Raw
History Blame Contribute Delete
7.3 kB
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"])
# 🚨 CRITICAL FIX: Handle missing Weather/Interaction Data
# 30% of weather data was missing due to district mismatches.
# NNs cannot process NaNs. We forward-fill the gaps, then fill remaining with 0
# (For features like temp_anomaly and momentum, 0 mathematically means 'neutral/average')
numeric_cols = df.select_dtypes(include=["number"]).columns
df[numeric_cols] = df.groupby(["Mandi", "Commodity"])[numeric_cols].ffill().fillna(0)
# 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)
# Holding out the last 30 days of the 10-month timeline for robust validation
training_cutoff = df["time_idx"].max() - 30
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 = 512 # Massively bumped for GPU parallelization
train_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=8)
val_dataloader = validation.to_dataloader(train=False, batch_size=batch_size * 2, num_workers=8)
# 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=128, # UPGRADED: 128 to handle 56 features
attention_head_size=4, # Expanded attention mapping
dropout=0.2, # Prevent Overfitting
hidden_continuous_size=32, # UPGRADED: 32 to prevent information bottleneck
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="gpu", # Specifically request GPU execution
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()