From 8ce63d412990e241b6c5a990b83642da455deb5c Mon Sep 17 00:00:00 2001 From: Mahimn Patel Date: Thu, 24 Sep 2026 20:19:51 -0400 Subject: [PATCH] feat: support multi-batch TFT explanations --- CHANGELOG.md | 1 + darts/explainability/tft_explainer.py | 189 ++++++++---------- .../explainability/test_tft_explainer.py | 123 ++++++++++-- 3 files changed, 186 insertions(+), 127 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 969f25ea24..faa364d5df 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -21,6 +21,7 @@ but cannot always guarantee backwards compatibility. Changes that may **break co - `FittableAnomalyScorer.fit_from_prediction()` now returns the fitted scorer object similar to `fit()`. [#3202](https://github.com/unit8co/darts/pull/3202) by [Venish Paneliya](https://github.com/VenishPaneliya). - Calling `ForecastingModel.historical_forecasts()` with a `start` value that is later than what is forecastable given the supplied covariates now raises an informative exception. [#3207](https://github.com/unit8co/darts/pull/3207) by [Dennis Bader](https://github.com/dennisbader). - Calling `ForecastingModel.gridsearch()` with a sequence of `TimeSeries` now raises an informative exception. [#3191](https://github.com/unit8co/darts/pull/3191) by [Geovanny Basantes](https://github.com/COMPUMAX-EC). +- `TFTExplainer` can now explain more series than the model's batch size. [#3182](https://github.com/unit8co/darts/pull/3182) by [Mahimn](https://github.com/mahimn01). **Fixed** diff --git a/darts/explainability/tft_explainer.py b/darts/explainability/tft_explainer.py index 8e9931a9ab..e91a1f7180 100644 --- a/darts/explainability/tft_explainer.py +++ b/darts/explainability/tft_explainer.py @@ -29,6 +29,8 @@ import matplotlib.pyplot as plt import numpy as np import pandas as pd +import pytorch_lightning as pl +import torch from matplotlib.figure import Figure from torch import Tensor @@ -38,12 +40,31 @@ from darts.logging import get_logger, raise_log from darts.models import TFTModel from darts.typing import TimeSeriesLike -from darts.utils.ts_utils import SeriesType, get_series_seq_type from darts.utils.utils import generate_index logger = get_logger(__name__) +class _TFTWeightsCollector(pl.Callback): + """Collects the attention and variable selection weights of every prediction batch.""" + + def __init__(self): + self.weights = [] + + def on_predict_batch_end( + self, trainer, pl_module, outputs, batch, batch_idx, dataloader_idx=0 + ): + self.weights.append([ + None if weight is None else weight.detach().cpu() + for weight in ( + pl_module._attn_out_weights, + pl_module._encoder_sparse_weights, + pl_module._decoder_sparse_weights, + pl_module._static_covariate_var, + ) + ]) + + class TFTExplainer(_ForecastingModelExplainer): model: TFTModel @@ -136,6 +157,9 @@ def explain( - encoder variable importances per timestep of the input chunk and decoder variable importances per timestep of the output chunk. + .. note:: + Multi-device and multi-node prediction is not supported. + Parameters ---------- foreground_series @@ -194,41 +218,71 @@ def explain( foreground_past_covariates, foreground_future_covariates, ) - if ( - get_series_seq_type(foreground_series) is SeriesType.SEQ - and len(foreground_series) > self.model.batch_size - ): + horizons, _ = self._process_horizons_and_targets(None, None) + + # collect the weights of every prediction batch, so that more series than the + # model's batch size can be explained + collector = _TFTWeightsCollector() + trainer_params = dict(self.model.trainer_params) + trainer_params["callbacks"] = [ + collector, + *(trainer_params.get("callbacks") or []), + ] + trainer = self.model._init_trainer( + trainer_params=trainer_params, max_epochs=self.model.n_epochs + ) + if trainer.strategy.launcher is not None or trainer.world_size > 1: raise_log( ValueError( - f"The number of back- or foreground series to explain ({len(foreground_series)}) " - f"must be smaller than or equal to the model's batch size ({self.model.batch_size})." + "`TFTExplainer` does not support multi-device or multi-node prediction." ), ) - - horizons, _ = self._process_horizons_and_targets(None, None) - preds = self.model.predict( - n=self.n, - series=foreground_series, - past_covariates=foreground_past_covariates, - future_covariates=foreground_future_covariates, + try: + preds = self.model.predict( + n=self.n, + series=foreground_series, + past_covariates=foreground_past_covariates, + future_covariates=foreground_future_covariates, + trainer=trainer, + ) + finally: + # the model keeps a reference to the trainer, so the collected weights + # should not stay attached to it + trainer.callbacks.remove(collector) + ( + attention_weights, + encoder_weights, + decoder_weights, + static_covariate_weights, + ) = ( + None if weights[0] is None else torch.cat(weights) + for weights in zip(*collector.weights) ) - # get the weights and the attention head from the trained model for the prediction + # aggregate over attention heads - attention_heads = ( - self.model.model._attn_out_weights.detach().cpu().numpy().sum(axis=-2) - ) + attention_heads = attention_weights.numpy().sum(axis=-2) # get the variable importances (pd.DataFrame with rows corresponding to the number of input series) - encoder_importance = self._encoder_importance - decoder_importance = self._decoder_importance - static_covariates_importance = self._static_covariates_importance + encoder_importance = self._get_importance( + weight=encoder_weights, names=self.model.model.encoder_variables + ) + decoder_importance = self._get_importance( + weight=decoder_weights, names=self.model.model.decoder_variables + ) + static_covariates_importance = self._get_importance( + weight=static_covariate_weights, names=self.model.model.static_variables + ) # get the encoder/decoder importances over time; # static covariates have no time dimension, so there is no "over time" variant for them encoder_importance_over_time, encoder_var_names = ( - self._encoder_importance_over_time + self._get_importance_over_time( + weight=encoder_weights, names=self.model.model.encoder_variables + ) ) decoder_importance_over_time, decoder_var_names = ( - self._decoder_importance_over_time + self._get_importance_over_time( + weight=decoder_weights, names=self.model.model.decoder_variables + ) ) horizon_idx = [h - 1 for h in horizons] @@ -468,68 +522,9 @@ def plot_attention( return plotted_figures[0] return plotted_figures - @property - def _encoder_importance(self) -> pd.DataFrame: - """Returns the encoder variable importance of the TFT model. - - The encoder_weights are calculated for the past inputs of the model. - The encoder_importance contains the weights of the encoder variable selection network. - The encoder variable selection network is used to select the most important static and time dependent - covariates. It provides insights which variable are most significant for the prediction problem. - See section 4.2 of the paper for more details. - - Returns - ------- - pd.DataFrame - The encoder variable importance. - """ - return self._get_importance( - weight=self.model.model._encoder_sparse_weights, - names=self.model.model.encoder_variables, - ) - - @property - def _decoder_importance(self) -> pd.DataFrame: - """Returns the decoder variable importance of the TFT model. - - The decoder_weights are calculated for the known future inputs of the model. - The decoder_importance contains the weights of the decoder variable selection network. - The decoder variable selection network is used to select the most important static and time dependent - covariates. It provides insights which variable are most significant for the prediction problem. - See section 4.2 of the paper for more details. - - Returns - ------- - pd.DataFrame - The importance of the decoder variables. - """ - return self._get_importance( - weight=self.model.model._decoder_sparse_weights, - names=self.model.model.decoder_variables, - ) - - @property - def _static_covariates_importance(self) -> pd.DataFrame: - """Returns the static covariates importance of the TFT model. - - The static covariate importances are calculated for the static inputs of the model (numeric and / or - categorical). The static variable selection network is used to select the most important static covariates. - It provides insights which variable are most significant for the prediction problem. - See section 4.2, and 4.3 of the paper for more details. - - Returns - ------- - pd.DataFrame - The static covariates importance. - """ - return self._get_importance( - weight=self.model.model._static_covariate_var, - names=self.model.model.static_variables, - ) - def _get_importance( self, - weight: Tensor, + weight: Tensor | None, names: list[str], n_decimals=3, ) -> pd.DataFrame: @@ -571,34 +566,6 @@ def _get_importance( # return the importance sorted descending return importance.transpose().sort_values(0, ascending=True).transpose() - @property - def _encoder_importance_over_time(self) -> tuple[np.ndarray, list[str]]: - """Returns the encoder variable importance over time of the TFT model. - - Returns - ------- - tuple[np.ndarray, list[str]] - The importance over time of the encoder variables as well as the variable names. - """ - return self._get_importance_over_time( - weight=self.model.model._encoder_sparse_weights, - names=self.model.model.encoder_variables, - ) - - @property - def _decoder_importance_over_time(self) -> tuple[np.ndarray, list[str]]: - """Returns the decoder variable importance over time of the TFT model. - - Returns - ------- - tuple[np.ndarray, list[str]] - The importance over time of the decoder variables as well as the variable names. - """ - return self._get_importance_over_time( - weight=self.model.model._decoder_sparse_weights, - names=self.model.model.decoder_variables, - ) - def _get_importance_over_time( self, weight: Tensor, diff --git a/darts/tests/explainability/test_tft_explainer.py b/darts/tests/explainability/test_tft_explainer.py index 6d1bd084cd..cc7af086e5 100644 --- a/darts/tests/explainability/test_tft_explainer.py +++ b/darts/tests/explainability/test_tft_explainer.py @@ -13,10 +13,22 @@ f"Torch not available. {__name__} tests will be skipped.", allow_module_level=True, ) +import pytorch_lightning as pl + from darts.explainability import TFTExplainabilityResult, TFTExplainer from darts.models import TFTModel +class _PredictionBatchSizes(pl.Callback): + def __init__(self): + self.batch_sizes = [] + + def on_predict_batch_end( + self, trainer, pl_module, outputs, batch, batch_idx, dataloader_idx=0 + ): + self.batch_sizes.append(batch.past_target.shape[0]) + + def helper_create_test_cases(series_options: list): covariates_options = [ {}, @@ -321,21 +333,71 @@ def test_explainer_multiple_multivariate_series(self, test_case): for ts in enc_imp_ot + dec_imp_ot: np.testing.assert_allclose(ts.values().sum(axis=1), 100.0, atol=0.1) - # cannot explain more series than the batch size - with pytest.raises(ValueError) as exc: - explainer.explain( - foreground_series=[series[0]] * (model.batch_size + 1), - foreground_past_covariates=[pc[0]] * (model.batch_size + 1) - if use_pc - else None, - foreground_future_covariates=[fc[0]] * (model.batch_size + 1) - if use_fc - else None, - ) - assert str(exc.value) == ( - "The number of back- or foreground series to explain (33) must be smaller than " - "or equal to the model's batch size (32)." + @pytest.mark.parametrize("with_static_covariates", [True, False]) + @pytest.mark.parametrize("n_series,batch_sizes", [(2, [2]), (5, [2, 2, 1])]) + def test_explainer_more_series_than_batch_size( + self, n_series, batch_sizes, with_static_covariates + ): + """Test that series spread over several prediction batches are explained the same as + when they are explained one at a time.""" + series, past_covariates, future_covariates = self.helper_get_distinct_input( + n_series, with_static_covariates + ) + callback = _PredictionBatchSizes() + model = self.helper_create_model(batch_size=2, callbacks=[callback]) + model.fit( + series, past_covariates=past_covariates, future_covariates=future_covariates + ) + explainer = TFTExplainer( + model, + background_series=series, + background_past_covariates=past_covariates, + background_future_covariates=future_covariates, ) + result = explainer.explain() + assert callback.batch_sizes == batch_sizes + + for idx in range(n_series): + expected = explainer.explain( + foreground_series=series[idx], + foreground_past_covariates=past_covariates[idx], + foreground_future_covariates=future_covariates[idx], + ) + for getter in [ + "get_attention", + "get_encoder_importance", + "get_decoder_importance", + "get_static_covariates_importance", + "get_encoder_importance_over_time", + "get_decoder_importance_over_time", + ]: + actual_value = getattr(result, getter)()[idx] + expected_value = getattr(expected, getter)() + if isinstance(expected_value, TimeSeries): + assert actual_value.time_index.equals(expected_value.time_index) + np.testing.assert_allclose( + actual_value.values(), + expected_value.values(), + rtol=1e-6, + atol=1e-6, + ) + else: + pd.testing.assert_frame_equal( + actual_value.reset_index(drop=True), + expected_value.reset_index(drop=True), + check_like=True, + rtol=1e-6, + atol=1e-6, + ) + + @pytest.mark.parametrize("devices", [1, 2]) + def test_explainer_distributed_prediction(self, devices): + """Test that multi-device and multi-node prediction is rejected.""" + model = self.helper_create_model() + model.fit(self.series_mv1, past_covariates=self.pc, future_covariates=self.fc) + model.trainer_params.update(devices=devices, strategy="ddp_spawn") + with pytest.raises(ValueError, match="multi-device or multi-node"): + TFTExplainer(model).explain() @pytest.mark.parametrize("n_series", [1, 2]) def test_variable_selection_explanation(self, n_series, mpl_safe_plotting): @@ -530,20 +592,49 @@ def _check_plot(n_figs_expected, n_axes_expected, **kwargs): _check_plot(n_series, 2, plot_type="invalid", show_index_as="time") def helper_create_model( - self, use_encoders=True, add_relative_idx=True, full_attention=False + self, + use_encoders=True, + add_relative_idx=True, + full_attention=False, + batch_size=32, + callbacks=None, ): add_encoders = ( {"cyclic": {"past": ["month"], "future": ["month"]}} if use_encoders else None ) + model_kwargs = dict(tfm_kwargs) + if callbacks is not None: + model_kwargs["pl_trainer_kwargs"] = { + **tfm_kwargs["pl_trainer_kwargs"], + "callbacks": callbacks, + } return TFTModel( input_chunk_length=5, output_chunk_length=2, n_epochs=1, + batch_size=batch_size, add_encoders=add_encoders, add_relative_index=add_relative_idx, full_attention=full_attention, random_state=42, - **tfm_kwargs, + **model_kwargs, ) + + def helper_get_distinct_input(self, n_series, with_static_covariates): + series, past_covariates, future_covariates = [], [], [] + for idx in range(n_series): + static_covariates = ( + pd.Series([idx % 2, idx + 0.25], index=["cat", "num"]) + if with_static_covariates + else None + ) + series.append( + (self.series_mv1 * (idx + 1) + idx) + .shift(idx * 12) + .with_static_covariates(static_covariates) + ) + past_covariates.append((self.pc * (idx + 1)).shift(idx * 12)) + future_covariates.append((self.fc * (idx + 2) + idx).shift(idx * 12)) + return series, past_covariates, future_covariates