Skip to main content

shap_rs/explainers/
permutation.rs

1use crate::{
2    evaluation::CoalitionEvaluator, Background, EvaluationConfig, Explainer, Explanation,
3    IndependentMasker, Link, Masker, Predict, Result, ShapError,
4};
5use ndarray::{Array2, Array3, ArrayView2};
6use rand::{rngs::StdRng, seq::SliceRandom, SeedableRng};
7/// Monte-Carlo Shapley estimator using random feature permutations.
8pub struct PermutationExplainer<M, K = IndependentMasker> {
9    model: M,
10    masker: K,
11    n_permutations: usize,
12    seed: u64,
13    antithetic: bool,
14    link: Link,
15    evaluation: EvaluationConfig,
16}
17impl<M> PermutationExplainer<M, IndependentMasker> {
18    pub fn new(model: M, background: Background) -> Self {
19        Self::from_masker(model, IndependentMasker::new(background))
20    }
21}
22impl<M, K> PermutationExplainer<M, K> {
23    pub fn from_masker(model: M, masker: K) -> Self {
24        Self {
25            model,
26            masker,
27            n_permutations: 128,
28            seed: 0,
29            antithetic: true,
30            link: Link::Identity,
31            evaluation: EvaluationConfig {
32                coalition_batch_size: 64,
33                cache_capacity: 65536,
34                max_model_rows: None,
35            },
36        }
37    }
38    pub fn with_n_permutations(mut self, n: usize) -> Self {
39        self.n_permutations = n;
40        self
41    }
42    pub fn with_seed(mut self, s: u64) -> Self {
43        self.seed = s;
44        self
45    }
46    /// Enables reverse-order pairing to reduce Monte Carlo variance.
47    pub fn with_antithetic(mut self, enabled: bool) -> Self {
48        self.antithetic = enabled;
49        self
50    }
51    pub fn with_link(mut self, link: Link) -> Self {
52        self.link = link;
53        self
54    }
55    pub fn with_evaluation_config(mut self, config: EvaluationConfig) -> Self {
56        self.evaluation = config;
57        self
58    }
59}
60impl<M: Predict, K: Masker> Explainer for PermutationExplainer<M, K> {
61    fn explain(&self, x: ArrayView2<'_, f64>) -> Result<Explanation> {
62        let m = self.masker.n_features();
63        if x.nrows() == 0 {
64            return Err(ShapError::EmptyData);
65        }
66        if x.ncols() != self.masker.n_input_features() {
67            return Err(ShapError::DimensionMismatch {
68                expected: format!("{} input features", self.masker.n_input_features()),
69                found: format!("{}", x.ncols()),
70            });
71        }
72        if self.n_permutations == 0 {
73            return Err(ShapError::InvalidConfiguration(
74                "n_permutations must be positive".into(),
75            ));
76        }
77        if m >= 63 {
78            return Err(ShapError::InvalidConfiguration(
79                "permutation SHAP currently supports at most 62 features".into(),
80            ));
81        }
82        let mut probe_evaluator =
83            CoalitionEvaluator::new(&self.model, &self.masker, self.evaluation)?;
84        let o = probe_evaluator.evaluate(x.row(0), &[0])?[0].len();
85        crate::error::checked_f64_shape(&[x.nrows(), m, o], "permutation explanation")?;
86        self.n_permutations.checked_mul(m).ok_or_else(|| {
87            ShapError::InvalidConfiguration("permutation step count overflowed".into())
88        })?;
89        let mut vals = Array3::zeros((x.nrows(), m, o));
90        let mut bases = Array2::zeros((x.nrows(), o));
91        for i in 0..x.nrows() {
92            let mut rng = StdRng::seed_from_u64(crate::coalition::sample_seed(self.seed, x.row(i)));
93            let mut requested = vec![0u64];
94            let mut steps = Vec::with_capacity(self.n_permutations * m);
95            let mut generated = 0;
96            while generated < self.n_permutations {
97                let mut order = (0..m).collect::<Vec<_>>();
98                order.shuffle(&mut rng);
99                append_order(&order, &mut requested, &mut steps);
100                generated += 1;
101                if self.antithetic && generated < self.n_permutations {
102                    order.reverse();
103                    append_order(&order, &mut requested, &mut steps);
104                    generated += 1;
105                }
106            }
107            let mut evaluator =
108                CoalitionEvaluator::new(&self.model, &self.masker, self.evaluation)?;
109            let evaluated = evaluator
110                .evaluate(x.row(i), &requested)?
111                .into_iter()
112                .map(|row| {
113                    row.into_iter()
114                        .map(|value| self.link.forward(value))
115                        .collect::<Result<Vec<_>>>()
116                })
117                .collect::<Result<Vec<_>>>()?;
118            let base = &evaluated[0];
119            for out in 0..o {
120                bases[[i, out]] = base[out]
121            }
122            for (j, before, after) in steps {
123                for out in 0..o {
124                    vals[[i, j, out]] += (evaluated[after][out] - evaluated[before][out])
125                        / self.n_permutations as f64
126                }
127            }
128        }
129        Explanation::new(vals, bases, self.masker.attribution_data(x)?)
130    }
131}
132
133fn append_order(order: &[usize], requested: &mut Vec<u64>, steps: &mut Vec<(usize, usize, usize)>) {
134    let mut mask = 0u64;
135    let mut before = 0usize;
136    for &feature in order {
137        mask |= 1u64 << feature;
138        requested.push(mask);
139        let after = requested.len() - 1;
140        steps.push((feature, before, after));
141        before = after;
142    }
143}
144
145#[cfg(test)]
146mod tests {
147    use super::*;
148    use crate::{FixedMasker, FnModel};
149    use ndarray::{array, Axis};
150
151    #[test]
152    fn antithetic_pair_is_exact_for_two_feature_interaction() {
153        let model = FnModel::new(|x: ArrayView2<'_, f64>| {
154            Ok(x.map_axis(Axis(1), |row| row[0] * row[1])
155                .insert_axis(Axis(1)))
156        });
157        let explanation =
158            PermutationExplainer::from_masker(model, FixedMasker::new(array![0., 0.]).unwrap())
159                .with_n_permutations(2)
160                .with_antithetic(true)
161                .explain(array![[2., 3.]].view())
162                .unwrap();
163        assert!((explanation.values()[[0, 0, 0]] - 3.).abs() < 1e-12);
164        assert!((explanation.values()[[0, 1, 0]] - 3.).abs() < 1e-12);
165    }
166
167    #[test]
168    fn odd_permutation_count_is_preserved() {
169        let model =
170            FnModel::new(|x: ArrayView2<'_, f64>| Ok(x.sum_axis(Axis(1)).insert_axis(Axis(1))));
171        let explanation =
172            PermutationExplainer::from_masker(model, FixedMasker::new(array![0., 0., 0.]).unwrap())
173                .with_n_permutations(3)
174                .explain(array![[1., 2., 3.]].view())
175                .unwrap();
176        assert_eq!(explanation.reconstructed(), array![[6.]]);
177    }
178
179    #[test]
180    fn permutation_logit_link_is_locally_accurate_in_log_odds() {
181        let model =
182            FnModel::new(|x: ArrayView2<'_, f64>| Ok(x.column(0).to_owned().insert_axis(Axis(1))));
183        let explanation =
184            PermutationExplainer::from_masker(model, FixedMasker::new(array![0.5]).unwrap())
185                .with_n_permutations(2)
186                .with_link(Link::Logit)
187                .explain(array![[0.8]].view())
188                .unwrap();
189        assert!((explanation.reconstructed()[[0, 0]] - 4f64.ln()).abs() < 1e-12);
190    }
191}