Global feature effect methods such as partial dependence (PD) and accumulated local effects (ALE) describe how a fitted model’s prediction changes, on average, when one feature is varied. Because they average over all observations, they can hide feature interactions: the effect of a feature may look flat globally while being strongly positive in one subgroup and strongly negative in another.
xplaineff implements the GADGET algorithm (Generalized Additive Decomposition of Global EffecTs; Herbinger et al., 2024). GADGET recursively partitions the feature space on so-called split features so that, within each region, the effects of the features of interest are as homogeneous as possible across observations. The result is a shallow tree whose nodes are regions with their own regional PD or ALE curves. Splits reveal which features interact, and the regional curves show how the effect changes between regions.
The package exposes three R6 classes:
| Class | Role |
|---|---|
GadgetTree |
The tree: $new(), $fit(),
$extract_split_info(), $plot_tree_structure(),
$plot() |
PdStrategy |
Heterogeneity and plots based on PD and individual conditional expectation (ICE) curves |
AleStrategy |
Heterogeneity and plots based on ALE |
A typical workflow has four steps:
PdStrategy or
AleStrategy) and create a GadgetTree with the
stopping rules you want.$fit() with the training data, the target column,
the model, and optionally the features of interest
(feature_set) and the candidate split features
(split_feature).$extract_split_info(),
$plot_tree_structure(), and $plot().We start with a small simulated data set in which the effect of
x2 is moderated by x3: the slope of
x2 is +3 when x3 > 0.3 and
-3 otherwise. A global PD curve for x2 would
average these two regimes and look almost flat.
set.seed(1)
n = 500
x1 = runif(n, -1, 1)
x2 = runif(n, -1, 1)
x3 = runif(n, -1, 1)
y = 0.2 * x1 + ifelse(x3 > 0.3, 3, -3) * x2 + rnorm(n, 0, 0.3)
syn_data = data.frame(x1, x2, x3, y)Any model with a predict() method can be explained. Here
we use a linear model with the correct interaction term so that the
example runs without extra packages.
We now grow a PD tree. feature_set names the feature of
interest whose effect we want to decompose, and
split_feature lists the features allowed to define regions.
The constructor arguments control when the tree stops growing:
n_split: maximum number of splits (tree depth).min_node_size: minimum number of observations per child
node.n_quantiles: number of quantile-based candidate split
points per numeric split feature.impr_par: minimum relative improvement in heterogeneity
required to accept a split.The PD-specific argument n_grid sets the number of grid
points at which the PD and ICE curves are evaluated.
syn_tree = GadgetTree$new(
strategy = PdStrategy$new(),
n_split = 2,
min_node_size = 50,
n_quantiles = 40
)
syn_tree$fit(
data = syn_data,
target_feature_name = "y",
model = syn_model,
feature_set = "x2",
split_feature = c("x1", "x3"),
n_grid = 20L
)extract_split_info() returns one row per node.
syn_split = syn_tree$extract_split_info()
print(
syn_split[, c("id", "depth", "n_obs", "node_type", "split_feature", "split_value", "int_imp", "is_final")],
row.names = FALSE, digits = 2
)
#> id depth n_obs node_type split_feature split_value int_imp is_final
#> 1 1 500 root x3 0.31 0.99 FALSE
#> 2 2 329 left <NA> NA NA TRUE
#> 3 2 171 right <NA> NA NA TRUEThe root split recovers the moderator x3 close to the
true threshold of 0.3. int_imp is the
improvement in heterogeneity achieved by the split, normalized by the
root heterogeneity; a value near one means the split removes almost all
of the heterogeneity of the x2 effect. Both child nodes are
leaves (is_final = TRUE): within each x3
regime the effect of x2 is homogeneous, so the tree stops
before reaching n_split.
plot() draws the regional PD and ICE curves for every
node. With show_plot = FALSE it only returns the plots as a
nested list named by depth and node id, which is convenient for choosing
what to display.
syn_plots = syn_tree$plot(
data = syn_data,
target_feature_name = "y",
show_plot = FALSE
)
syn_plots$Depth_1$Node_1At the root the ICE curves fan out in two directions and the averaged PD curve is nearly flat. After the split, each region shows one clean linear effect with the expected sign.
The remaining examples use the Bikeshare data from
ISLR2 and a random forest trained with
mlr3. They are evaluated only if these packages are
installed.
library(mlr3)
#> Warning: Paket 'mlr3' wurde unter R Version 4.5.2 erstellt
library(mlr3learners)
#> Warning: Paket 'mlr3learners' wurde unter R Version 4.5.2 erstellt
data("Bikeshare", package = "ISLR2")
set.seed(123)
bike = Bikeshare[sample(seq_len(nrow(Bikeshare)), 1000), ]
factor_features = c("season", "mnth", "weekday", "workingday", "holiday", "weathersit")
bike[factor_features] = lapply(bike[factor_features], as.factor)
bike_data = bike[, c("hr", "temp", "workingday", "season", "mnth", "day", "holiday", "weekday",
"weathersit", "atemp", "hum", "windspeed", "bikers")]
names(bike_data)[names(bike_data) == "bikers"] = "target"
task = TaskRegr$new(id = "bike", backend = bike_data, target = "target")
learner = lrn("regr.ranger")
learner$train(task)The target is the hourly number of bikers. We decompose the effects
of hr, temp, workingday, and
season, and allow temp,
workingday, and season to define regions.
Categorical split features must be factors; the package converts
character columns automatically but treats numeric 0/1 indicators as
numeric unless they are encoded as factors beforehand.
effect_features = c("hr", "temp", "workingday", "season")
split_features = c("temp", "workingday", "season")
bike_pd = GadgetTree$new(
strategy = PdStrategy$new(),
n_split = 2,
min_node_size = 50
)
bike_pd$fit(
data = bike_data,
target_feature_name = "target",
model = learner,
feature_set = effect_features,
split_feature = split_features,
n_grid = 20L
)plot_tree_structure() draws the tree with the split
conditions and node sizes.
bike_pd_split = bike_pd$extract_split_info()
print(
bike_pd_split[, c("id", "depth", "n_obs", "node_type", "split_feature", "split_value", "int_imp", "is_final")],
row.names = FALSE, digits = 2
)
#> id depth n_obs node_type split_feature split_value int_imp is_final
#> 1 1 1000 root temp 0.49 0.44 FALSE
#> 2 2 489 left season 1 0.10 FALSE
#> 3 2 511 right workingday 0 0.08 FALSE
#> 4 3 224 left <NA> <NA> NA TRUE
#> 5 3 265 right <NA> <NA> NA TRUE
#> 6 3 152 left <NA> <NA> NA TRUE
#> 7 3 359 right <NA> <NA> NA TRUEThe root split separates cooler and warmer hours on
temp. Each branch is then refined by a categorical feature.
For a categorical split, split_value is the level that
defines the partition. With PdStrategy (one-vs-rest splits)
it is the single level sent to the left child; with
AleStrategy (ordered-prefix splits) it is the last level of
the ordered prefix sent to the left child. The tree plot shows the
complete level sets of both children.
Regional curves for a single feature are obtained with the
features argument. We look at the hourly pattern at the
root and in its two children.
bike_pd_plots = bike_pd$plot(
data = bike_data,
target_feature_name = "target",
features = "hr",
show_plot = FALSE
)
bike_pd_plots$Depth_1$Node_1At the root the ICE curves for hr spread widely around
the PD curve. Within each temperature region the curves are more
concentrated, and the two regions differ in the strength of the evening
peak.
AleStrategy uses the same interface. ALE is computed
internally from the model; the ALE-specific argument
n_intervals sets the number of intervals used for
accumulation. Here we also set impr_par so that only splits
with a relative improvement of at least one percent are accepted.
bike_ale = GadgetTree$new(
strategy = AleStrategy$new(),
n_split = 2,
impr_par = 0.01,
min_node_size = 50
)
bike_ale$fit(
data = bike_data,
target_feature_name = "target",
model = learner,
feature_set = effect_features,
split_feature = split_features,
n_intervals = 10
)
bike_ale_split = bike_ale$extract_split_info()
print(
bike_ale_split[, c("id", "depth", "n_obs", "node_type", "split_feature", "split_value", "int_imp", "is_final")],
row.names = FALSE, digits = 2
)
#> id depth n_obs node_type split_feature split_value int_imp is_final
#> 1 1 1000 root workingday 0 0.53 FALSE
#> 2 2 316 left season 1 0.02 FALSE
#> 3 2 684 right temp 0.47 0.16 FALSE
#> 4 3 72 left <NA> <NA> NA TRUE
#> 5 3 244 right <NA> <NA> NA TRUE
#> 6 3 316 left <NA> <NA> NA TRUE
#> 7 3 368 right <NA> <NA> NA TRUEThe ALE tree may choose a different root split than the PD tree, because PD and ALE measure heterogeneity on different scales: PD compares whole ICE curves, whereas ALE compares local prediction differences within intervals. Both views are valid; comparing them is a useful robustness check.
For ALE the int_imp_<feature> columns of the split
table additionally report, for each feature of interest, the reduction
of that feature’s own heterogeneity achieved by the split, relative to
its heterogeneity at the root.
print(
bike_ale_split[!bike_ale_split$is_final,
c("id", "split_feature", setdiff(grep("^int_imp_", names(bike_ale_split), value = TRUE), "int_imp_parent"))],
row.names = FALSE, digits = 2
)
#> id split_feature int_imp_hr int_imp_temp int_imp_workingday int_imp_season
#> 1 workingday 0.21 0.09 1 0.02
#> 2 season 0.03 0.01 0 0.09
#> 3 temp 0.38 0.00 0 0.03Regional ALE curves are drawn in the same way.
mean_center = TRUE (the default) centers each curve, which
makes shapes easier to compare across regions.
bike_ale_plots = bike_ale$plot(
data = bike_data,
target_feature_name = "target",
features = "hr",
show_plot = FALSE
)
bike_ale_plots$Depth_1$Node_1plot() accepts a few arguments that are shared by both
strategies:
depth and node_id restrict the output to
particular tree levels or nodes; node ids are the ones shown by
plot_tree_structure() and in the id column of
the split table.features selects the features of interest to draw.show_plot = FALSE suppresses printing and only returns
the list of ggplot objects.show_point overlays the observed data points, and
mean_center toggles centering of the curves.one_plot = bike_pd$plot(
data = bike_data,
target_feature_name = "target",
features = "temp",
depth = 2,
node_id = 3,
show_point = TRUE,
mean_center = FALSE,
show_plot = FALSE
)
one_plot$Depth_2$Node_3extract_split_info(include_timing = TRUE) adds the time
spent searching each split, which is helpful when tuning
n_quantiles or n_grid on larger data.
For PdStrategy you can pass ICE curves computed with
iml instead of a model. This is useful if you already
have a FeatureEffects object or want full control over the
prediction function.
library(iml)
predictor = Predictor$new(
model = learner,
data = bike_data[, setdiff(names(bike_data), "target")],
y = bike_data$target
)
effects = FeatureEffects$new(predictor, features = effect_features, method = "ice", grid.size = 20)
bike_pd_iml = GadgetTree$new(strategy = PdStrategy$new(), n_split = 2, min_node_size = 50)
bike_pd_iml$fit(
data = bike_data,
target_feature_name = "target",
effect = effects,
split_feature = split_features
)
bike_pd_iml$plot(effect = effects, data = bike_data, target_feature_name = "target", features = "hr")AleStrategy always computes ALE internally, because the
split search needs observation-level local effects that interval-level
ALE summaries from other packages do not provide.
Both strategies can split on categorical features. How the candidate
splits are formed is set in the strategy constructor through
categorical_split:
PdStrategy$new(categorical_split = "one_vs_rest")
(default) tries one level against all others.AleStrategy$new(categorical_split = "ordered_prefix")
(default) orders the levels and tries all prefix partitions of that
order.categorical_split = "exhaustive" searches all level
subsets for either strategy; max_exhaustive_levels caps the
number of levels for which this is attempted.For AleStrategy, the level order used by
"ordered_prefix" is chosen through the
order_method argument of $fit():
"raw" (default) keeps the factor level order."mds" and "pca" derive a data-driven order
from the other features, which helps when the levels have no natural
order."random" uses a random order and is mainly useful as a
baseline.Herbinger, J., Wright, M. N., Nagler, T., Bischl, B., and Casalicchio, G. (2024). Decomposing global feature effects based on feature interactions. Journal of Machine Learning Research, 25(381), 1–65. https://jmlr.org/papers/v25/23-0699.html