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};
7pub 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 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}