Skip to main content

shap_rs/plot/
interaction.rs

1use crate::{interactions::InteractionExplanation, Result, ShapError};
2use ndarray::Array2;
3
4#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
5pub struct InteractionHeatmapData {
6    pub sample: usize,
7    pub output: usize,
8    pub values: Array2<f64>,
9    pub feature_names: Option<Vec<String>>,
10}
11
12#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
13pub struct InteractionDependencePoint {
14    pub sample: usize,
15    pub feature_value: f64,
16    pub interaction_value: f64,
17    pub color_value: Option<f64>,
18}
19
20/// Returns one symmetric feature-by-feature interaction matrix.
21pub fn heatmap_data(
22    explanation: &InteractionExplanation,
23    sample: usize,
24    output: usize,
25) -> Result<InteractionHeatmapData> {
26    explanation.validate()?;
27    validate(explanation, sample, 0, output)?;
28    Ok(InteractionHeatmapData {
29        sample,
30        output,
31        values: Array2::from_shape_fn(
32            (explanation.n_features(), explanation.n_features()),
33            |(first, second)| explanation.values()[[sample, first, second, output]],
34        ),
35        feature_names: explanation.feature_names().map(<[String]>::to_vec),
36    })
37}
38
39/// Returns interaction strength against a feature value over all samples.
40pub fn dependence_data(
41    explanation: &InteractionExplanation,
42    feature: usize,
43    interacting_feature: usize,
44    output: usize,
45    color_feature: Option<usize>,
46) -> Result<Vec<InteractionDependencePoint>> {
47    explanation.validate()?;
48    validate(explanation, 0, feature, output)?;
49    if interacting_feature >= explanation.n_features() {
50        return Err(ShapError::InvalidFeatureIndex {
51            index: interacting_feature,
52            n_features: explanation.n_features(),
53        });
54    }
55    if let Some(index) = color_feature.filter(|&index| index >= explanation.n_features()) {
56        return Err(ShapError::InvalidFeatureIndex {
57            index,
58            n_features: explanation.n_features(),
59        });
60    }
61    Ok((0..explanation.n_samples())
62        .map(|sample| InteractionDependencePoint {
63            sample,
64            feature_value: explanation.data()[[sample, feature]],
65            interaction_value: explanation.values()[[sample, feature, interacting_feature, output]],
66            color_value: color_feature.map(|index| explanation.data()[[sample, index]]),
67        })
68        .collect())
69}
70
71fn validate(
72    explanation: &InteractionExplanation,
73    sample: usize,
74    feature: usize,
75    output: usize,
76) -> Result<()> {
77    if sample >= explanation.n_samples() {
78        return Err(ShapError::InvalidSampleIndex {
79            index: sample,
80            n_samples: explanation.n_samples(),
81        });
82    }
83    if feature >= explanation.n_features() {
84        return Err(ShapError::InvalidFeatureIndex {
85            index: feature,
86            n_features: explanation.n_features(),
87        });
88    }
89    if output >= explanation.n_outputs() {
90        return Err(ShapError::InvalidOutputIndex {
91            index: output,
92            n_outputs: explanation.n_outputs(),
93        });
94    }
95    Ok(())
96}
97
98#[cfg(test)]
99mod tests {
100    use super::*;
101    use crate::FeatureMetadata;
102    use ndarray::{array, Array4};
103
104    #[test]
105    fn creates_metadata_aware_interaction_plot_data() {
106        let explanation = InteractionExplanation::new(
107            Array4::from_shape_vec((2, 2, 2, 1), vec![1., 0.5, 0.5, 2., 3., 1., 1., 4.]).unwrap(),
108            array![[0.], [0.]],
109            array![[10., 20.], [30., 40.]],
110        )
111        .unwrap()
112        .with_feature_metadata(FeatureMetadata::new(vec!["a".into(), "b".into()]).unwrap())
113        .unwrap();
114        let heatmap = heatmap_data(&explanation, 1, 0).unwrap();
115        assert_eq!(heatmap.values[[0, 1]], 1.0);
116        assert_eq!(heatmap.feature_names.unwrap(), ["a", "b"]);
117        let points = dependence_data(&explanation, 0, 1, 0, Some(1)).unwrap();
118        assert_eq!(points[1].feature_value, 30.0);
119        assert_eq!(points[1].color_value, Some(40.0));
120    }
121}