1use crate::{Explanation, FeatureMetadata, OutputMetadata, Result};
2use ndarray::ArrayView2;
3pub trait Explainer {
5 fn explain(&self, x: ArrayView2<'_, f64>) -> Result<Explanation>;
6}
7
8pub 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
54pub 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}