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
20pub 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
39pub 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}