The goal of xplainfi is to collect common feature importance methods
under a unified and extensible interface.
It is built around mlr3 as available abstractions for learners, tasks, measures, etc. greatly simplify the implementation of importance measures.
Install xplainfi from CRAN:
install.packages("xplainfi")Or install from R-universe:
install.packages("xplainfi", repos = c("https://mlr-org.r-universe.dev", "https://cloud.r-project.org"))The latest development version of xplainfi can be installed with
pak:
# install.packages(pak)
pak::pak("mlr-org/xplainfi")Here is a basic example on how to calculate PFI for an untrained learner
and task, using cross-validation for resampling and computing PFI within
each resampling iteration 10 times on the friedman1 task (see
?mlbench::mlbench.friedman1).
The friedman1 task has the following structure:
Where important1 through important5 in
the Task, with additional numbered unimportant features without
effect on
library(xplainfi)
library(mlr3learners)
#> Loading required package: mlr3
task = tgen("friedman1")$generate(1000)
learner = lrn("regr.ranger", num.trees = 100)
measure = msr("regr.mse")
pfi = PFI$new(
task = task,
learner = learner,
measure = measure,
resampling = rsmp("cv", folds = 3),
n_repeats = 30
)Compute and print PFI scores:
pfi$compute()
pfi$importance()
#> Key: <feature>
#> feature importance
#> <char> <num>
#> 1: important1 8.130730121
#> 2: important2 7.587050771
#> 3: important3 1.603608069
#> 4: important4 12.547878920
#> 5: important5 2.816002479
#> 6: unimportant1 0.026536401
#> 7: unimportant2 0.002251584
#> 8: unimportant3 -0.039315390
#> 9: unimportant4 -0.057172791
#> 10: unimportant5 0.056256867If it aids interpretation, importances can also be calculated as the ratio rather than the difference between the baseline and post-permutation losses:
pfi$importance(relation = "ratio")
#> Key: <feature>
#> feature importance
#> <char> <num>
#> 1: important1 2.6880636
#> 2: important2 2.5822356
#> 3: important3 1.3366899
#> 4: important4 3.6191329
#> 5: important5 1.5878385
#> 6: unimportant1 1.0059545
#> 7: unimportant2 1.0004247
#> 8: unimportant3 0.9917426
#> 9: unimportant4 0.9881006
#> 10: unimportant5 1.0116920When PFI is computed based on resampling with multiple iterations, and /
or multiple permutation iterations, the individual scores can be
retrieved as a data.table:
str(pfi$scores())
#> Classes 'data.table' and 'data.frame': 900 obs. of 6 variables:
#> $ feature : chr "important1" "important1" "important1" "important1" ...
#> $ iter_rsmp : int 1 1 1 1 1 1 1 1 1 1 ...
#> $ iter_repeat : int 1 2 3 4 5 6 7 8 9 10 ...
#> $ regr.mse_baseline: num 4.56 4.56 4.56 4.56 4.56 ...
#> $ regr.mse_post : num 12.8 12.9 12.6 12.5 11.3 ...
#> $ importance : num 8.28 8.35 8.08 7.96 6.73 ...
#> - attr(*, ".internal.selfref")=<pointer: 0x105bd1420>Where iter_rsmp corresponds to the resampling iteration, i.e., 3 for
3-fold cross-validation, and iter_repeat corresponds to the
permutation iteration within each resampling iteration, 5 in this case.
While pfi$importance() contains the means across all iterations,
pfi$scores() allows you to manually visualize or aggregate them in any
way you see fit.
For example:
library(ggplot2)
ggplot(
pfi$scores(),
aes(x = importance, y = reorder(feature, importance))
) +
geom_boxplot(color = "#f44560", fill = alpha("#f44560", 0.4)) +
labs(
title = "Permutation Feature Importance on Friedman1",
subtitle = "Computed over 3-fold CV with 5 permutations per iteration using Random Forest",
x = "Importance",
y = "Feature"
) +
theme_minimal(base_size = 16) +
theme(
plot.title.position = "plot",
panel.grid.major.y = element_blank()
)If the measure in question needs to be maximized rather than minimized
(like $minimize property of the measure and calculates
importances such that the intuition “performance improvement” ->
“higher importance score” still holds:
pfi = PFI$new(
task = task,
learner = learner,
measure = msr("regr.rsq")
)
#> ℹ No <Resampling> provided, using `resampling = rsmp("holdout", ratio = 2/3)`
#> (test set size: 333)
pfi$compute()
pfi$importance()
#> Key: <feature>
#> feature importance
#> <char> <num>
#> 1: important1 0.3144810393
#> 2: important2 0.3134162904
#> 3: important3 0.0620583655
#> 4: important4 0.5415552397
#> 5: important5 0.1350006512
#> 6: unimportant1 -0.0010198832
#> 7: unimportant2 0.0010749042
#> 8: unimportant3 -0.0021041761
#> 9: unimportant4 0.0006529191
#> 10: unimportant5 0.0009255579See vignette("xplainfi") for more examples.
