1#![allow(dead_code)]
7type Result<T, E = torsh_core::error::TorshError> = std::result::Result<T, E>;
10
11use crate::{GraphData, GraphLayer};
12use std::collections::HashMap;
13use torsh_tensor::{
14 creation::{ones, zeros},
15 Tensor,
16};
17
18#[derive(Debug, Clone)]
23pub struct GraphLRP {
24 pub alpha: f32,
26 pub beta: f32,
28 pub epsilon: f32,
30 pub activations: HashMap<String, Tensor>,
32 pub relevances: HashMap<String, Tensor>,
34}
35
36impl GraphLRP {
37 pub fn new() -> Self {
39 Self {
40 alpha: 1.0,
41 beta: 0.0,
42 epsilon: 1e-6,
43 activations: HashMap::new(),
44 relevances: HashMap::new(),
45 }
46 }
47
48 pub fn with_alpha_beta(alpha: f32, beta: f32) -> Self {
50 Self {
51 alpha,
52 beta,
53 epsilon: 1e-6,
54 activations: HashMap::new(),
55 relevances: HashMap::new(),
56 }
57 }
58
59 pub fn store_activation(&mut self, layer_name: String, activation: Tensor) {
61 self.activations.insert(layer_name, activation);
62 }
63
64 pub fn compute_relevance_epsilon(
66 &self,
67 input: &Tensor,
68 _output: &Tensor,
69 weight: &Tensor,
70 output_relevance: &Tensor,
71 ) -> Result<Tensor, Box<dyn std::error::Error>> {
72 let z = input.matmul(weight)?;
76
77 let z_with_eps = self.add_epsilon_with_sign(&z)?;
79
80 let weighted_input = input.matmul(weight)?;
82 let relevance_factors = weighted_input.div(&z_with_eps)?;
83 let input_relevance = relevance_factors.mul(output_relevance)?;
84
85 Ok(input_relevance)
86 }
87
88 pub fn compute_relevance_alpha_beta(
90 &self,
91 input: &Tensor,
92 weight: &Tensor,
93 output_relevance: &Tensor,
94 ) -> Result<Tensor, Box<dyn std::error::Error>> {
95 let weight_pos = self.relu_tensor(weight)?;
99 let weight_neg = weight.sub(&weight_pos)?;
100
101 let z_pos = input.matmul(&weight_pos)?;
103 let z_neg = input.matmul(&weight_neg)?;
104
105 let alpha_contrib = z_pos.mul_scalar(self.alpha)?;
107 let beta_contrib = z_neg.mul_scalar(self.beta)?;
108 let z_combined = alpha_contrib.sub(&beta_contrib)?;
109
110 let z_with_eps = self.add_epsilon_with_sign(&z_combined)?;
112
113 let relevance_factors = z_combined.div(&z_with_eps)?;
115 let input_relevance = relevance_factors.mul(output_relevance)?;
116
117 Ok(input_relevance)
118 }
119
120 pub fn compute_graph_relevance(
122 &self,
123 graph: &GraphData,
124 node_relevance: &Tensor,
125 layer_name: &str,
126 ) -> Result<GraphRelevanceResult, Box<dyn std::error::Error>> {
127 let _num_nodes = graph.num_nodes;
128 let num_edges = graph.num_edges;
129
130 let mut edge_relevance = zeros(&[num_edges])?;
132 let node_relevance_propagated = node_relevance.clone();
133
134 let edge_data = graph.edge_index.to_vec()?;
136
137 for edge_idx in 0..num_edges {
139 let src_idx = edge_data[edge_idx] as usize;
140 let dst_idx = edge_data[edge_idx + num_edges] as usize;
141
142 let src_relevance = self.get_node_relevance(node_relevance, src_idx)?;
144 let dst_relevance = self.get_node_relevance(node_relevance, dst_idx)?;
145
146 let edge_rel = (src_relevance + dst_relevance) / 2.0;
148 edge_relevance = self.set_edge_relevance(edge_relevance, edge_idx, edge_rel)?;
149 }
150
151 Ok(GraphRelevanceResult {
152 node_relevance: node_relevance_propagated,
153 edge_relevance,
154 layer_name: layer_name.to_string(),
155 })
156 }
157
158 pub fn analyze_relevance_patterns(
160 &self,
161 _graph: &GraphData,
162 relevance_result: &GraphRelevanceResult,
163 ) -> RelevanceAnalysis {
164 let node_stats = self.compute_node_relevance_stats(&relevance_result.node_relevance);
165 let edge_stats = self.compute_edge_relevance_stats(&relevance_result.edge_relevance);
166
167 let important_nodes = self.find_important_nodes(&relevance_result.node_relevance, 0.8);
169 let important_edges = self.find_important_edges(&relevance_result.edge_relevance, 0.8);
170
171 RelevanceAnalysis {
172 node_stats: node_stats.clone(),
173 edge_stats: edge_stats.clone(),
174 important_nodes,
175 important_edges,
176 total_relevance: node_stats.sum + edge_stats.sum,
177 }
178 }
179
180 fn add_epsilon_with_sign(&self, tensor: &Tensor) -> Result<Tensor, Box<dyn std::error::Error>> {
183 let sign_tensor = self.sign_tensor(tensor)?;
185 let epsilon_term = sign_tensor.mul_scalar(self.epsilon)?;
186 Ok(tensor.add(&epsilon_term)?)
187 }
188
189 fn relu_tensor(&self, tensor: &Tensor) -> Result<Tensor, Box<dyn std::error::Error>> {
190 let zeros_tensor = zeros(tensor.shape().dims())?;
192 Ok(tensor.maximum(&zeros_tensor)?)
193 }
194
195 fn sign_tensor(&self, tensor: &Tensor) -> Result<Tensor, Box<dyn std::error::Error>> {
196 let zeros_tensor = zeros(tensor.shape().dims())?;
198 let ones_tensor = ones(tensor.shape().dims())?;
199 let neg_ones = ones_tensor.mul_scalar(-1.0)?;
200
201 let positive_mask = tensor.gt(&zeros_tensor)?;
203 let negative_mask = tensor.lt(&zeros_tensor)?;
204
205 let pos_mask_data = positive_mask.to_vec()?;
207 let neg_mask_data = negative_mask.to_vec()?;
208
209 let pos_float: Vec<f32> = pos_mask_data
210 .iter()
211 .map(|&b| if b { 1.0 } else { 0.0 })
212 .collect();
213 let neg_float: Vec<f32> = neg_mask_data
214 .iter()
215 .map(|&b| if b { 1.0 } else { 0.0 })
216 .collect();
217
218 let pos_tensor =
219 torsh_tensor::creation::from_vec(pos_float, tensor.shape().dims(), tensor.device())?;
220 let neg_tensor =
221 torsh_tensor::creation::from_vec(neg_float, tensor.shape().dims(), tensor.device())?;
222
223 let mut result = zeros(tensor.shape().dims())?;
224 result = result.add(&pos_tensor.mul(&ones_tensor)?)?;
225 result = result.add(&neg_tensor.mul(&neg_ones)?)?;
226
227 Ok(result)
228 }
229
230 fn get_node_relevance(
231 &self,
232 relevance: &Tensor,
233 node_idx: usize,
234 ) -> Result<f32, Box<dyn std::error::Error>> {
235 let node_rel = relevance.slice_tensor(0, node_idx, node_idx + 1)?;
236 let rel_vec = node_rel.to_vec()?;
237 Ok(rel_vec[0])
238 }
239
240 fn set_edge_relevance(
241 &self,
242 edge_relevance: Tensor,
243 _edge_idx: usize,
244 _value: f32,
245 ) -> Result<Tensor, Box<dyn std::error::Error>> {
246 Ok(edge_relevance)
248 }
249
250 fn compute_node_relevance_stats(&self, _relevance: &Tensor) -> RelevanceStats {
251 RelevanceStats {
253 mean: 0.0,
254 std: 0.0,
255 min: 0.0,
256 max: 0.0,
257 sum: 0.0,
258 }
259 }
260
261 fn compute_edge_relevance_stats(&self, _relevance: &Tensor) -> RelevanceStats {
262 RelevanceStats {
264 mean: 0.0,
265 std: 0.0,
266 min: 0.0,
267 max: 0.0,
268 sum: 0.0,
269 }
270 }
271
272 fn find_important_nodes(&self, _relevance: &Tensor, _threshold: f32) -> Vec<usize> {
273 Vec::new()
275 }
276
277 fn find_important_edges(&self, _relevance: &Tensor, _threshold: f32) -> Vec<usize> {
278 Vec::new()
280 }
281}
282
283impl Default for GraphLRP {
284 fn default() -> Self {
285 Self::new()
286 }
287}
288
289#[derive(Debug, Clone)]
291pub struct GraphRelevanceResult {
292 pub node_relevance: Tensor,
293 pub edge_relevance: Tensor,
294 pub layer_name: String,
295}
296
297#[derive(Debug, Clone)]
299pub struct RelevanceStats {
300 pub mean: f32,
301 pub std: f32,
302 pub min: f32,
303 pub max: f32,
304 pub sum: f32,
305}
306
307#[derive(Debug, Clone)]
309pub struct RelevanceAnalysis {
310 pub node_stats: RelevanceStats,
311 pub edge_stats: RelevanceStats,
312 pub important_nodes: Vec<usize>,
313 pub important_edges: Vec<usize>,
314 pub total_relevance: f32,
315}
316
317#[derive(Debug, Clone)]
319pub struct GraphGradientAttribution {
320 pub gradients: HashMap<String, Tensor>,
322 pub smooth_steps: usize,
324}
325
326impl GraphGradientAttribution {
327 pub fn new() -> Self {
328 Self {
329 gradients: HashMap::new(),
330 smooth_steps: 50,
331 }
332 }
333
334 pub fn integrated_gradients(
336 &self,
337 graph: &GraphData,
338 _baseline_graph: &GraphData,
339 _target_class: usize,
340 ) -> Result<GraphData, Box<dyn std::error::Error>> {
341 let integrated_features = graph.x.clone();
343 let integrated_edges = graph.edge_index.clone();
344
345 Ok(GraphData::new(integrated_features, integrated_edges))
349 }
350
351 pub fn gradient_saliency(
353 &self,
354 graph: &GraphData,
355 _target_output: &Tensor,
356 ) -> Result<Tensor, Box<dyn std::error::Error>> {
357 Ok(zeros(&[graph.num_nodes])?)
360 }
361}
362
363impl Default for GraphGradientAttribution {
364 fn default() -> Self {
365 Self::new()
366 }
367}
368
369pub struct GraphExplainer {
371 pub lrp: GraphLRP,
372 pub gradient_attribution: GraphGradientAttribution,
373}
374
375impl GraphExplainer {
376 pub fn new() -> Self {
377 Self {
378 lrp: GraphLRP::new(),
379 gradient_attribution: GraphGradientAttribution::new(),
380 }
381 }
382
383 pub fn explain_prediction(
385 &mut self,
386 graph: &GraphData,
387 model_layers: &[Box<dyn GraphLayer>],
388 target_class: usize,
389 ) -> Result<GraphExplanation, Box<dyn std::error::Error>> {
390 let mut current_graph = graph.clone();
392 for (i, layer) in model_layers.iter().enumerate() {
393 current_graph = layer.forward(¤t_graph)?;
394 self.lrp
395 .store_activation(format!("layer_{}", i), current_graph.x.clone());
396 }
397
398 let final_relevance = ones(&[graph.num_nodes])?; let lrp_result = self
401 .lrp
402 .compute_graph_relevance(graph, &final_relevance, "final")?;
403
404 let gradient_saliency = self
406 .gradient_attribution
407 .gradient_saliency(graph, &final_relevance)?;
408
409 let relevance_analysis = self.lrp.analyze_relevance_patterns(graph, &lrp_result);
411
412 Ok(GraphExplanation {
413 lrp_result,
414 gradient_saliency,
415 relevance_analysis,
416 target_class,
417 })
418 }
419}
420
421impl Default for GraphExplainer {
422 fn default() -> Self {
423 Self::new()
424 }
425}
426
427#[derive(Debug, Clone)]
429pub struct GraphExplanation {
430 pub lrp_result: GraphRelevanceResult,
431 pub gradient_saliency: Tensor,
432 pub relevance_analysis: RelevanceAnalysis,
433 pub target_class: usize,
434}
435
436#[cfg(test)]
437mod tests {
438 use super::*;
439 use torsh_core::device::DeviceType;
440 use torsh_tensor::creation::{from_vec, randn};
441
442 #[test]
443 fn test_graph_lrp_creation() {
444 let lrp = GraphLRP::new();
445 assert_eq!(lrp.alpha, 1.0);
446 assert_eq!(lrp.beta, 0.0);
447 assert_eq!(lrp.epsilon, 1e-6);
448 }
449
450 #[test]
451 fn test_graph_lrp_alpha_beta() {
452 let lrp = GraphLRP::with_alpha_beta(2.0, -1.0);
453 assert_eq!(lrp.alpha, 2.0);
454 assert_eq!(lrp.beta, -1.0);
455 }
456
457 #[test]
458 fn test_graph_explainer_creation() {
459 let explainer = GraphExplainer::new();
460 assert_eq!(explainer.lrp.alpha, 1.0);
461 assert_eq!(explainer.gradient_attribution.smooth_steps, 50);
462 }
463
464 #[test]
465 fn test_relevance_epsilon_computation() {
466 let lrp = GraphLRP::new();
467
468 let input = randn(&[3, 4]).unwrap();
470 let weight = randn(&[4, 2]).unwrap();
471 let output = input.matmul(&weight).unwrap();
472 let output_relevance = ones(&[3, 2]).unwrap();
473
474 let _result = lrp.compute_relevance_epsilon(&input, &output, &weight, &output_relevance);
476 }
478
479 #[test]
480 fn test_graph_relevance_structure() {
481 let lrp = GraphLRP::new();
482
483 let x = randn(&[4, 3]).unwrap();
485 let edge_index = from_vec(
486 vec![0.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 0.0],
487 &[2, 4],
488 DeviceType::Cpu,
489 )
490 .unwrap();
491 let graph = GraphData::new(x, edge_index);
492
493 let node_relevance = ones(&[4]).unwrap();
494
495 let _result = lrp.compute_graph_relevance(&graph, &node_relevance, "test_layer");
497 }
499}