graph-explain 0.7.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (42) hide show
  1. graph_explain/__init__.py +79 -0
  2. graph_explain/backends/__init__.py +4 -0
  3. graph_explain/backends/base.py +103 -0
  4. graph_explain/backends/dgl.py +121 -0
  5. graph_explain/benchmarks/__init__.py +3 -0
  6. graph_explain/benchmarks/synthetic.py +246 -0
  7. graph_explain/cli.py +459 -0
  8. graph_explain/core/__init__.py +14 -0
  9. graph_explain/core/benchmark.py +284 -0
  10. graph_explain/core/evaluation.py +391 -0
  11. graph_explain/core/explainer.py +83 -0
  12. graph_explain/core/explanation.py +72 -0
  13. graph_explain/core/model_utils.py +44 -0
  14. graph_explain/core/registry.py +55 -0
  15. graph_explain/methods/__init__.py +39 -0
  16. graph_explain/methods/attention/attention.py +147 -0
  17. graph_explain/methods/base.py +25 -0
  18. graph_explain/methods/baseline/random_baseline.py +78 -0
  19. graph_explain/methods/counterfactual/counterfactual.py +304 -0
  20. graph_explain/methods/feature/graph_lime.py +141 -0
  21. graph_explain/methods/gradient/__init__.py +0 -0
  22. graph_explain/methods/gradient/grad_x_input.py +110 -0
  23. graph_explain/methods/gradient/guided_backprop.py +117 -0
  24. graph_explain/methods/gradient/integrated_gradients.py +115 -0
  25. graph_explain/methods/gradient/saliency.py +93 -0
  26. graph_explain/methods/perturbation/__init__.py +0 -0
  27. graph_explain/methods/perturbation/gnn_explainer.py +265 -0
  28. graph_explain/methods/perturbation/node_mask.py +136 -0
  29. graph_explain/methods/perturbation/pg_explainer.py +162 -0
  30. graph_explain/methods/perturbation/subgraphx.py +393 -0
  31. graph_explain/methods/relevance/deeplift.py +262 -0
  32. graph_explain/methods/relevance/gnn_lrp.py +219 -0
  33. graph_explain/narration/__init__.py +3 -0
  34. graph_explain/narration/narrator.py +185 -0
  35. graph_explain/visualization/__init__.py +4 -0
  36. graph_explain/visualization/interactive.py +73 -0
  37. graph_explain/visualization/static.py +90 -0
  38. graph_explain-0.7.0.dist-info/METADATA +332 -0
  39. graph_explain-0.7.0.dist-info/RECORD +42 -0
  40. graph_explain-0.7.0.dist-info/WHEEL +5 -0
  41. graph_explain-0.7.0.dist-info/entry_points.txt +2 -0
  42. graph_explain-0.7.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,304 @@
1
+ from __future__ import annotations
2
+
3
+ from collections import deque
4
+ from typing import Any
5
+
6
+ import torch
7
+ import torch.nn.functional as F
8
+
9
+ from ...core.explanation import Explanation
10
+ from ...core.registry import register
11
+ from ..base import ExplanationAlgorithm
12
+
13
+
14
+ def _softmax(logits, c: int) -> float:
15
+ return float(F.softmax(logits[None], dim=-1)[0, c].item())
16
+
17
+
18
+ @register("counterfactual", "counterfactual_explainer", "cf")
19
+ class Counterfactual(ExplanationAlgorithm):
20
+ """Counterfactual explanation: minimal perturbation that changes the prediction.
21
+
22
+ Finds the minimal set of edges (mode='edge') or feature coordinates
23
+ (mode='feature') whose removal/re-scaling makes the node's prediction change
24
+ class (or reach `flip_to`). The search is greedy and deterministic: at each
25
+ step it removes the candidate element that most reduces `P(original class)`;
26
+ if the class does not change within `max_steps` steps it returns the current
27
+ state (prediction unchanged).
28
+
29
+ The returned importance marks the modified elements (edges 0/1, nodes from
30
+ their incidence on removed edges, changed features 0/1), with
31
+ `prediction_explanation` = logits after the perturbation.
32
+ """
33
+
34
+ def __init__(
35
+ self,
36
+ mode: str = "edge",
37
+ flip_to: int | None = None,
38
+ max_steps: int = 10,
39
+ hops: int = 2,
40
+ eps: float = 0.0,
41
+ seed: int = 0,
42
+ ):
43
+ if mode not in ("edge", "feature"):
44
+ raise ValueError("mode debe ser 'edge' o 'feature'.")
45
+ self.mode = mode
46
+ self.flip_to = flip_to
47
+ self.max_steps = int(max_steps)
48
+ self.hops = int(hops)
49
+ self.eps = float(eps)
50
+ torch.manual_seed(seed)
51
+
52
+ def explain(
53
+ self,
54
+ backend: Any,
55
+ model: Any,
56
+ data: Any,
57
+ index: int | list[int] | torch.Tensor | None = None,
58
+ target_class: int | None = None,
59
+ **kwargs,
60
+ ) -> Explanation:
61
+ model.eval()
62
+ node = self._single_node(index)
63
+ x = backend.node_features(data).detach().clone()
64
+ edge_index = backend.edge_index(data)
65
+ edge_weight = backend.edge_weight(data)
66
+ num_nodes = int(x.size(0))
67
+
68
+ orig_logits = backend.forward(model, x, edge_index, edge_weight=edge_weight)
69
+ orig_class = int(orig_logits[node].argmax().item())
70
+ flip_class = target_class if self.flip_to is None else self.flip_to
71
+ target_new = None if flip_class is None else int(flip_class)
72
+
73
+ device = x.device
74
+ if self.mode == "edge":
75
+ edge_importance, final_logits = self._flip_edges(
76
+ backend,
77
+ model,
78
+ x,
79
+ edge_index,
80
+ edge_weight,
81
+ node,
82
+ orig_class,
83
+ target_new,
84
+ device,
85
+ )
86
+ return self._build(
87
+ backend,
88
+ data,
89
+ node,
90
+ orig_class,
91
+ orig_logits,
92
+ final_logits,
93
+ node_importance=torch.bincount(
94
+ torch.cat(
95
+ [
96
+ edge_index[0][edge_importance.bool()],
97
+ edge_index[1][edge_importance.bool()],
98
+ ]
99
+ ),
100
+ minlength=num_nodes,
101
+ ).float(),
102
+ edge_importance=edge_importance,
103
+ feature_importance=None,
104
+ )
105
+ feature_importance, _final_x, final_logits = self._flip_features(
106
+ backend,
107
+ model,
108
+ x,
109
+ edge_index,
110
+ edge_weight,
111
+ node,
112
+ orig_class,
113
+ target_new,
114
+ device,
115
+ )
116
+ return self._build(
117
+ backend,
118
+ data,
119
+ node,
120
+ orig_class,
121
+ orig_logits,
122
+ final_logits,
123
+ node_importance=torch.zeros(num_nodes).index_fill(
124
+ 0, torch.tensor([node], device=device), 1.0
125
+ ),
126
+ edge_importance=None,
127
+ feature_importance=feature_importance,
128
+ )
129
+
130
+ # ------------------------------------------------------------------ búsquedas
131
+ def _flip_edges(
132
+ self,
133
+ backend,
134
+ model,
135
+ x,
136
+ edge_index,
137
+ edge_weight,
138
+ node,
139
+ orig_class,
140
+ target_new,
141
+ device,
142
+ ):
143
+ num_nodes = int(x.size(0))
144
+ weight = (
145
+ edge_weight.detach().clone()
146
+ if edge_weight is not None
147
+ else torch.ones(edge_index.size(1), device=device)
148
+ )
149
+ candidates = self._candidate_edges(edge_index, num_nodes, node, self.hops)
150
+ removed: list[int] = []
151
+ pred = backend.forward(model, x, edge_index, edge_weight=weight)[node]
152
+
153
+ def flipped(logits, tc):
154
+ pred_cls = int(logits.argmax().item())
155
+ return pred_cls == tc if tc is not None else pred_cls != orig_class
156
+
157
+ for _ in range(self.max_steps):
158
+ if flipped(pred, target_new):
159
+ break
160
+ remaining = [e for e in candidates if e not in removed]
161
+ if not remaining:
162
+ break
163
+ best_e, best_logits, best_p = None, None, None
164
+ with torch.no_grad():
165
+ for e in remaining:
166
+ w = weight.clone()
167
+ w[e] = 0.0
168
+ lg = backend.forward(model, x, edge_index, edge_weight=w)[node]
169
+ p = float(_softmax(lg, orig_class))
170
+ if best_p is None or p < best_p:
171
+ best_e, best_logits, best_p = e, lg, p
172
+ if best_e is None or (
173
+ best_p is not None
174
+ and best_p >= float(_softmax(pred, orig_class)) - self.eps
175
+ ):
176
+ break
177
+ removed.append(best_e)
178
+ weight[best_e] = 0.0
179
+ pred = best_logits
180
+
181
+ importance = torch.zeros(edge_index.size(1), device=device)
182
+ importance[removed] = 1.0
183
+ return importance, pred.detach()
184
+
185
+ def _flip_features(
186
+ self,
187
+ backend,
188
+ model,
189
+ x,
190
+ edge_index,
191
+ edge_weight,
192
+ node,
193
+ orig_class,
194
+ target_new,
195
+ device,
196
+ ):
197
+ baseline = x.mean(dim=0, keepdim=True)
198
+ x_cur = x.detach().clone()
199
+ changed: list[int] = []
200
+ pred = backend.forward(model, x_cur, edge_index, edge_weight=edge_weight)[node]
201
+
202
+ def flipped(logits, tc):
203
+ pred_cls = int(logits.argmax().item())
204
+ return pred_cls == tc if tc is not None else pred_cls != orig_class
205
+
206
+ for _ in range(self.max_steps):
207
+ if flipped(pred, target_new):
208
+ break
209
+ best_c, best_logits, best_p = None, None, None
210
+ with torch.no_grad():
211
+ for c in range(x.size(1)):
212
+ if c in changed:
213
+ continue
214
+ xn = x_cur.clone()
215
+ xn[node, c] = baseline[0, c]
216
+ lg = backend.forward(
217
+ model, xn, edge_index, edge_weight=edge_weight
218
+ )[node]
219
+ p = float(_softmax(lg, orig_class))
220
+ if best_p is None or p < best_p:
221
+ best_c, best_logits, best_p = c, lg, p
222
+ if best_c is None or (
223
+ best_p is not None
224
+ and best_p >= float(_softmax(pred, orig_class)) - self.eps
225
+ ):
226
+ break
227
+ changed.append(best_c)
228
+ x_cur[node, best_c] = baseline[0, best_c]
229
+ pred = best_logits
230
+
231
+ importance = torch.zeros(x.size(1), device=device)
232
+ importance[changed] = 1.0
233
+ return importance, x_cur, pred.detach()
234
+
235
+ # ------------------------------------------------------------------ utilidades
236
+ @staticmethod
237
+ def _single_node(index) -> int:
238
+ if index is None:
239
+ return 0
240
+ if isinstance(index, (list, tuple)):
241
+ index = index[0]
242
+ return int(torch.as_tensor(index).reshape(-1)[0].item())
243
+
244
+ @staticmethod
245
+ def _candidate_edges(edge_index, num_nodes: int, node: int, hops: int) -> list[int]:
246
+ adj: dict[int, list[int]] = {}
247
+ for u, v in zip(edge_index[0].tolist(), edge_index[1].tolist()):
248
+ adj.setdefault(u, []).append(v)
249
+ distance = {node: 0}
250
+ queue = deque([node])
251
+ while queue:
252
+ cur = queue.popleft()
253
+ if distance[cur] >= hops:
254
+ continue
255
+ for nb in adj.get(cur, []):
256
+ if nb not in distance:
257
+ distance[nb] = distance[cur] + 1
258
+ queue.append(nb)
259
+ near = set(distance)
260
+ edges = []
261
+ for e in range(edge_index.size(1)):
262
+ u = int(edge_index[0, e].item())
263
+ v = int(edge_index[1, e].item())
264
+ if u in near or v in near:
265
+ edges.append(e)
266
+ return edges
267
+
268
+ def _build(
269
+ self,
270
+ backend,
271
+ data,
272
+ node,
273
+ orig_class,
274
+ orig_logits,
275
+ final_logits,
276
+ node_importance,
277
+ edge_importance,
278
+ feature_importance,
279
+ ) -> Explanation:
280
+ metadata = {
281
+ "backend": backend,
282
+ "backing_data": data,
283
+ "counterfactual": True,
284
+ "original_class": orig_class,
285
+ }
286
+ return Explanation(
287
+ node_importance=node_importance.detach().cpu(),
288
+ edge_importance=edge_importance.detach().cpu()
289
+ if edge_importance is not None
290
+ else None,
291
+ feature_importance=(
292
+ feature_importance.detach().cpu()
293
+ if feature_importance is not None
294
+ else None
295
+ ),
296
+ prediction_original=orig_logits[node].detach().cpu(),
297
+ prediction_explanation=final_logits.detach().cpu(),
298
+ node_idx=node,
299
+ target_class=int(final_logits.argmax().item())
300
+ if final_logits.numel()
301
+ else None,
302
+ metadata=metadata,
303
+ mask_threshold=0.5,
304
+ )
@@ -0,0 +1,141 @@
1
+ from __future__ import annotations
2
+
3
+ from collections import defaultdict
4
+ from typing import Any
5
+
6
+ import torch
7
+
8
+ from ...core.explanation import Explanation
9
+ from ...core.registry import register
10
+ from ..base import ExplanationAlgorithm
11
+
12
+
13
+ @register("graph_lime", "glime", "gl")
14
+ class GraphLIME(ExplanationAlgorithm):
15
+ """GraphLIME: feature attribution via weighted local regression.
16
+
17
+ Fits a linear regression (ridge, closed form) over the k-hop neighbors'
18
+ features, weighting each neighbor by its similarity to the target node's
19
+ feature (Gaussian kernel). The coefficients explain the probability
20
+ (softmax) of the target class; node importance matches the kernel
21
+ similarity.
22
+ """
23
+
24
+ def __init__(
25
+ self,
26
+ hops: int = 2,
27
+ lambda_: float = 1.0,
28
+ sigma: float | None = None,
29
+ normalize: bool = True,
30
+ **kwargs,
31
+ ):
32
+ self.hops = hops
33
+ self.lambda_ = lambda_
34
+ self.sigma = sigma
35
+ self.normalize = normalize
36
+
37
+ def explain(
38
+ self,
39
+ backend: Any,
40
+ model: Any,
41
+ data: Any,
42
+ index: int | torch.Tensor | None = None,
43
+ target_class: int | None = None,
44
+ **kwargs,
45
+ ) -> Explanation:
46
+ model.eval()
47
+ x = backend.node_features(data)
48
+ edge_index = backend.edge_index(data)
49
+ num_nodes = int(x.size(0))
50
+
51
+ nodes = self._to_node_ids(index, num_nodes)
52
+ root = int(nodes[0])
53
+ with torch.no_grad():
54
+ logits = backend.forward(model, x, edge_index)
55
+ if logits.dim() != 2:
56
+ raise ValueError(
57
+ "GraphLIME requiere predicciones node-level (logits (N, C))."
58
+ )
59
+ target = target_class
60
+ if target is None:
61
+ target = int(logits[root].argmax().item())
62
+ target = max(0, min(int(target), logits.size(1) - 1))
63
+
64
+ probs = torch.softmax(logits, dim=-1)[:, target].detach().cpu()
65
+
66
+ neighbors = self._khop_neighbors(edge_index, root, self.hops)
67
+ if root not in neighbors:
68
+ neighbors.append(root)
69
+ nb = torch.tensor(neighbors, dtype=torch.long)
70
+ x_nb = x[nb].detach().cpu()
71
+ y = probs[nb]
72
+
73
+ dists = torch.norm(x_nb - x[root : root + 1].detach().cpu(), dim=-1)
74
+ sigma = self.sigma
75
+ if sigma is None:
76
+ tail = dists[1:]
77
+ sigma = float(tail.mean().item()) if tail.numel() else 1.0
78
+ sigma = max(sigma, 1e-4)
79
+ weights = torch.exp(-(dists**2) / (2.0 * sigma**2))
80
+
81
+ coef = self._ridge(x_nb, y, weights, self.lambda_)
82
+
83
+ node_importance = torch.zeros(num_nodes)
84
+ node_importance[nb] = weights
85
+ node_importance = node_importance.cpu()
86
+
87
+ if self.normalize and coef.numel():
88
+ denom = coef.abs().max()
89
+ if denom > 1e-12:
90
+ coef = coef / denom
91
+
92
+ return Explanation(
93
+ node_importance=node_importance,
94
+ edge_importance=None,
95
+ feature_importance=coef,
96
+ prediction_original=logits[root].detach().cpu(),
97
+ prediction_explanation=None,
98
+ node_idx=root,
99
+ target_class=target,
100
+ metadata={"neighborhood": neighbors},
101
+ )
102
+
103
+ @staticmethod
104
+ def _ridge(
105
+ x: torch.Tensor, y: torch.Tensor, weights: torch.Tensor, lambda_: float
106
+ ) -> torch.Tensor:
107
+ ones = torch.ones(x.size(0), 1)
108
+ X = torch.cat([x, ones], dim=-1).double()
109
+ yw = (y * weights).double()
110
+ Xw = X * weights.unsqueeze(-1).double()
111
+ n_features = X.size(1)
112
+ gram = Xw.t() @ X + lambda_ * torch.eye(n_features, dtype=torch.double)
113
+ beta = torch.linalg.solve(gram, X.t() @ yw)
114
+ return beta[:-1].float()
115
+
116
+ @staticmethod
117
+ def _khop_neighbors(edge_index: torch.Tensor, node: int, hops: int) -> list[int]:
118
+ adj: dict[int, set[int]] = defaultdict(set)
119
+ src = edge_index[0].tolist()
120
+ dst = edge_index[1].tolist()
121
+ for s, d in zip(src, dst):
122
+ adj[s].add(d)
123
+ adj[d].add(s)
124
+ seen = {node}
125
+ frontier = {node}
126
+ for _ in range(hops):
127
+ nxt: set[int] = set()
128
+ for n in frontier:
129
+ nxt |= adj.get(n, set())
130
+ frontier = nxt - seen
131
+ seen |= frontier
132
+ return sorted(seen)
133
+
134
+ @staticmethod
135
+ def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
136
+ if index is None:
137
+ return torch.zeros(1, dtype=torch.long)
138
+ if isinstance(index, int):
139
+ return torch.tensor([index], dtype=torch.long)
140
+ idx = torch.as_tensor(index, dtype=torch.long)
141
+ return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
File without changes
@@ -0,0 +1,110 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+
7
+ from ...core.explanation import Explanation
8
+ from ...core.registry import register
9
+ from ..base import ExplanationAlgorithm
10
+
11
+
12
+ @register("grad_x_input", "gradient_x_input", "gx")
13
+ class GradXInput(ExplanationAlgorithm):
14
+ """Gradient x Input: attribution as gradient scaled by the activation.
15
+
16
+ The importance of each feature (and of each edge, if the backend supports
17
+ edge weights) is the gradient of the target-class logit multiplied by the
18
+ input-baseline difference (zero baseline by default). Node importance is the
19
+ sum of `abs(grad * Δx)` over features.
20
+ """
21
+
22
+ graph_level = True
23
+
24
+ def __init__(
25
+ self,
26
+ baseline: str = "zero",
27
+ edge_grads: bool = True,
28
+ node_mask_type: str | None = "attributes",
29
+ **kwargs,
30
+ ):
31
+ self.baseline_name = baseline
32
+ self.edge_grads = edge_grads
33
+ self.node_mask_type = node_mask_type
34
+
35
+ def explain(
36
+ self,
37
+ backend: Any,
38
+ model: Any,
39
+ data: Any,
40
+ index: int | list[int] | torch.Tensor | None = None,
41
+ target_class: int | None = None,
42
+ **kwargs,
43
+ ) -> Explanation:
44
+ model.eval()
45
+ x = backend.node_features(data)
46
+ edge_index = backend.edge_index(data)
47
+ num_nodes = int(x.size(0))
48
+
49
+ nodes = self._to_node_ids(index, num_nodes)
50
+ with torch.no_grad():
51
+ logits = backend.forward(model, x, edge_index)
52
+ if logits.dim() != 2:
53
+ raise ValueError(
54
+ "GradXInput requiere predicciones node-level (logits (N, C))."
55
+ )
56
+ target = target_class
57
+ if target is None:
58
+ target = int(logits[nodes[0]].argmax().item())
59
+ target_cls = max(0, min(int(target), logits.size(1) - 1))
60
+
61
+ x_in = x.detach().clone().requires_grad_(True)
62
+ ew = None
63
+ compute_edge = self.edge_grads and backend.supports_edge_weight(model)
64
+ if compute_edge:
65
+ ew = torch.ones(
66
+ edge_index.size(1), dtype=torch.float32, device=x.device
67
+ ).requires_grad_(True)
68
+
69
+ out = backend.forward(model, x_in, edge_index, edge_weight=ew)
70
+ if out.dim() != 2:
71
+ raise ValueError("GradXInput requiere predicciones node-level.")
72
+ score = out[nodes, target_cls].sum()
73
+ model.zero_grad()
74
+ score.backward()
75
+
76
+ grad_x = x_in.grad.detach()
77
+ baseline = torch.zeros_like(x_in)
78
+ contrib = grad_x * (x_in - baseline)
79
+
80
+ node_importance = contrib.abs().sum(dim=-1)
81
+ feature_importance = contrib.detach()
82
+
83
+ edge_importance = None
84
+ if compute_edge and ew is not None and ew.grad is not None:
85
+ edge_importance = ew.grad.detach().abs()
86
+
87
+ return Explanation(
88
+ node_importance=node_importance.cpu(),
89
+ edge_importance=(
90
+ edge_importance.cpu() if edge_importance is not None else None
91
+ ),
92
+ feature_importance=(
93
+ feature_importance.cpu()
94
+ if self.node_mask_type == "attributes"
95
+ else None
96
+ ),
97
+ prediction_original=logits[nodes[0]].detach().cpu(),
98
+ prediction_explanation=None,
99
+ node_idx=int(nodes[0]) if nodes.shape[0] == 1 else index,
100
+ target_class=target_cls,
101
+ )
102
+
103
+ @staticmethod
104
+ def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
105
+ if index is None:
106
+ return torch.zeros(1, dtype=torch.long)
107
+ if isinstance(index, int):
108
+ return torch.tensor([index], dtype=torch.long)
109
+ idx = torch.as_tensor(index, dtype=torch.long)
110
+ return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
@@ -0,0 +1,117 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ from torch import nn
7
+
8
+ from ...core.explanation import Explanation
9
+ from ...core.registry import register
10
+ from ..base import ExplanationAlgorithm
11
+
12
+
13
+ @register("guided_backprop", "guided-backprop", "gbp")
14
+ class GuidedBackprop(ExplanationAlgorithm):
15
+ """Guided Backpropagation: gradients guided by the ReLU mask.
16
+
17
+ During backpropagation the gradient is filtered: it only propagates where
18
+ the ReLU activation was positive (negative gradients are discarded),
19
+ highlighting the features that positively contribute to the class. Temporary
20
+ hooks are registered on the `nn.ReLU` modules; if the model has none, it
21
+ falls back to standard gradients (metadata `guided=False`).
22
+ """
23
+
24
+ graph_level = True
25
+
26
+ def __init__(self, fallback_to_gradient: bool = True, **kwargs):
27
+ self.fallback_to_gradient = fallback_to_gradient
28
+
29
+ def explain(
30
+ self,
31
+ backend: Any,
32
+ model: Any,
33
+ data: Any,
34
+ index: int | torch.Tensor | None = None,
35
+ target_class: int | None = None,
36
+ **kwargs,
37
+ ) -> Explanation:
38
+ model.eval()
39
+ x = backend.node_features(data)
40
+ edge_index = backend.edge_index(data)
41
+ num_nodes = int(x.size(0))
42
+
43
+ nodes = self._to_node_ids(index, num_nodes)
44
+ root = int(nodes[0])
45
+ with torch.no_grad():
46
+ logits = backend.forward(model, x, edge_index)
47
+ if logits.dim() != 2:
48
+ raise ValueError(
49
+ "GuidedBackprop requiere predicciones node-level (logits (N, C))."
50
+ )
51
+ target = target_class
52
+ if target is None:
53
+ target = int(logits[root].argmax().item())
54
+ target = max(0, min(int(target), logits.size(1) - 1))
55
+
56
+ relus = [m for m in model.modules() if isinstance(m, nn.ReLU)]
57
+ hooks, guided = [], False
58
+ if relus and self.fallback_to_gradient:
59
+ guided = True
60
+ for module in relus:
61
+ fw = module.register_forward_hook(self._mask_forward)
62
+ bw = module.register_full_backward_hook(self._guide_backward)
63
+ hooks.extend((fw, bw))
64
+
65
+ try:
66
+ x_in = x.detach().clone().requires_grad_(True)
67
+ out = backend.forward(model, x_in, edge_index)
68
+ if out.dim() != 2:
69
+ raise ValueError("GuidedBackprop requiere predicciones node-level.")
70
+ score = out[root, target]
71
+ model.zero_grad()
72
+ score.backward()
73
+ grad = (
74
+ x_in.grad.detach() if x_in.grad is not None else torch.zeros_like(x_in)
75
+ )
76
+ finally:
77
+ for hook in hooks:
78
+ hook.remove()
79
+
80
+ node_importance = grad.abs().sum(dim=-1)
81
+ return Explanation(
82
+ node_importance=node_importance.cpu(),
83
+ edge_importance=None,
84
+ feature_importance=grad.cpu(),
85
+ prediction_original=logits[root].detach().cpu(),
86
+ prediction_explanation=None,
87
+ node_idx=root,
88
+ target_class=target,
89
+ metadata={"guided": guided},
90
+ )
91
+
92
+ @staticmethod
93
+ def _mask_forward(module, inp, out):
94
+ mask = torch.where(out > 0, torch.ones_like(out), torch.zeros_like(out))
95
+ module._gbp_mask = mask.detach()
96
+
97
+ @staticmethod
98
+ def _guide_backward(module, grad_input, grad_output):
99
+ mask = module._gbp_mask
100
+ if mask is None or not grad_output or grad_output[0] is None:
101
+ return grad_input
102
+ g = grad_output[0]
103
+ if g.shape != mask.shape:
104
+ return grad_input
105
+ guided = (g * mask).clamp(min=0)
106
+ if len(grad_input) == 1:
107
+ return (guided,)
108
+ return grad_input
109
+
110
+ @staticmethod
111
+ def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
112
+ if index is None:
113
+ return torch.zeros(1, dtype=torch.long)
114
+ if isinstance(index, int):
115
+ return torch.tensor([index], dtype=torch.long)
116
+ idx = torch.as_tensor(index, dtype=torch.long)
117
+ return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)