Skip to main content

shap_rs/
lib.rs

1#![doc = include_str!("../README.md")]
2#![forbid(unsafe_code)]
3pub mod analysis;
4pub mod background;
5#[cfg(feature = "burn-adapter")]
6pub mod burn_adapter;
7pub mod coalition;
8pub mod error;
9pub mod evaluation;
10pub mod explainer;
11pub mod explainers;
12pub mod explanation;
13pub mod interactions;
14pub mod link;
15pub mod masker;
16pub mod metadata;
17pub mod metrics;
18pub mod model;
19#[cfg(feature = "parallel")]
20pub mod parallel;
21pub mod plot;
22pub mod sparse;
23pub mod tree;
24pub use background::Background;
25pub use error::{Result, ShapError};
26pub use evaluation::EvaluationConfig;
27pub use explainer::{Explainer, ExplainerExt, MetadataExplainer};
28pub use explanation::{AttributionSemantics, Explanation, UncertainExplanation};
29pub use link::Link;
30pub use masker::{
31    ConditionalTabularMasker, FixedMasker, FnMasker, FnStreamingMasker, GroupedMasker,
32    ImageBaseline, ImageMasker, IndependentMasker, InpaintingImageMasker, Masker,
33    SegmentedImageMasker, SpecialTokenPolicy, TextMasker, TextToken, TokenizedTextMasker,
34};
35pub use metadata::{FeatureKind, FeatureMetadata, OutputKind, OutputMetadata};
36pub use model::{
37    AcceleratedPredict, CachedModel, DeepAttribution, DeviceModel, DifferentiablePredict,
38    ExecutionDevice, FnAcceleratedModel, FnModel, FnOwnedAcceleratedModel, Predict,
39};
40#[cfg(feature = "parallel")]
41pub use parallel::ParallelExplainerExt;
42pub use sparse::{
43    FnSparseModel, SparseIndependentMasker, SparseMatrix, SparsePermutationExplainer,
44    SparsePredict, SparseRowView,
45};
46pub use tree::{
47    MissingBranch, MissingValuePolicy, Node, SplitComparison, Tree, TreeArrays, TreeEnsemble,
48};
49
50/// Compatibility helper for scalar-output Kernel SHAP.
51pub fn explain_sample<F>(
52    predict: F,
53    sample: &[f64],
54    background: &[Vec<f64>],
55    nsamples: usize,
56) -> Result<Vec<f64>>
57where
58    F: Fn(&[Vec<f64>]) -> Vec<f64>,
59{
60    use ndarray::{Array2, ArrayView2};
61    let m = sample.len();
62    if background.is_empty() {
63        return Err(ShapError::EmptyBackground);
64    }
65    if background.iter().any(|r| r.len() != m) {
66        return Err(ShapError::DimensionMismatch {
67            expected: format!("{m} features"),
68            found: "ragged background".into(),
69        });
70    }
71    let bg = Background::new(
72        Array2::from_shape_vec(
73            (background.len(), m),
74            background.iter().flatten().copied().collect(),
75        )
76        .map_err(|e| ShapError::Other(e.to_string()))?,
77    )?;
78    let model = FnModel::new(move |x: ArrayView2<'_, f64>| {
79        let rows = x.rows().into_iter().map(|r| r.to_vec()).collect::<Vec<_>>();
80        let y = predict(&rows);
81        if y.len() != x.nrows() {
82            return Err(ShapError::OutputDimensionMismatch {
83                expected: x.nrows(),
84                found: y.len(),
85            });
86        }
87        Array2::from_shape_vec((y.len(), 1), y).map_err(|e| ShapError::ModelError(e.to_string()))
88    });
89    let e = explainers::KernelExplainer::new(model, bg)
90        .with_nsamples(nsamples)
91        .explain(
92            Array2::from_shape_vec((1, m), sample.to_vec())
93                .unwrap()
94                .view(),
95        )?;
96    Ok((0..m).map(|j| e.values()[[0, j, 0]]).collect())
97}