SHAP Explainer for SKLearn and Torch Models#
A SHAP explainer for Darts’ SKLearnModel and TorchForecastingModel
instances.
For detailed examples and tutorials, see:
ShapExplainer computes SHAP values, which measure each input feature’s contribution to a prediction
relative to a baseline (average prediction).
Depending on the model and training data, features can include:
lags of the target series (input chunk for torch models)
lags of past covariates (input chunk for torch models)
lags of future covariates (input and output chunk for torch models)
static covariates (global or component-specific)
Note
All input features except static covariates are named according to the convention
"{name}_{type_of_cov}_lag{idx}", where:
{name}is the component name from the original foreground series (target, past covariates, or future covariates).{type_of_cov}is the covariates type. It can take 3 different values:"target","pastcov","futcov".{idx}is the lag index, where0represents the position of the first predicted step.
Static covariates are named according to the convention: "{name}_statcov_target_{comp}", where:
{name}is the variable name of the static covariate.{comp}is the component name of the target series if static covariates are component-specific, or"global_components"if they are global.
Note
SHAP uses a feature-independence assumption. Indirect effects between features are not captured.
ShapExplainer provides the following methods for explaining multiple forecasts in batches:
explain()computes SHAP values per forecast horizon and target component.summary_plot()shows SHAP value distributions by feature.force_plot()shows additive SHAP contributions for one target component and horizon.
ShapExplainer also provides explain_single() for explaining
a single forecast (equivalent to calling model.predict(n=output_chunk_length)).
Note
All above methods can use optional foreground data to explain forecasts, with background data as reference. If foreground data is not provided, background data is used for both.
- class darts.explainability.shap_explainer.ShapExplainer(model, background_series=None, background_past_covariates=None, background_future_covariates=None, background_num_samples=None, shap_method=None, batch_size=None, test_stationarity=True, **kwargs)[source]#
Bases:
_ForecastingModelExplainerSHAP Explainer for SKLearn and Torch Models.
Definitions:
A background series is a
TimeSeriesused to train the SHAP explainer.A foreground series is a
TimeSeriesthat can be explained by a SHAP explainer after it has been fitted.
The number of explained horizons (t+1, t+2, …) cannot be greater than
output_chunk_lengthofmodel.- Parameters:
model (
SKLearnModel|TorchForecastingModel) – TheSKLearnModelorTorchForecastingModelto be explained. It must be fitted first.background_series (
Union[TimeSeries,Sequence[TimeSeries],None]) – One or several series to train theShapExplaineras reference for explanations. Consider using a reduced well-chosen background to reduce computation time. Optional ifmodelwas fit on a single target series. By default, it is theseriesused at fitting time. Mandatory ifmodelwas fit on multiple (list of) target series.background_past_covariates (
Union[TimeSeries,Sequence[TimeSeries],None]) – A past covariates series or list of series that the model needs once fitted.background_future_covariates (
Union[TimeSeries,Sequence[TimeSeries],None]) – A future covariates series or list of series that the model needs once fitted.background_num_samples (
int|None) – Optionally, whether to sample a subset of the original background. Randomly picks samples of the constructed training dataset. Generally used for faster computation, especially whenshap_methodis"kernel"or"permutation".shap_method (
str|None) – Optionally, the SHAP method to apply. By default, an attempt is made to select the most appropriate method based on a pre-defined set of known models internal mapping. Supported values forSKLearnModel:["tree", "kernel", "partition", "linear", "permutation", "additive"]. Supported valuesTorchForecastingModel:["kernel", "partition", "sampling", "permutation"].batch_size (
int|None) – Optionally, the batch size to use whenmodelis aTorchForecastingModel. Increasing the batch size can significantly reduce computation time.test_stationarity (
bool) – Whether to perform stationarity checks and raise a warning if not all background_series are stationary.**kwargs – Optionally, additional keyword arguments passed to
shap_method.
Examples
For
SKLearnModel:>>> from darts.datasets import AusBeerDataset >>> from darts.explainability import ShapExplainer >>> from darts.models import LinearRegressionModel >>> series = AusBeerDataset().load().astype("float32")[:-36] >>> model = LinearRegressionModel(lags=12, output_chunk_length=1).fit(series) >>> explainer = ShapExplainer(model) >>> result = explainer.explain() >>> explainer.summary_plot() >>> explainer.force_plot()
For
TorchForecastingModel:>>> from darts.datasets import AusBeerDataset >>> from darts.explainability import ShapExplainer >>> from darts.models import TiDEModel >>> series = AusBeerDataset().load().astype("float32")[:-36] >>> model = TiDEModel(input_chunk_length=12, output_chunk_length=1).fit(series) >>> explainer = ShapExplainer(model, batch_size=2048) >>> result = explainer.explain() >>> explainer.summary_plot() >>> explainer.force_plot()
Methods
explain([foreground_series, ...])Explains all possible foreground series forecasts (or background, if foreground is not provided) and returns a
ShapExplainabilityResultof SHAP values.explain_single([foreground_series, ...])Explains the last forecast of a foreground series (or background, if foreground is not provided) and returns a
ShapSingleExplainabilityResultof SHAP values.force_plot([foreground_series, ...])Display a SHAP "Force Plot" for one target and one horizon.
summary_plot([foreground_series, ...])Display a SHAP "Summary Plot" for each horizon and each component dimension of the target.
- explain(foreground_series=None, foreground_past_covariates=None, foreground_future_covariates=None, horizons=None, target_components=None, **kwargs)[source]#
Explains all possible foreground series forecasts (or background, if foreground is not provided) and returns a
ShapExplainabilityResultof SHAP values.The results can then be retrieved with method
get_explanation(), which returns a multivariateTimeSeriesinstance containing the SHAP values for the(horizon, target_component)forecasts at all timestamps forecastable in the foreground series.The components of the
TimeSeriescorrespond to the input features used by the model to produce the forecasts. See above for the naming convention.- Parameters:
foreground_series (
Union[TimeSeries,Sequence[TimeSeries],None]) – Optionally, one or a sequence of targetTimeSeriesto be explained. Can be multivariate. Default:None, which means that the background series will be used as foreground.foreground_past_covariates (
Union[TimeSeries,Sequence[TimeSeries],None]) – Optionally, one or a sequence of past covariatesTimeSeriesif required by the forecasting model.foreground_future_covariates (
Union[TimeSeries,Sequence[TimeSeries],None]) – Optionally, one or a sequence of future covariatesTimeSeriesif required by the forecasting model.horizons (
int|Sequence[int] |None) – Optionally, an integer or sequence of integers representing the future time steps to be explained.1corresponds to the first timestamp being forecasted. All values must be no greater thanoutput_chunk_lengthof the explained forecasting model.target_components (
Sequence[str] |None) – Optionally, a string or sequence of strings with the target components to explain.**kwargs – Other keyword arguments to be passed to the SHAP explainer.
- Returns:
The forecast explanations of the specified horizons and target components.
- Return type:
Examples
Say we have a
SKLearnModelinstance with:1 target component named
"Y",1 future covariate named
"month",lags = 2, andlags_future_covariates = [-1, 0].
Let’s explain the background series that the model was trained on:
>>> from darts.datasets import AusBeerDataset >>> from darts.explainability import ShapExplainer >>> from darts.models import LinearRegressionModel >>> from darts.utils.timeseries_generation import datetime_attribute_timeseries as dta >>> >>> # load a target series and create future covariates holding the calendar month values >>> series = AusBeerDataset().load() >>> fc = dta(series, attribute="month", add_length=12) >>> >>> # create and fit a model >>> model = LinearRegressionModel(lags=2, lags_future_covariates=[-1, 0]) >>> model.fit(series, future_covariates=fc) >>> >>> # create an explainer; requires background series if the model was trained on multiple series >>> explainer = ShapExplainer(model) >>> # explain the background series (or foreground if passed to `explain()`) >>> result = explainer.explain() >>> >>> # get explanations for a specific horizon (and optional `component` for multivariate models) >>> # the feature SHAP values for all possible forecast start points >>> result.get_explanation(horizon=1) Y_target_lag-2 Y_target_lag-1 month_futcov_lag-1 month_futcov_lag0 1956-07-01 -56.332566 -106.927156 -24.253184 33.064478 1956-10-01 -88.545937 -99.569541 13.642416 84.727725 1957-01-01 -82.194005 -57.000488 51.538016 -70.262016 1957-04-01 -45.443539 -81.175506 -62.148784 -18.598769 1957-07-01 -66.314174 -99.043998 -24.253184 33.064478 ... ... ... ... ... 2007-10-01 -11.415330 -11.803715 13.642416 84.727725 2008-01-01 -6.424526 29.714250 51.538016 -70.262016 2008-04-01 29.418521 1.860425 -62.148784 -18.598769 2008-07-01 5.371920 -13.905891 -24.253184 33.064478 2008-10-01 -8.239364 -3.395013 13.642416 84.727725 shape: (210, 4, 1), freq: QS-OCT, size: 6.56 KB
The explanation has length 210, containing the feature SHAP values for all possible forecast start points over the background series.
Now, let’s get the feature values that were used as model input to forecast the series:
>>> result.get_feature_values(horizon=1) Y_target_lag-2 Y_target_lag-1 month_futcov_lag-1 month_futcov_lag0 1956-07-01 284.0 213.0 3.0 6.0 1956-10-01 213.0 227.0 6.0 9.0 1957-01-01 227.0 308.0 9.0 0.0 1957-04-01 308.0 262.0 0.0 3.0 1957-07-01 262.0 228.0 3.0 6.0 ... ... ... ... ... 2007-10-01 383.0 394.0 6.0 9.0 2008-01-01 394.0 473.0 9.0 0.0 2008-04-01 473.0 420.0 0.0 3.0 2008-07-01 420.0 390.0 3.0 6.0 2008-10-01 390.0 410.0 6.0 9.0 shape: (210, 4, 1), freq: QS-OCT, size: 6.56 KB
And also, we can get the raw shap.Explanation object for further processing:
>>> shap_object = result.get_shap_explanation_object(horizon=1)
- explain_single(foreground_series=None, foreground_past_covariates=None, foreground_future_covariates=None, target_components=None, **kwargs)[source]#
Explains the last forecast of a foreground series (or background, if foreground is not provided) and returns a
ShapSingleExplainabilityResultof SHAP values.The results can then be retrieved with method
get_explanation(), which returns a multivariateTimeSeriesinstance containing the SHAP values fortarget_componentstarting from the last forecastable timestamp.The components of the
TimeSeriescorrespond to the input features used by the model to produce the forecast. See above for the naming convention.Note
The forecast explained by this method is equivalent to the one obtained by calling
model.predict(n=output_chunk_length, series=series, ...)whereseriesis eitherforeground_seriesorbackground_seriesdepending on what was used when callingexplain_single().- Parameters:
foreground_series (TimeSeries | None) – Optionally, one or a sequence of target
TimeSeriesto be explained. Can be multivariate. Default:None, which means that the background series will be used as foreground.foreground_past_covariates (TimeSeries | None) – Optionally, one or a sequence of past covariates
TimeSeriesif required by the forecasting model.foreground_future_covariates (TimeSeries | None) – Optionally, one or a sequence of future covariates
TimeSeriesif required by the forecasting model.target_components (Sequence[str] | None) – Optionally, a string or sequence of strings with the target components to explain.
**kwargs – Other keyword arguments to be passed to the SHAP explainer.
- Returns:
The forecast explanations of the specified target components for the single forecasted timestamp.
- Return type:
Examples
Say we have a
SKLearnModelinstance with:1 target component named
"Y",1 future covariate named
"month",lags = 2, andlags_future_covariates = [-1, 0].
Let’s explain the background series that the model was trained on:
>>> from darts.datasets import AusBeerDataset >>> from darts.explainability import ShapExplainer >>> from darts.models import LinearRegressionModel >>> from darts.utils.timeseries_generation import datetime_attribute_timeseries as dta >>> >>> # load a target series and create future covariates holding the calendar month values >>> series = AusBeerDataset().load() >>> fc = dta(series, attribute="month", add_length=12) >>> >>> # create and fit a model >>> model = LinearRegressionModel(lags=2, lags_future_covariates=[-1, 0]) >>> model.fit(series, future_covariates=fc) >>> >>> # create an explainer; requires background series if the model was trained on multiple series >>> explainer = ShapExplainer(model) >>> # explain the background forecast (or foreground forecast if passed to `explain_single()`) >>> result = explainer.explain_single() >>> >>> # get explanations for that forecast (and optional component for multivariate models) >>> # the feature SHAP values for that forecast >>> result.get_explanation() Y_target_lag-2 Y_target_lag-1 month_futcov_lag-1 month_futcov_lag0 2008-10-01 -8.239364 -3.395013 13.642416 84.727725 shape: (1, 4, 1), freq: QS-OCT, size: 6.56 KB
The explanation has length 210, containing the feature SHAP values for all possible forecast start points over the background series.
Now, let’s get the feature values that were used as model input to forecast the series:
>>> result.get_feature_values() Y_target_lag-2 Y_target_lag-1 month_futcov_lag-1 month_futcov_lag0 2008-10-01 390.0 410.0 6.0 9.0 shape: (1, 4, 1), freq: QS-OCT, size: 6.56 KB
And also, we can get the raw shap.Explanation object for further processing:
>>> shap_object = result.get_shap_explanation_object()
- force_plot(foreground_series=None, foreground_past_covariates=None, foreground_future_covariates=None, horizon=1, target_component=None, plot_kwargs=None, **kwargs)[source]#
Display a SHAP “Force Plot” for one target and one horizon.
It shows SHAP values of all input features with an additive force layout for each forecastable timestamp in the foreground series. At each timestamp, SHAP values of all features and the base value would sum up to the model prediction.
Note
Once the plot is displayed, select “original sample ordering” to observe the forecasted timestamps chronologically.
- Parameters:
foreground_series (TimeSeries | None) – Optionally, the target series to explain. Can be multivariate. Default:
None, which means that the background series will be used as foreground.foreground_past_covariates (TimeSeries | None) – Optionally, a past covariate series if required by the forecasting model.
foreground_future_covariates (TimeSeries | None) – Optionally, a future covariate series if required by the forecasting model.
horizon (int | None) – Optionally, an integer for the point/step in the future to explain, starting from the first prediction step at 1. Must not be larger than
output_chunk_lengthof the model.target_component (str | None) – Optionally, the target component to plot. If the target series is multivariate, the target component must be specified.
plot_kwargs (dict[str, Any] | None) – Optionally, a dictionary of keyword arguments to be passed to
shap.force_plot().**kwargs – Other keyword arguments to be passed to the SHAP explainer.
- summary_plot(foreground_series=None, foreground_past_covariates=None, foreground_future_covariates=None, horizons=None, target_components=None, num_samples=None, plot_type='dot', plot_kwargs=None, **kwargs)[source]#
Display a SHAP “Summary Plot” for each horizon and each component dimension of the target.
On each summary plot, SHAP values of each input feature are plotted with dots (
plot_type="dot", each dot corresponds to a forecasted timestamp), a bar (plot_type="bar"), or a violin (plot_type="violin"). The input features are sorted by importance, defined as the mean absolute SHAP value.- Parameters:
foreground_series (
Union[TimeSeries,Sequence[TimeSeries],None]) – Optionally, one or a sequence of targetTimeSeriesto be explained. Can be multivariate. Default:None, which means that the background series will be used as foreground.foreground_past_covariates (
Union[TimeSeries,Sequence[TimeSeries],None]) – Optionally, one or a sequence of past covariatesTimeSeriesif required by the forecasting model.foreground_future_covariates (
Union[TimeSeries,Sequence[TimeSeries],None]) – Optionally, one or a sequence of future covariatesTimeSeriesif required by the forecasting model.horizons (
int|Sequence[int] |None) – Optionally, an integer or sequence of integers representing which points/steps in the future to explain, starting from the first prediction step at 1. Each horizon must be no greater thanoutput_chunk_lengthof the explained forecasting model. Default:None, which means that all horizons will be plotted.target_components (
str|Sequence[str] |None) – Optionally, a string or sequence of strings with the target components to explain. Default:None, which means that all target components will be plotted.num_samples (
int|None) – Optionally, an integer for sampling the foreground series for the sake of performance.plot_type (
str|None) – Optionally, specify which of the SHAP library plot type to use. Can be one of"dot","bar","violin".plot_kwargs (
dict[str,Any] |None) – Optionally, a dictionary of keyword arguments to be passed toshap.summary_plot().**kwargs – Other keyword arguments to be passed to the SHAP explainer.
- Returns:
A nested dictionary
{horizon : {component : shap.Explanation}}containing the raw Explanation objects for all the horizons and components.- Return type:
dict[int, dict[str, shap.Explanation]]