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,115 @@
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("integrated_gradients", "ig")
13
+ class IntegratedGradients(ExplanationAlgorithm):
14
+ graph_level = True
15
+
16
+ def __init__(
17
+ self,
18
+ steps: int = 50,
19
+ method: str = "riemann",
20
+ edge_grads: bool = True,
21
+ **kwargs,
22
+ ):
23
+ self.steps = steps
24
+ self.method = method
25
+ self.edge_grads = edge_grads
26
+
27
+ def explain(
28
+ self,
29
+ backend: Any,
30
+ model: Any,
31
+ data: Any,
32
+ index: int | list[int] | torch.Tensor,
33
+ target_class: int | None = None,
34
+ **kwargs,
35
+ ) -> Explanation:
36
+ model.eval()
37
+ x = backend.node_features(data)
38
+ edge_index = backend.edge_index(data)
39
+
40
+ if index is None:
41
+ idx = torch.zeros(1, dtype=torch.long, device=x.device)
42
+ elif isinstance(index, int):
43
+ idx = torch.tensor([index], device=x.device)
44
+ else:
45
+ idx = torch.as_tensor(index, device=x.device)
46
+
47
+ with torch.no_grad():
48
+ logits = backend.forward(model, x, edge_index)
49
+ if target_class is None and logits.dim() == 2:
50
+ target_class = int(logits[idx[0]].argmax().item())
51
+
52
+ baseline = torch.zeros_like(x)
53
+ alphas, weights = self._alphas(self.steps, self.method, device=x.device)
54
+
55
+ ig_feat = torch.zeros_like(x)
56
+ ig_edge = None
57
+ compute_edge = self.edge_grads and backend.supports_edge_weight(model)
58
+ if compute_edge:
59
+ ig_edge = torch.zeros(
60
+ edge_index.size(1), dtype=torch.float32, device=x.device
61
+ )
62
+
63
+ for alpha, w in zip(alphas, weights):
64
+ x_step = (baseline + alpha * (x - baseline)).requires_grad_(True)
65
+ ew = None
66
+ if compute_edge:
67
+ ew = (
68
+ torch.ones(edge_index.size(1), device=x.device) * (1.0 - alpha)
69
+ ).requires_grad_(True)
70
+ out = backend.forward(model, x_step, edge_index, edge_weight=ew)
71
+ if out.dim() == 2:
72
+ score = out[idx, target_class].sum()
73
+ else:
74
+ score = out[idx].sum()
75
+ model.zero_grad()
76
+ score.backward()
77
+ ig_feat = ig_feat + w * x_step.grad
78
+ if compute_edge and ew is not None and ew.grad is not None:
79
+ ig_edge = ig_edge + w * ew.grad
80
+
81
+ ig_feat = ig_feat * (x - baseline)
82
+ if compute_edge and ig_edge is not None:
83
+ ig_edge = ig_edge * (-1.0)
84
+
85
+ grad = ig_feat.detach()
86
+ node_importance = grad.abs().sum(dim=-1)
87
+ feature_importance = grad
88
+
89
+ return Explanation(
90
+ node_importance=node_importance.cpu(),
91
+ edge_importance=(
92
+ ig_edge.detach().abs().cpu() if ig_edge is not None else None
93
+ ),
94
+ feature_importance=feature_importance.cpu(),
95
+ prediction_original=logits[idx[0]].detach().reshape(1, -1).cpu(),
96
+ prediction_explanation=None,
97
+ node_idx=int(idx[0].item()) if isinstance(index, int) else index,
98
+ target_class=target_class,
99
+ )
100
+
101
+ def _alphas(self, steps: int, method: str, device):
102
+ if method in ("riemann", "left"):
103
+ alphas = torch.arange(0.0, 1.0, 1.0 / steps, device=device)
104
+ return alphas, torch.full_like(alphas, 1.0 / steps)
105
+ if method == "right":
106
+ alphas = torch.arange(1.0 / steps, 1.0 + 1e-6, 1.0 / steps, device=device)
107
+ return alphas, torch.full_like(alphas, 1.0 / steps)
108
+ if method == "gausslegendre":
109
+ from numpy.polynomial.legendre import leggauss
110
+
111
+ xs, ws = leggauss(steps)
112
+ alphas = torch.as_tensor((xs + 1) / 2, device=device, dtype=torch.float32)
113
+ weights = torch.as_tensor(ws / 2, device=device, dtype=torch.float32)
114
+ return alphas, weights
115
+ raise ValueError(f"method desconocido: {method}")
@@ -0,0 +1,93 @@
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("saliency", "gradient", "grad")
13
+ class Saliency(ExplanationAlgorithm):
14
+ graph_level = True
15
+
16
+ def __init__(
17
+ self,
18
+ absolute: bool = True,
19
+ aggregate: str = "sum",
20
+ node_mask_type: str | None = None,
21
+ ):
22
+ self.absolute = absolute
23
+ self.aggregate = aggregate
24
+ self.node_mask_type = node_mask_type
25
+
26
+ def explain(
27
+ self,
28
+ backend: Any,
29
+ model: Any,
30
+ data: Any,
31
+ index: int | list[int] | torch.Tensor,
32
+ target_class: int | None = None,
33
+ **kwargs,
34
+ ) -> Explanation:
35
+ model.eval()
36
+ x = backend.node_features(data).detach().clone().requires_grad_(True)
37
+ edge_index = backend.edge_index(data)
38
+ edge_weight = backend.edge_weight(data)
39
+
40
+ out = backend.forward(model, x, edge_index, edge_weight=edge_weight)
41
+ logits = out
42
+
43
+ if index is None:
44
+ idx = torch.zeros(1, dtype=torch.long, device=x.device)
45
+ elif isinstance(index, int):
46
+ idx = torch.tensor([index], device=x.device)
47
+ else:
48
+ idx = torch.as_tensor(index, device=x.device)
49
+
50
+ if target_class is None and logits.dim() == 2:
51
+ target_class = int(logits[idx].argmax(dim=-1)[0].item())
52
+
53
+ if logits.dim() == 2:
54
+ score = logits[idx, target_class].sum()
55
+ else:
56
+ score = logits[idx].sum()
57
+
58
+ model.zero_grad()
59
+ score.backward()
60
+
61
+ grad = x.grad
62
+ if grad is None:
63
+ raise RuntimeError("No se obtuvieron gradientes del modelo.")
64
+
65
+ if self.absolute:
66
+ grad = grad.abs()
67
+ if self.aggregate == "sum":
68
+ node_importance = grad.sum(dim=-1)
69
+ elif self.aggregate == "mean":
70
+ node_importance = grad.mean(dim=-1)
71
+ elif self.aggregate == "max":
72
+ node_importance = grad.max(dim=-1).values
73
+ else:
74
+ raise ValueError(f"aggregate desconocido: {self.aggregate}")
75
+
76
+ prediction_original = logits[idx].detach()
77
+ feature_importance = (
78
+ grad.detach() if self.node_mask_type == "attributes" else None
79
+ )
80
+
81
+ return Explanation(
82
+ node_importance=node_importance.detach().cpu(),
83
+ feature_importance=(
84
+ feature_importance.detach().cpu()
85
+ if feature_importance is not None
86
+ else None
87
+ ),
88
+ edge_importance=None,
89
+ prediction_original=prediction_original.cpu(),
90
+ prediction_explanation=None,
91
+ node_idx=int(idx[0].item()) if isinstance(index, int) else index,
92
+ target_class=target_class,
93
+ )
File without changes
@@ -0,0 +1,265 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ import torch.nn.functional as F
7
+
8
+ from ...core.explanation import Explanation
9
+ from ...core.registry import register
10
+ from ..base import ExplanationAlgorithm
11
+
12
+
13
+ @register("gnn_explainer", "gnnexplainer")
14
+ class GNNExplainer(ExplanationAlgorithm):
15
+ graph_level = True
16
+
17
+ def __init__(
18
+ self,
19
+ epochs: int = 200,
20
+ lr: float = 0.01,
21
+ edge_entropy: float = 0.001,
22
+ node_entropy: float = 0.001,
23
+ node_mask_type: str | None = "attributes",
24
+ edge_mask_type: str | None = "object",
25
+ prints: int = 20,
26
+ **kwargs,
27
+ ):
28
+ self.epochs = epochs
29
+ self.lr = lr
30
+ self.edge_entropy = edge_entropy
31
+ self.node_entropy = node_entropy
32
+ self.node_mask_type = node_mask_type
33
+ self.edge_mask_type = edge_mask_type
34
+ self.prints = prints
35
+
36
+ def explain(
37
+ self,
38
+ backend: Any,
39
+ model: Any,
40
+ data: Any,
41
+ index: int | torch.Tensor,
42
+ target_class: int | None = None,
43
+ **kwargs,
44
+ ) -> Explanation:
45
+ node_mask_type = kwargs.get("node_mask_type", self.node_mask_type)
46
+ edge_mask_type = kwargs.get("edge_mask_type", self.edge_mask_type)
47
+
48
+ model.eval()
49
+ x = backend.node_features(data)
50
+ edge_index = backend.edge_index(data)
51
+ num_nodes = backend.num_nodes(data)
52
+
53
+ sub_nodes, sub_edge_index, mapping, sub_edge_mask = self._extract_subgraph(
54
+ backend, data, index, edge_index
55
+ )
56
+ device = x.device
57
+ x_sub = x[sub_nodes].to(device)
58
+ sub_graph_level = index is None
59
+
60
+ node_mask = None
61
+ if node_mask_type is not None and not sub_graph_level:
62
+ node_mask = torch.nn.Parameter(torch.randn(x_sub.size(0), device=device))
63
+
64
+ edge_mask = None
65
+ if edge_mask_type is not None:
66
+ edge_mask = torch.nn.Parameter(
67
+ torch.randn(sub_edge_index.size(1), device=device)
68
+ )
69
+
70
+ params = [p for p in (node_mask, edge_mask) if p is not None]
71
+ optimizer = torch.optim.Adam(params, lr=self.lr)
72
+
73
+ with torch.no_grad():
74
+ orig_logits = backend.forward(model, x_sub, sub_edge_index)
75
+ if sub_graph_level:
76
+ if target_class is None:
77
+ target_class = (
78
+ int(orig_logits[0].argmax().item()) if orig_logits.dim() == 2 else 0
79
+ )
80
+ tgt_idx = 0
81
+ else:
82
+ ni = int(mapping.item() if torch.is_tensor(mapping) else mapping)
83
+ if target_class is None and orig_logits.dim() == 2:
84
+ target_class = int(orig_logits[ni].argmax().item())
85
+ tgt_idx = ni
86
+
87
+ for epoch in range(self.epochs):
88
+ optimizer.zero_grad()
89
+ mask_node = None
90
+ if node_mask is not None:
91
+ mask_node = torch.sigmoid(node_mask)
92
+ eweight = None
93
+ if edge_mask is not None:
94
+ eweight = torch.sigmoid(edge_mask)
95
+
96
+ pred = backend.forward(
97
+ model, x_sub, sub_edge_index, edge_weight=eweight, node_mask=mask_node
98
+ )
99
+ loss = self._loss(
100
+ pred, tgt_idx, sub_graph_level, target_class, node_mask, edge_mask
101
+ )
102
+ loss.backward()
103
+ optimizer.step()
104
+
105
+ logits = backend.forward(model, x_sub, sub_edge_index)
106
+
107
+ if sub_graph_level:
108
+ target_class = target_class or (
109
+ int(logits.argmax(-1)[0].item()) if logits.dim() == 2 else 0
110
+ )
111
+ ni = 0
112
+ pred_orig = logits[0].detach()
113
+ final_mask_node = None
114
+ final_mask_edge = None
115
+ if node_mask is not None:
116
+ final_mask_node = torch.sigmoid(node_mask)
117
+ if edge_mask is not None:
118
+ final_mask_edge = torch.sigmoid(edge_mask)
119
+ pred_masked = backend.forward(
120
+ model,
121
+ x_sub,
122
+ sub_edge_index,
123
+ edge_weight=final_mask_edge,
124
+ node_mask=final_mask_node,
125
+ )[0].detach()
126
+ else:
127
+ ni = int(mapping.item() if torch.is_tensor(mapping) else mapping)
128
+ if target_class is None and logits.dim() == 2:
129
+ target_class = int(logits[ni].argmax().item())
130
+ pred_orig = logits[ni].detach()
131
+ final_mask_node = None
132
+ final_mask_edge = None
133
+ if node_mask is not None:
134
+ final_mask_node = torch.sigmoid(node_mask)
135
+ if edge_mask is not None:
136
+ final_mask_edge = torch.sigmoid(edge_mask)
137
+ pred_masked = backend.forward(
138
+ model,
139
+ x_sub,
140
+ sub_edge_index,
141
+ edge_weight=final_mask_edge,
142
+ node_mask=final_mask_node,
143
+ )[ni].detach()
144
+
145
+ full_num_nodes = num_nodes
146
+ full_edge_count = edge_index.size(1)
147
+
148
+ node_full = self._scatter_node(node_mask, sub_nodes, full_num_nodes)
149
+ edge_full = self._scatter_edge(edge_mask, sub_edge_mask, full_edge_count)
150
+
151
+ return Explanation(
152
+ node_importance=node_full,
153
+ edge_importance=edge_full,
154
+ feature_importance=None,
155
+ prediction_original=pred_orig.cpu(),
156
+ prediction_explanation=pred_masked.cpu(),
157
+ node_idx=None
158
+ if sub_graph_level
159
+ else (int(index[0]) if torch.is_tensor(index) else int(index)),
160
+ target_class=target_class,
161
+ metadata={
162
+ "sub_nodes": sub_nodes,
163
+ "sub_edge_index": sub_edge_index,
164
+ "sub_edge_mask": sub_edge_mask,
165
+ },
166
+ )
167
+
168
+ @staticmethod
169
+ def _scatter_node(
170
+ mask: torch.nn.Parameter | None,
171
+ sub_nodes: torch.Tensor,
172
+ num_nodes: int,
173
+ ) -> torch.Tensor | None:
174
+ if mask is None:
175
+ return None
176
+ full = torch.zeros(num_nodes, dtype=torch.float32)
177
+ vals = torch.sigmoid(mask).detach().cpu()
178
+ full[sub_nodes.cpu()] = vals
179
+ return full
180
+
181
+ @staticmethod
182
+ def _scatter_edge(
183
+ mask: torch.nn.Parameter | None,
184
+ sub_edge_mask: torch.Tensor,
185
+ num_edges: int,
186
+ ) -> torch.Tensor | None:
187
+ if mask is None:
188
+ return None
189
+ full = torch.zeros(num_edges, dtype=torch.float32)
190
+ idx = sub_edge_mask.nonzero(as_tuple=False).view(-1)
191
+ full[idx.cpu()] = torch.sigmoid(mask).detach().cpu()
192
+ return full
193
+
194
+ def _loss(
195
+ self,
196
+ pred: torch.Tensor,
197
+ node_idx: int | torch.Tensor,
198
+ sub_graph_level: bool,
199
+ target_class: int | None,
200
+ node_mask: torch.nn.Parameter | None,
201
+ edge_mask: torch.nn.Parameter | None,
202
+ ) -> torch.Tensor:
203
+ if pred.dim() == 2:
204
+ log_logits = pred.log_softmax(dim=-1)
205
+ else:
206
+ log_logits = pred
207
+
208
+ if sub_graph_level:
209
+ idx = torch.zeros(1, dtype=torch.long, device=pred.device)
210
+ if target_class is None:
211
+ target_class = int(pred[0].argmax().item())
212
+ loss = F.nll_loss(
213
+ log_logits[0].unsqueeze(0),
214
+ torch.tensor([target_class], device=pred.device),
215
+ )
216
+ else:
217
+ if isinstance(node_idx, torch.Tensor) and node_idx.dim() == 0:
218
+ idx = node_idx.unsqueeze(0)
219
+ else:
220
+ idx = (
221
+ torch.as_tensor([node_idx], device=pred.device)
222
+ if not torch.is_tensor(node_idx)
223
+ else node_idx.reshape(-1)
224
+ )
225
+ if target_class is None:
226
+ target_class = int(pred[idx].argmax(dim=-1)[0].item())
227
+ loss = F.nll_loss(
228
+ log_logits[idx], torch.tensor([target_class], device=pred.device)
229
+ )
230
+
231
+ if edge_mask is not None and self.edge_entropy > 0:
232
+ loss += self.edge_entropy * self._entropy(torch.sigmoid(edge_mask))
233
+ if node_mask is not None and self.node_entropy > 0:
234
+ loss += self.node_entropy * self._entropy(torch.sigmoid(node_mask))
235
+ return loss
236
+
237
+ @staticmethod
238
+ def _entropy(p: torch.Tensor) -> torch.Tensor:
239
+ eps = 1e-8
240
+ return -(p * torch.log(p + eps) + (1 - p) * torch.log(1 - p + eps)).mean()
241
+
242
+ def _extract_subgraph(
243
+ self,
244
+ backend: Any,
245
+ data: Any,
246
+ index: int | torch.Tensor | None,
247
+ edge_index: torch.Tensor,
248
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
249
+ from torch_geometric.utils import k_hop_subgraph
250
+
251
+ if index is None:
252
+ n = edge_index.max().item() + 1
253
+ node_idx = torch.arange(n, device=edge_index.device)
254
+ return (
255
+ node_idx,
256
+ edge_index,
257
+ node_idx,
258
+ torch.ones(
259
+ edge_index.size(1), dtype=torch.bool, device=edge_index.device
260
+ ),
261
+ )
262
+ node_idx = torch.as_tensor([index], device=edge_index.device).reshape(-1)
263
+ return k_hop_subgraph(
264
+ node_idx, num_hops=3, edge_index=edge_index, relabel_nodes=True
265
+ )
@@ -0,0 +1,136 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ import torch.nn.functional as F
7
+
8
+ from ...core.explanation import Explanation
9
+ from ...core.registry import register
10
+ from ..base import ExplanationAlgorithm
11
+
12
+
13
+ @register("node_mask", "nodemask", "nm")
14
+ class NodeMask(ExplanationAlgorithm):
15
+ """NodeMask: node mask learned by optimization.
16
+
17
+ Optimizes a (sigmoid) mask over the nodes of the target node's k-hop
18
+ subgraph so the model keeps its prediction, with an entropy regularizer to
19
+ force sparsity. The resulting node importance is re-projected onto the full
20
+ graph (0 outside the neighborhood).
21
+ """
22
+
23
+ def __init__(
24
+ self,
25
+ epochs: int = 200,
26
+ lr: float = 0.05,
27
+ hops: int = 3,
28
+ suppress_ratio: float = 0.8,
29
+ entropy: float = 0.05,
30
+ **kwargs,
31
+ ):
32
+ self.epochs = epochs
33
+ self.lr = lr
34
+ self.hops = hops
35
+ self.suppress_ratio = suppress_ratio
36
+ self.entropy = entropy
37
+
38
+ def explain(
39
+ self,
40
+ backend: Any,
41
+ model: Any,
42
+ data: Any,
43
+ index: int | torch.Tensor | None = None,
44
+ target_class: int | None = None,
45
+ **kwargs,
46
+ ) -> Explanation:
47
+ model.eval()
48
+ x = backend.node_features(data)
49
+ edge_index = backend.edge_index(data)
50
+ num_nodes = backend.num_nodes(data)
51
+
52
+ nodes = self._to_node_ids(index, num_nodes)
53
+ root = int(nodes[0])
54
+
55
+ from torch_geometric.utils import k_hop_subgraph
56
+
57
+ device = x.device
58
+ sub_nodes, sub_edge_index, mapping, _ = k_hop_subgraph(
59
+ [root],
60
+ num_hops=self.hops,
61
+ edge_index=edge_index,
62
+ relabel_nodes=True,
63
+ )
64
+ x_sub = x[sub_nodes].to(device)
65
+
66
+ with torch.no_grad():
67
+ orig = backend.forward(model, x_sub, sub_edge_index)
68
+ ni = int(mapping.item() if torch.is_tensor(mapping) else mapping)
69
+ if target_class is None and orig.dim() == 2:
70
+ target_class = int(orig[ni].argmax().item())
71
+ target = target_class if target_class is not None else 0
72
+
73
+ if self.suppress_ratio > 0:
74
+ k = int(self.suppress_ratio * sub_nodes.numel())
75
+ k = max(0, k)
76
+ else:
77
+ k = max(0, sub_nodes.numel() - 1)
78
+
79
+ mask = torch.nn.Parameter(torch.zeros(sub_nodes.numel(), device=device))
80
+ optimizer = torch.optim.Adam([mask], lr=self.lr)
81
+
82
+ for _ in range(self.epochs):
83
+ optimizer.zero_grad()
84
+ node_mask = torch.sigmoid(mask)
85
+ pred = backend.forward(model, x_sub, sub_edge_index, node_mask=node_mask)
86
+ topk = node_mask.topk(max(k, 1)).values.min()
87
+ loss = self._loss(pred, ni, target, node_mask, topk)
88
+ loss.backward()
89
+ optimizer.step()
90
+
91
+ final = torch.sigmoid(mask).detach()
92
+ full = torch.zeros(num_nodes, dtype=torch.float32)
93
+ full[sub_nodes.cpu()] = final.cpu()
94
+
95
+ with torch.no_grad():
96
+ pred_masked = backend.forward(
97
+ model, x_sub, sub_edge_index, node_mask=final
98
+ )[ni]
99
+
100
+ return Explanation(
101
+ node_importance=full,
102
+ edge_importance=None,
103
+ feature_importance=None,
104
+ prediction_original=orig[ni].detach().cpu(),
105
+ prediction_explanation=pred_masked.detach().cpu(),
106
+ node_idx=root,
107
+ target_class=target,
108
+ )
109
+
110
+ def _loss(self, pred, node_idx, target_class, node_mask, topk) -> torch.Tensor:
111
+ if pred.dim() == 2:
112
+ log_logits = pred.log_softmax(dim=-1)
113
+ else:
114
+ log_logits = pred
115
+ idx = torch.as_tensor([node_idx], device=pred.device).reshape(-1)
116
+ loss = F.nll_loss(
117
+ log_logits[idx],
118
+ torch.tensor([target_class], device=pred.device),
119
+ )
120
+ loss += self.entropy * self._entropy(node_mask)
121
+ loss += (node_mask - topk.detach()).relu().mean()
122
+ return loss
123
+
124
+ @staticmethod
125
+ def _entropy(p: torch.Tensor) -> torch.Tensor:
126
+ eps = 1e-8
127
+ return -(p * torch.log(p + eps) + (1 - p) * torch.log(1 - p + eps)).mean()
128
+
129
+ @staticmethod
130
+ def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
131
+ if index is None:
132
+ return torch.zeros(1, dtype=torch.long)
133
+ if isinstance(index, int):
134
+ return torch.tensor([index], dtype=torch.long)
135
+ idx = torch.as_tensor(index, dtype=torch.long)
136
+ return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)