Skip to main content

shap_rs/explainers/
kernel.rs

1use crate::{
2    coalition, evaluation::CoalitionEvaluator, Background, EvaluationConfig, Explainer,
3    Explanation, IndependentMasker, Link, Masker, Predict, Result, ShapError,
4};
5use ndarray::{Array2, Array3, ArrayView2};
6use rand::{rngs::StdRng, Rng, SeedableRng};
7use std::collections::BTreeSet;
8
9/// Linear solver used by Kernel SHAP's constrained weighted least squares.
10#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
11pub enum KernelSolver {
12    /// Fast, allocation-light solve of the normal equations.
13    #[default]
14    NormalEquations,
15    /// Householder QR on the weighted design matrix. This uses more memory but
16    /// avoids squaring the design matrix's condition number.
17    HouseholderQr,
18}
19
20/// Kernel SHAP with Shapley-kernel weighted least squares and an exact
21/// efficiency constraint. Sampled coalitions are complement-paired.
22pub struct KernelExplainer<M, K = IndependentMasker> {
23    model: M,
24    masker: K,
25    nsamples: usize,
26    seed: u64,
27    exact_threshold: usize,
28    ridge: f64,
29    solver: KernelSolver,
30    evaluation: EvaluationConfig,
31    link: Link,
32}
33impl<M> KernelExplainer<M, IndependentMasker> {
34    pub fn new(model: M, background: Background) -> Self {
35        Self::from_masker(model, IndependentMasker::new(background))
36    }
37}
38impl<M, K> KernelExplainer<M, K> {
39    pub fn from_masker(model: M, masker: K) -> Self {
40        Self {
41            model,
42            masker,
43            nsamples: 512,
44            seed: 0,
45            exact_threshold: 12,
46            ridge: 1e-10,
47            solver: KernelSolver::NormalEquations,
48            evaluation: EvaluationConfig::default(),
49            link: Link::Identity,
50        }
51    }
52    pub fn with_nsamples(mut self, n: usize) -> Self {
53        self.nsamples = n;
54        self
55    }
56    pub fn with_seed(mut self, s: u64) -> Self {
57        self.seed = s;
58        self
59    }
60    pub fn with_exact_threshold(mut self, n: usize) -> Self {
61        self.exact_threshold = n;
62        self
63    }
64    pub fn with_ridge(mut self, ridge: f64) -> Self {
65        self.ridge = ridge;
66        self
67    }
68    pub fn with_solver(mut self, solver: KernelSolver) -> Self {
69        self.solver = solver;
70        self
71    }
72    pub fn with_evaluation_config(mut self, config: EvaluationConfig) -> Self {
73        self.evaluation = config;
74        self
75    }
76    pub fn with_link(mut self, link: Link) -> Self {
77        self.link = link;
78        self
79    }
80    fn coalitions(&self, m: usize) -> Result<Vec<u64>> {
81        if self.nsamples == 0 {
82            return Err(ShapError::InvalidConfiguration(
83                "nsamples must be positive".into(),
84            ));
85        }
86        if m < 63 && m <= self.exact_threshold {
87            let count = usize::try_from((1u64 << m) - 2).map_err(|_| {
88                ShapError::InvalidConfiguration(
89                    "exact Kernel SHAP coalition count exceeds usize".into(),
90                )
91            })?;
92            crate::error::checked_f64_shape(&[count], "Kernel SHAP coalition set")?;
93            return Ok((1..(1u64 << m) - 1).collect());
94        }
95        let full = if m < 64 {
96            u64::MAX >> (64 - m)
97        } else {
98            u64::MAX
99        };
100        let target = self.nsamples.min(if m < 63 {
101            usize::try_from((1u64 << m) - 2).unwrap_or(usize::MAX)
102        } else {
103            usize::MAX
104        });
105        crate::error::checked_f64_shape(&[target], "Kernel SHAP coalition set")?;
106        let mut set = BTreeSet::new();
107        let mut rng = StdRng::seed_from_u64(self.seed);
108        while set.len() + 2 <= target {
109            let z = rng.gen::<u64>() & full;
110            let complement = full ^ z;
111            if z != 0 && z != full && !set.contains(&z) && !set.contains(&complement) {
112                set.insert(z);
113                set.insert(complement);
114            }
115        }
116        while set.len() < target {
117            let z = rng.gen::<u64>() & full;
118            if z != 0 && z != full {
119                set.insert(z);
120            }
121        }
122        Ok(set.into_iter().collect())
123    }
124}
125impl<M: Predict, K: Masker> Explainer for KernelExplainer<M, K> {
126    fn explain(&self, x: ArrayView2<'_, f64>) -> Result<Explanation> {
127        let m = self.masker.n_features();
128        if x.nrows() == 0 {
129            return Err(ShapError::EmptyData);
130        }
131        if m >= 63 {
132            return Err(ShapError::InvalidConfiguration(
133                "Kernel SHAP currently supports at most 62 features".into(),
134            ));
135        }
136        if !self.ridge.is_finite() || self.ridge < 0.0 {
137            return Err(ShapError::InvalidConfiguration(
138                "ridge must be finite and non-negative".into(),
139            ));
140        }
141        if x.ncols() != self.masker.n_input_features() {
142            return Err(ShapError::DimensionMismatch {
143                expected: format!("{} input features", self.masker.n_input_features()),
144                found: format!("{}", x.ncols()),
145            });
146        }
147        let masks = self.coalitions(m)?;
148        let full_mask = (1u64 << m) - 1;
149        let mut first_eval = CoalitionEvaluator::new(&self.model, &self.masker, self.evaluation)?;
150        let probe = first_eval.evaluate(x.row(0), &[0])?.remove(0);
151        let o = probe.len();
152        crate::error::checked_f64_shape(&[x.nrows(), m, o], "kernel explanation")?;
153        let mut v = Array3::zeros((x.nrows(), m, o));
154        let mut bases = Array2::zeros((x.nrows(), o));
155        for n in 0..x.nrows() {
156            let mut requested = Vec::with_capacity(masks.len() + 2);
157            requested.push(0);
158            requested.push(full_mask);
159            requested.extend_from_slice(&masks);
160            let mut evaluator =
161                CoalitionEvaluator::new(&self.model, &self.masker, self.evaluation)?;
162            let evaluated = evaluator.evaluate(x.row(n), &requested)?;
163            let base = evaluated[0]
164                .iter()
165                .map(|&z| self.link.forward(z))
166                .collect::<Result<Vec<_>>>()?;
167            let full = evaluated[1]
168                .iter()
169                .map(|&z| self.link.forward(z))
170                .collect::<Result<Vec<_>>>()?;
171            for k in 0..o {
172                bases[[n, k]] = base[k]
173            }
174            if m == 1 {
175                for k in 0..o {
176                    v[[n, 0, k]] = full[k] - base[k]
177                }
178                continue;
179            }
180            let p = m - 1;
181            crate::error::checked_f64_shape(&[p, p], "Kernel SHAP linear system")?;
182            crate::error::checked_f64_shape(&[p, o], "Kernel SHAP right-hand side")?;
183            let qr_rows = masks.len().checked_add(p).ok_or_else(|| {
184                ShapError::InvalidConfiguration("Kernel SHAP QR row count overflow".into())
185            })?;
186            if self.solver == KernelSolver::HouseholderQr {
187                crate::error::checked_f64_shape(&[qr_rows, p], "Kernel SHAP QR design")?;
188                crate::error::checked_f64_shape(&[qr_rows, o], "Kernel SHAP QR response")?;
189            }
190            let mut a = vec![vec![0.; p]; p];
191            let mut b = vec![vec![0.; o]; p];
192            let mut qr_a = Vec::with_capacity(qr_rows);
193            let mut qr_b = Vec::with_capacity(qr_rows);
194            for (row, &mask) in masks.iter().enumerate() {
195                let z = coalition::members(mask, m);
196                let y = evaluated[row + 2]
197                    .iter()
198                    .map(|&value| self.link.forward(value))
199                    .collect::<Result<Vec<_>>>()?;
200                let w = coalition::kernel_weight(m, mask.count_ones() as usize);
201                let use_qr = self.solver == KernelSolver::HouseholderQr;
202                let sqrt_weight = w.sqrt();
203                let mut design_row = if use_qr { vec![0.0; p] } else { Vec::new() };
204                let mut response_row = if use_qr { vec![0.0; o] } else { Vec::new() };
205                for i in 0..p {
206                    let xi = (z[i] as u8 as f64) - (z[m - 1] as u8 as f64);
207                    if use_qr {
208                        design_row[i] = sqrt_weight * xi;
209                    }
210                    for j in 0..p {
211                        a[i][j] += w * xi * ((z[j] as u8 as f64) - (z[m - 1] as u8 as f64))
212                    }
213                    for k in 0..o {
214                        let target = y[k] - base[k] - (z[m - 1] as u8 as f64) * (full[k] - base[k]);
215                        b[i][k] += w * xi * target;
216                        if use_qr {
217                            response_row[k] = sqrt_weight * target;
218                        }
219                    }
220                }
221                if use_qr {
222                    qr_a.push(design_row);
223                    qr_b.push(response_row);
224                }
225            }
226            let beta = match self.solver {
227                KernelSolver::NormalEquations => {
228                    for (i, row) in a.iter_mut().enumerate().take(p) {
229                        row[i] += self.ridge
230                    }
231                    solve(a, b)?
232                }
233                KernelSolver::HouseholderQr => {
234                    if self.ridge > 0.0 {
235                        let scale = self.ridge.sqrt();
236                        for column in 0..p {
237                            let mut row = vec![0.0; p];
238                            row[column] = scale;
239                            qr_a.push(row);
240                            qr_b.push(vec![0.0; o]);
241                        }
242                    }
243                    solve_qr(qr_a, qr_b, p)?
244                }
245            };
246            for k in 0..o {
247                let mut sum = 0.;
248                for j in 0..p {
249                    v[[n, j, k]] = beta[j][k];
250                    sum += beta[j][k]
251                }
252                v[[n, m - 1, k]] = full[k] - base[k] - sum
253            }
254        }
255        Explanation::new(v, bases, self.masker.attribution_data(x)?)
256    }
257}
258#[allow(clippy::needless_range_loop)]
259fn solve(mut a: Vec<Vec<f64>>, mut b: Vec<Vec<f64>>) -> Result<Vec<Vec<f64>>> {
260    let n = a.len();
261    let o = b[0].len();
262    for c in 0..n {
263        let p = (c..n)
264            .max_by(|&i, &j| a[i][c].abs().total_cmp(&a[j][c].abs()))
265            .unwrap();
266        if a[p][c].abs() < 1e-14 {
267            return Err(ShapError::SolverError(
268                "singular Kernel SHAP design; increase nsamples or ridge".into(),
269            ));
270        }
271        a.swap(c, p);
272        b.swap(c, p);
273        let d = a[c][c];
274        for j in c..n {
275            a[c][j] /= d
276        }
277        for k in 0..o {
278            b[c][k] /= d
279        }
280        for i in 0..n {
281            if i == c {
282                continue;
283            }
284            let f = a[i][c];
285            for j in c..n {
286                a[i][j] -= f * a[c][j]
287            }
288            for k in 0..o {
289                b[i][k] -= f * b[c][k]
290            }
291        }
292    }
293    Ok(b)
294}
295
296#[allow(clippy::needless_range_loop)]
297fn solve_qr(mut a: Vec<Vec<f64>>, mut b: Vec<Vec<f64>>, columns: usize) -> Result<Vec<Vec<f64>>> {
298    let rows = a.len();
299    if rows < columns || columns == 0 || b.len() != rows {
300        return Err(ShapError::SolverError(
301            "Kernel SHAP QR design is underdetermined".into(),
302        ));
303    }
304    let outputs = b.first().map_or(0, Vec::len);
305    if outputs == 0
306        || a.iter().any(|row| row.len() != columns)
307        || b.iter().any(|row| row.len() != outputs)
308    {
309        return Err(ShapError::SolverError(
310            "Kernel SHAP QR design is ragged or empty".into(),
311        ));
312    }
313    for column in 0..columns {
314        let norm = a[column..]
315            .iter()
316            .map(|row| row[column])
317            .fold(0.0_f64, f64::hypot);
318        if !norm.is_finite() || norm == 0.0 {
319            return Err(ShapError::SolverError(
320                "rank-deficient Kernel SHAP design; increase nsamples or ridge".into(),
321            ));
322        }
323        let alpha = if a[column][column] >= 0.0 {
324            -norm
325        } else {
326            norm
327        };
328        let mut reflector = a[column..]
329            .iter()
330            .map(|row| row[column])
331            .collect::<Vec<_>>();
332        reflector[0] -= alpha;
333        let reflector_norm = reflector.iter().copied().fold(0.0_f64, f64::hypot);
334        if !reflector_norm.is_finite() || reflector_norm == 0.0 {
335            return Err(ShapError::SolverError(
336                "failed to construct Kernel SHAP QR reflector".into(),
337            ));
338        }
339        for value in &mut reflector {
340            *value /= reflector_norm;
341        }
342        for target_column in column..columns {
343            let projection = (column..rows)
344                .map(|row| reflector[row - column] * a[row][target_column])
345                .sum::<f64>();
346            for row in column..rows {
347                a[row][target_column] -= 2.0 * reflector[row - column] * projection;
348            }
349        }
350        for output in 0..outputs {
351            let projection = (column..rows)
352                .map(|row| reflector[row - column] * b[row][output])
353                .sum::<f64>();
354            for row in column..rows {
355                b[row][output] -= 2.0 * reflector[row - column] * projection;
356            }
357        }
358        a[column][column] = alpha;
359        for row in column + 1..rows {
360            a[row][column] = 0.0;
361        }
362    }
363    let scale = (0..columns)
364        .map(|index| a[index][index].abs())
365        .fold(0.0_f64, f64::max);
366    let tolerance = f64::EPSILON * rows.max(columns) as f64 * scale.max(1.0);
367    let mut solution = vec![vec![0.0; outputs]; columns];
368    for row in (0..columns).rev() {
369        if a[row][row].abs() <= tolerance {
370            return Err(ShapError::SolverError(
371                "rank-deficient Kernel SHAP design; increase nsamples or ridge".into(),
372            ));
373        }
374        for output in 0..outputs {
375            let remainder = (row + 1..columns)
376                .map(|column| a[row][column] * solution[column][output])
377                .sum::<f64>();
378            solution[row][output] = (b[row][output] - remainder) / a[row][row];
379        }
380    }
381    if solution.iter().flatten().any(|value| !value.is_finite()) {
382        return Err(ShapError::SolverError(
383            "Kernel SHAP QR solution is non-finite".into(),
384        ));
385    }
386    Ok(solution)
387}
388#[cfg(test)]
389mod tests {
390    use super::*;
391    use crate::explainers::ExactExplainer;
392    use crate::{metrics::check_additivity, FixedMasker, FnModel};
393    use ndarray::{array, Array2, Axis};
394
395    #[test]
396    fn kernel_wls_recovers_linear_shap_values() {
397        let model = FnModel::new(|x: ArrayView2<'_, f64>| {
398            Ok(x.map_axis(Axis(1), |r| 2.0 * r[0] - 3.0 * r[1] + r[2])
399                .insert_axis(Axis(1)))
400        });
401        let background = Background::new(array![[0., 0., 0.], [2., 2., 2.]]).unwrap();
402        let x = array![[3., 4., 5.]];
403        let explanation = KernelExplainer::new(model, background)
404            .explain(x.view())
405            .unwrap();
406        assert!((explanation.values()[[0, 0, 0]] - 4.0).abs() < 1e-7);
407        assert!((explanation.values()[[0, 1, 0]] + 9.0).abs() < 1e-7);
408        assert!((explanation.values()[[0, 2, 0]] - 4.0).abs() < 1e-7);
409        check_additivity(&explanation, array![[-1.]].view(), 1e-9).unwrap();
410    }
411    #[test]
412    fn accepts_custom_maskers() {
413        let model =
414            FnModel::new(|x: ArrayView2<'_, f64>| Ok(x.sum_axis(Axis(1)).insert_axis(Axis(1))));
415        let masker = FixedMasker::new(array![0., 0.]).unwrap();
416        let e = KernelExplainer::from_masker(model, masker)
417            .explain(array![[2., 3.]].view())
418            .unwrap();
419        assert!((e.values().sum() - 5.).abs() < 1e-9);
420    }
421    #[test]
422    fn logit_link_explains_log_odds() {
423        let model =
424            FnModel::new(|x: ArrayView2<'_, f64>| Ok(x.column(0).mapv(|z| z).insert_axis(Axis(1))));
425        let e = KernelExplainer::from_masker(model, FixedMasker::new(array![0.5]).unwrap())
426            .with_link(Link::Logit)
427            .explain(array![[0.8]].view())
428            .unwrap();
429        assert!((e.base_values()[[0, 0]]).abs() < 1e-12);
430        assert!((e.values()[[0, 0, 0]] - 4f64.ln()).abs() < 1e-12);
431    }
432
433    #[test]
434    fn exact_coalitions_match_exact_shap_for_nonlinear_multi_output_model() {
435        fn predict(x: ArrayView2<'_, f64>) -> Result<Array2<f64>> {
436            Ok(Array2::from_shape_fn((x.nrows(), 2), |(i, output)| {
437                let r = x.row(i);
438                match output {
439                    0 => r[0] * r[1] + r[2].sin() - 0.5 * r[3].powi(2),
440                    _ => (r[0] - r[2]) * (r[1] + r[3]) + r[0].exp(),
441                }
442            }))
443        }
444
445        let background = Background::new(array![
446            [0.0, -1.0, 0.5, 2.0],
447            [1.0, 0.5, -0.5, -1.0],
448            [-2.0, 1.5, 1.0, 0.25]
449        ])
450        .unwrap();
451        let samples = array![[0.25, 2.0, -1.0, 0.75], [1.5, -0.25, 0.3, -2.0]];
452
453        let exact = ExactExplainer::new(FnModel::new(predict), background.clone())
454            .explain(samples.view())
455            .unwrap();
456        let kernel = KernelExplainer::new(FnModel::new(predict), background)
457            .with_exact_threshold(4)
458            .with_ridge(0.0)
459            .explain(samples.view())
460            .unwrap();
461
462        for (actual, expected) in kernel.values().iter().zip(exact.values()) {
463            assert!((actual - expected).abs() < 1e-9, "{actual} != {expected}");
464        }
465        for (actual, expected) in kernel.base_values().iter().zip(exact.base_values()) {
466            assert!((actual - expected).abs() < 1e-12);
467        }
468    }
469
470    #[test]
471    fn householder_qr_kernel_matches_exact_multi_output_values() {
472        fn predict(x: ArrayView2<'_, f64>) -> Result<Array2<f64>> {
473            Ok(Array2::from_shape_fn((x.nrows(), 2), |(row, output)| {
474                let values = x.row(row);
475                if output == 0 {
476                    values[0] * values[1] + values[2]
477                } else {
478                    values[0] - values[1] * values[2]
479                }
480            }))
481        }
482        let background = Background::new(array![[0., 0., 0.], [1., -1., 2.]]).unwrap();
483        let samples = array![[2., 3., -1.]];
484        let exact = ExactExplainer::new(FnModel::new(predict), background.clone())
485            .explain(samples.view())
486            .unwrap();
487        let kernel = KernelExplainer::new(FnModel::new(predict), background)
488            .with_solver(KernelSolver::HouseholderQr)
489            .with_ridge(0.0)
490            .explain(samples.view())
491            .unwrap();
492        for (actual, expected) in kernel.values().iter().zip(exact.values()) {
493            assert!((actual - expected).abs() < 1e-9, "{actual} != {expected}");
494        }
495    }
496
497    #[test]
498    fn householder_qr_solves_overdetermined_multi_output_system() {
499        let design = vec![
500            vec![1.0, 1.0],
501            vec![1.0, 1.0 + 1e-8],
502            vec![1.0, 1.0 - 1e-8],
503            vec![1.0, -1.0],
504        ];
505        let expected = [[2.0, -1.0], [-3.0, 4.0]];
506        let response = design
507            .iter()
508            .map(|row| {
509                (0..2)
510                    .map(|output| row[0] * expected[0][output] + row[1] * expected[1][output])
511                    .collect::<Vec<_>>()
512            })
513            .collect::<Vec<_>>();
514        let solution = solve_qr(design, response, 2).unwrap();
515        for row in 0..2 {
516            for output in 0..2 {
517                assert!((solution[row][output] - expected[row][output]).abs() < 1e-9);
518            }
519        }
520    }
521
522    #[test]
523    fn householder_qr_rejects_rank_deficient_design() {
524        assert!(solve_qr(
525            vec![vec![1.0, 1.0], vec![2.0, 2.0]],
526            vec![vec![1.0], vec![2.0]],
527            2,
528        )
529        .is_err());
530    }
531}