Skip to main content

shap_rs/
explainer.rs

1use crate::{Explanation, FeatureMetadata, OutputMetadata, Result};
2use ndarray::ArrayView2;
3/// Common interface implemented by SHAP explainers.
4pub trait Explainer {
5    fn explain(&self, x: ArrayView2<'_, f64>) -> Result<Explanation>;
6}
7
8/// Adds feature and output metadata to any explainer without changing its algorithm.
9pub struct MetadataExplainer<E> {
10    inner: E,
11    feature_metadata: Option<FeatureMetadata>,
12    output_metadata: Option<OutputMetadata>,
13}
14
15impl<E> MetadataExplainer<E> {
16    pub fn new(inner: E) -> Self {
17        Self {
18            inner,
19            feature_metadata: None,
20            output_metadata: None,
21        }
22    }
23    pub fn with_feature_metadata(mut self, metadata: FeatureMetadata) -> Result<Self> {
24        metadata.validate()?;
25        self.feature_metadata = Some(metadata);
26        Ok(self)
27    }
28    pub fn with_output_metadata(mut self, metadata: OutputMetadata) -> Result<Self> {
29        metadata.validate()?;
30        self.output_metadata = Some(metadata);
31        Ok(self)
32    }
33    pub fn inner(&self) -> &E {
34        &self.inner
35    }
36    pub fn into_inner(self) -> E {
37        self.inner
38    }
39}
40
41impl<E: Explainer> Explainer for MetadataExplainer<E> {
42    fn explain(&self, x: ArrayView2<'_, f64>) -> Result<Explanation> {
43        let mut explanation = self.inner.explain(x)?;
44        if let Some(metadata) = &self.feature_metadata {
45            explanation = explanation.with_feature_metadata(metadata.clone())?;
46        }
47        if let Some(metadata) = &self.output_metadata {
48            explanation = explanation.with_output_metadata(metadata.clone())?;
49        }
50        Ok(explanation)
51    }
52}
53
54/// Convenience methods available on every concrete explainer.
55pub trait ExplainerExt: Explainer + Sized {
56    fn with_metadata(self) -> MetadataExplainer<Self> {
57        MetadataExplainer::new(self)
58    }
59}
60
61impl<E: Explainer> ExplainerExt for E {}
62
63#[cfg(test)]
64mod tests {
65    use super::*;
66    use crate::explainers::ExactExplainer;
67    use crate::{Background, FeatureKind, FnModel, OutputKind};
68    use ndarray::{array, Axis};
69
70    #[test]
71    fn decorates_any_explainer_with_validated_metadata() {
72        let model =
73            FnModel::new(|x: ArrayView2<'_, f64>| Ok(x.sum_axis(Axis(1)).insert_axis(Axis(1))));
74        let explainer = ExactExplainer::new(model, Background::new(array![[0., 0.]]).unwrap())
75            .with_metadata()
76            .with_feature_metadata(
77                FeatureMetadata::new(vec!["age".into(), "income".into()])
78                    .unwrap()
79                    .with_kinds(vec![FeatureKind::Continuous, FeatureKind::Continuous])
80                    .unwrap(),
81            )
82            .unwrap()
83            .with_output_metadata(
84                OutputMetadata::new(vec!["score".into()])
85                    .unwrap()
86                    .with_kinds(vec![OutputKind::Regression])
87                    .unwrap(),
88            )
89            .unwrap();
90        let explanation = explainer.explain(array![[2., 3.]].view()).unwrap();
91        assert_eq!(explanation.feature_names().unwrap(), ["age", "income"]);
92        assert_eq!(explanation.output_names().unwrap(), ["score"]);
93    }
94
95    #[test]
96    fn reports_metadata_dimension_mismatch_at_explanation_time() {
97        let model =
98            FnModel::new(|x: ArrayView2<'_, f64>| Ok(x.sum_axis(Axis(1)).insert_axis(Axis(1))));
99        let explainer = ExactExplainer::new(model, Background::new(array![[0., 0.]]).unwrap())
100            .with_metadata()
101            .with_feature_metadata(FeatureMetadata::new(vec!["only_one".into()]).unwrap())
102            .unwrap();
103        assert!(explainer.explain(array![[2., 3.]].view()).is_err());
104    }
105}