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#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
11pub enum KernelSolver {
12 #[default]
14 NormalEquations,
15 HouseholderQr,
18}
19
20pub 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}