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