Base class for SAGE (Shapley Additive Global Importance) feature importance based on Shapley values with marginalization. This is an abstract class - use MarginalSAGE or ConditionalSAGE.
Details
SAGE uses Shapley values to fairly distribute the total prediction performance among all features. Unlike perturbation-based methods, SAGE marginalizes features by integrating over their distribution. This is approximated by averaging predictions over a reference dataset.
SAGE values are reductions in the measure's score relative to the empty coalition,
score(empty) - score(S), so that positive values mean the feature improves performance.
For measures that are maximized (measure$minimize = FALSE, e.g. classif.acc) the scores are
negated internally, so the sign convention is the same for all measures.
Standard errors: The standard errors reported in $convergence_history and $convergence() are
Monte Carlo standard errors of the Shapley estimator: how much the estimates would still move if
more permutations or coalitions were sampled, for the fixed trained model, the fixed test set, and (for
MarginalSAGE) the fixed reference subsample.
They are convergence diagnostics in the sense of Covert & Lee (2021, Section 4.3), not inference about
feature importance: they say nothing about variability across train/test splits or refits (see the
resampling-based ci_methods of $importance() for that), and a feature's SAGE value being several
SEs from zero only means the computation has converged, not that the feature matters.
For the permutation estimator the SE is the sample standard error of the per-permutation marginal
contributions.
For the kernel estimator it follows from the multivariate central limit theorem for the estimated
regression targets, Cov(phi) = C Cov(b) C^T / n (their Eqs. 10-13), so it covers the coalition draws
and the test observations they are paired with (both are sampled with replacement from the fixed test
set).
The exact estimator has no coalition-sampling error and reports no SE.
Estimators: estimator = "permutation" (the default) is the permutation-sampling estimator of
Covert et al. (2020), budgeted by n_permutations.
estimator = "kernel" is the regression-based estimator of Covert & Lee (2021), budgeted by
n_coalitions: Shapley values are the solution of a weighted least squares problem, approximated from
sampled coalitions (see below).
estimator = "exact" enumerates all 2^n_features coalitions and computes the Shapley values in closed
form, so it has no coalition-sampling error and serves as a ground-truth reference for the sampling
estimators on small feature sets (capped by max_features).
"Exact" refers to coalition sampling only: the marginalization error controlled by n_samples remains,
and for ConditionalSAGE the value function itself is a Monte Carlo estimate of the sampler.
Kernel estimator: This implements the unbiased KernelSHAP estimator of Covert & Lee (2021, Eq. 9)
for the stochastic cooperative game of SAGE, mirroring the KernelEstimator of the reference Python
sage package.
Coalitions are drawn from the Shapley kernel (size k with probability proportional to 1 / (k (p - k)),
uniform within size) together with their complement (paired sampling, their Section 4.2), and each
draw is paired with a single test observation drawn with replacement, whose observation-wise loss is
the value-function sample.
Consequently the kernel estimator requires a measure with an observation-wise loss
("obs_loss" in measure$properties, e.g. regr.mse or classif.logloss, but not classif.auc), and
each coalition evaluation costs n_samples model rows rather than n_test * n_samples as for the other
estimators, which evaluate every coalition on the whole test set.
For such measures the kernel estimator targets the same SAGE values as the permutation and exact
estimators (the Shapley value is linear in the value function, and the test-set loss is the mean of the
observation-wise losses), so estimator = "exact" remains the reference; $budget$n_rows compares
their costs in model rows.
The kernel estimator is currently available for MarginalSAGE only.
Convergence and budget: With early_stopping = TRUE, sampling stops once the largest SE,
relative to the spread of the SAGE values (max(se) / (max(phi) - min(phi))), falls below
se_threshold.
This is the criterion of the reference Python sage package.
This applies to both sampling estimators; the exact estimator has no criterion.
The budget argument (n_permutations or n_coalitions) then acts as an upper bound rather than a planned cost:
exhausting it without meeting the criterion returns the values with a warning.
$budget reports what was actually spent and whether the criterion was met, and
$plot_convergence() shows the trajectory that led there.
Under resampling, only the first iteration runs the criterion and the remaining iterations
reuse its budget, which keeps them comparable and avoids re-deriving the standard errors in
every iteration.
References
Covert I, Lundberg S, Lee S (2020). “Understanding Global Feature Contributions With Additive Importance Measures.” In Advances in Neural Information Processing Systems, volume 33, 17212–17223. https://proceedings.neurips.cc/paper/2020/hash/c7bf0b7c1a86d5eb3be2c722cf2cf746-Abstract.html.
Covert I, Lee S (2021). “Improving KernelSHAP: Practical Shapley Value Estimation Using Linear Regression.” In Proceedings of the 24th International Conference on Artificial Intelligence and Statistics, volume 130, 3457–3465. https://proceedings.mlr.press/v130/covert21a.html.
Super class
FeatureImportanceMethod -> SAGE
Public fields
convergence_history(
data.table) History of SAGE values during computation. Columnsbudget(sampling effort in the estimator's own units),n_evals(the corresponding number of evaluated coalitions), andn_rows(model rows predicted) index the checkpoints; see$budget.converged(
logical(1)) Whether the convergence criterion was met (early_stopping = TRUE).NAfor the exact estimator, which enumerates all coalitions and has no criterion to meet.
Active bindings
budget(
data.table) Read-only one-row summary of the sampling effort: theestimator, itsunitof budget, therequestedupper bound, the amountused(below the request only with early stopping), the resulting number of coalition evaluationsn_evals(one empty-coalition baseline plusn_featuresper permutation; two anchors plus two per coalition draw for the kernel estimator;2^n_featuresfor the exact estimator), the number of model rows predictedn_rows, and whether the computationconverged.n_evalscounts coalition evaluations, which differ in cost between estimators (the kernel estimator evaluates a coalition on one test observation, the others on the whole test set), son_rowsis the unit in which estimators are comparable.used,n_evals, andn_rowsareNAbefore$compute();convergedisNAfor the exact estimator, which has no criterion to meet. With multiple resampling iterations it describes the first iteration, whose budget the remaining ones reuse (seeearly_stopping).n_permutations_usedDefunct. Use
$budgetinstead, which reports the effort spent alongside its unit and the implied number of coalition evaluations.n_permutations(
integer(1)) Deprecated. The permutation budget lives in the param_set; use$param_set$values$n_permutationsinstead. This alias is kept for backward compatibility with the field of the same name in earlier releases and warns on every access.
Methods
SAGE$new()
Creates a new instance of the SAGE class.
Usage
SAGE$new(
task,
learner,
measure = NULL,
resampling = NULL,
features = NULL,
estimator = c("permutation", "kernel", "exact"),
n_permutations = NULL,
n_coalitions = NULL,
max_features = 12L,
batch_size = 5000L,
n_samples = 100L,
early_stopping = FALSE,
se_threshold = 0.025,
min_permutations = 10L,
check_interval = 1L
)Arguments
task, learner, measure, resampling, featuresPassed to FeatureImportanceMethod.
estimator(
character(1):"permutation") Shapley-value estimator."permutation"is the permutation-sampling estimator of Covert et al. (2020), budgeted byn_permutations;"kernel"is the regression-based estimator of Covert & Lee (2021), budgeted byn_coalitions;"exact"enumerates all2^n_featurescoalitions (capped bymax_features) and takes no budget. All approximate the same SAGE values; setting the budget argument of a different estimator is an error. Their costs are comparable through$budget$n_rows, the number of model rows predicted; see Details for why the kernel estimator's coalition evaluations are cheaper than the others'.$compute()points out in a message (silenced byxplain_opt(verbose = FALSE)) when the sampling budget meets or exceeds the exact estimator's cost, since enumeration then removes the coalition-sampling error at no extra cost.n_permutations(
integer(1):NULL) Number of permutations forestimator = "permutation". Each permutation evaluates one coalition per feature, so the cost is1 + n_permutations * n_featuresevaluated coalitions. If unset, defaults to10L.n_coalitions(
integer(1):NULL) Number of paired coalition draws forestimator = "kernel". Each draw evaluates a coalition and its complement on one test observation, so the cost is2 + 2 * n_coalitionsevaluated coalitions. If unset, defaults to2048L. Check whether the budget suffices with$convergence()or$plot_convergence(), or letearly_stoppingdecide.max_features(
integer(1):12L) Cap on the number of features forestimator = "exact", whose cost grows as2^n_features; construction aborts above it.batch_size(
integer(1):5000L) Maximum number of observations to process in a single prediction call.n_samples(
integer(1):100L) Number of samples to use for marginalizing out-of-coalition features. For MarginalSAGE, this is the number of marginal data samples ("background data" in other implementations). For ConditionalSAGE, this is the number of conditional samples per test instance retrieved fromsampler.early_stopping(
logical(1):FALSE) Whether to stop once the convergence criterion is met, rather than spending the full budget. Applies to the permutation and kernel estimators; setting it forestimator = "exact"is a warning. The budget then acts as an upper bound: if the criterion is not met within it, the values are returned with a warning and$budgetreportsconverged = FALSE.se_threshold(
numeric(1):0.025) Convergence threshold for relative standard error. Convergence is detected when the maximum relative SE across all features falls below this threshold. Relative SE is calculated as SE divided by the range of importance values (max - min), making it scale-invariant across different loss metrics. The default of0.025(convergence once the relative SE is below 2.5% of the importance range) is the default of the Pythonsagepackage; the examples in Covert et al. (2020) and Covert & Lee (2021) use0.01to0.02. The same threshold buys different budgets across estimators, since their standard errors are constructed differently (see Details).min_permutations(
integer(1):10L) Minimum permutations before checking for convergence. Convergence is judged based on the standard errors of the estimated SAGE values, which requires a sufficiently large number of samples (i.e., evaluated coalitions). Permutation estimator only; the kernel estimator checks after every chunk of 512 draws.check_interval(
integer(1):1L) Check convergence every N permutations. Permutation estimator only.
SAGE$compute()
Compute SAGE values.
Usage
SAGE$compute(
store_backends = TRUE,
batch_size = NULL,
early_stopping = NULL,
se_threshold = NULL,
min_permutations = NULL,
check_interval = NULL
)Arguments
store_backends(
logical(1)) Whether to store data backends.batch_size(
integer(1):5000L) Maximum number of observations to process in a single prediction call.early_stopping(
logical(1):FALSE) Whether to check for convergence and stop early.se_threshold(
numeric(1):0.025) Convergence threshold for relative standard error. SE is normalized by the range of importance values (max - min) to make convergence detection scale-invariant. Default0.025means convergence when relative SE < 2.5%.min_permutations(
integer(1):10L) Minimum permutations before checking convergence.check_interval(
integer(1):1L) Check convergence every N permutations. The convergence arguments only apply toestimator = "permutation"; passing them for another estimator is a warning.
SAGE$reset()
Resets all stored fields populated by $compute(), including the convergence tracking
($convergence_history, $converged, $budget).
SAGE$convergence()
Monte Carlo standard errors of the final SAGE estimates, i.e. the last checkpoint of
$convergence_history, together with the convergence ratio max(se) / (max(importance) - min(importance))
that early_stopping compares against se_threshold.
These quantify how converged the computation is for the fixed model, not feature importance;
see the Standard errors section in Details.
Returns
A data.table with columns feature, importance, se, and ratio
(the same value in every row), or NULL before $compute() and for the exact estimator.