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,162 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ import torch.nn.functional as F
7
+ from torch import nn
8
+
9
+ from ...core.explanation import Explanation
10
+ from ...core.model_utils import capture_node_embeddings, edge_embeddings
11
+ from ...core.registry import register
12
+ from ..base import ExplanationAlgorithm
13
+
14
+
15
+ class _EdgeMaskMLP(nn.Module):
16
+ def __init__(self, in_dim: int, hidden: int = 64):
17
+ super().__init__()
18
+ self.mlp = nn.Sequential(
19
+ nn.Linear(in_dim, hidden),
20
+ nn.ReLU(),
21
+ nn.Linear(hidden, hidden),
22
+ nn.ReLU(),
23
+ nn.Linear(hidden, 1),
24
+ )
25
+
26
+ def forward(self, edge_emb: torch.Tensor) -> torch.Tensor:
27
+ return self.mlp(edge_emb).squeeze(-1)
28
+
29
+
30
+ @register("pg_explainer", "pgexplainer")
31
+ class PGExplainer(ExplanationAlgorithm):
32
+ def __init__(
33
+ self,
34
+ epochs: int = 100,
35
+ lr: float = 0.01,
36
+ hidden: int = 64,
37
+ temp: float = 1.0,
38
+ loss_coeff: float = 0.5,
39
+ entropy_coeff: float = 0.005,
40
+ batch_nodes: int = 32,
41
+ **kwargs,
42
+ ):
43
+ self.epochs = epochs
44
+ self.lr = lr
45
+ self.hidden = hidden
46
+ self.temp = temp
47
+ self.loss_coeff = loss_coeff
48
+ self.entropy_coeff = entropy_coeff
49
+ self.batch_nodes = batch_nodes
50
+
51
+ def explain(
52
+ self,
53
+ backend: Any,
54
+ model: Any,
55
+ data: Any,
56
+ index: int | torch.Tensor | None,
57
+ target_class: int | None = None,
58
+ **kwargs,
59
+ ) -> Explanation:
60
+ model.eval()
61
+ x = backend.node_features(data)
62
+ edge_index = backend.edge_index(data)
63
+ num_nodes = backend.num_nodes(data)
64
+
65
+ embeddings = capture_node_embeddings(model, backend, data)
66
+ edge_emb = edge_embeddings(embeddings, edge_index).detach()
67
+ in_dim = edge_emb.size(-1)
68
+
69
+ with torch.no_grad():
70
+ base_logits = backend.forward(model, x, edge_index)
71
+
72
+ mlp = _EdgeMaskMLP(in_dim, hidden=self.hidden)
73
+ optimizer = torch.optim.Adam(mlp.parameters(), lr=self.lr)
74
+
75
+ train_idx = (
76
+ torch.arange(num_nodes, device=x.device)
77
+ if not hasattr(data, "train_mask")
78
+ else data.train_mask.nonzero(as_tuple=False).view(-1)
79
+ )
80
+
81
+ for epoch in range(self.epochs):
82
+ mlp.train()
83
+ optimizer.zero_grad()
84
+ idx = train_idx[
85
+ torch.randperm(len(train_idx), device=x.device)[: self.batch_nodes]
86
+ ]
87
+ edge_weight = self._sample_mask(mlp, edge_emb, x.device)
88
+ pred = backend.forward(model, x, edge_index, edge_weight=edge_weight)
89
+ if pred.dim() == 2:
90
+ target = base_logits[idx].argmax(dim=-1)
91
+ ce = F.nll_loss(pred[idx].log_softmax(dim=-1), target)
92
+ else:
93
+ ce = -pred[idx].mean()
94
+ logit = mlp(edge_emb)
95
+ p = torch.sigmoid(logit)
96
+ entropy = (
97
+ -(p * torch.clamp(p, 1e-8, 1).log())
98
+ - ((1 - p) * torch.clamp(1 - p, 1e-8, 1).log())
99
+ ).mean()
100
+ sparsity = edge_weight.mean()
101
+ loss = ce + self.loss_coeff * sparsity + self.entropy_coeff * entropy
102
+ loss.backward()
103
+ optimizer.step()
104
+
105
+ mlp.eval()
106
+ with torch.no_grad():
107
+ edge_mask = torch.sigmoid(mlp(edge_emb))
108
+
109
+ ni = (
110
+ 0
111
+ if index is None
112
+ else (int(index[0]) if torch.is_tensor(index) else int(index))
113
+ )
114
+ pred_masked = backend.forward(model, x, edge_index, edge_weight=edge_mask)
115
+ pred_node = (
116
+ pred_masked[ni] if pred_masked.dim() == 2 else pred_masked.unsqueeze(0)
117
+ )
118
+
119
+ node_importance = self._node_importance(edge_mask, edge_index, num_nodes)
120
+
121
+ return Explanation(
122
+ node_importance=node_importance,
123
+ edge_importance=edge_mask.detach().cpu(),
124
+ feature_importance=None,
125
+ prediction_original=(
126
+ base_logits[ni].detach().reshape(1, -1).cpu()
127
+ if base_logits.dim() == 2
128
+ else base_logits.detach().reshape(1, -1).cpu()
129
+ ),
130
+ prediction_explanation=pred_node.detach().reshape(1, -1).cpu(),
131
+ node_idx=None if index is None else ni,
132
+ target_class=target_class,
133
+ metadata={},
134
+ )
135
+
136
+ def _sample_mask(
137
+ self,
138
+ mlp: _EdgeMaskMLP,
139
+ edge_emb: torch.Tensor,
140
+ device: torch.device,
141
+ ) -> torch.Tensor:
142
+ logits = mlp(edge_emb)
143
+ noise = torch.rand_like(logits)
144
+ gumbel_noise = torch.log(noise) - torch.log(1 - noise + 1e-8)
145
+ return torch.sigmoid((logits + gumbel_noise) / self.temp)
146
+
147
+ @staticmethod
148
+ def _node_importance(
149
+ edge_mask: torch.Tensor,
150
+ edge_index: torch.Tensor,
151
+ num_nodes: int,
152
+ ) -> torch.Tensor:
153
+ w = edge_mask.detach()
154
+ src, dst = edge_index[0], edge_index[1]
155
+ node_imp = torch.zeros(num_nodes, dtype=torch.float32, device=edge_mask.device)
156
+ node_imp.scatter_add_(0, src, w)
157
+ node_imp.scatter_add_(0, dst, w)
158
+ degree = torch.zeros(num_nodes, dtype=torch.float32, device=edge_mask.device)
159
+ degree.scatter_add_(0, src, torch.ones_like(src, dtype=torch.float32))
160
+ degree.scatter_add_(0, dst, torch.ones_like(dst, dtype=torch.float32))
161
+ degree.clamp_(min=1)
162
+ return (node_imp / degree).cpu()
@@ -0,0 +1,393 @@
1
+ from __future__ import annotations
2
+
3
+ import math
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
+ class MCTSNode:
14
+ __slots__ = ("children", "parent", "reward", "state", "value", "visits")
15
+
16
+ def __init__(self, state, parent=None):
17
+ self.state = frozenset(state)
18
+ self.parent = parent
19
+ self.children: dict[int, MCTSNode] = {}
20
+ self.visits = 0
21
+ self.value = 0.0
22
+ self.reward = -math.inf
23
+
24
+
25
+ @register("subgraphx", "subgraph_x")
26
+ class SubgraphX(ExplanationAlgorithm):
27
+ def __init__(
28
+ self,
29
+ num_hops: int = 3,
30
+ rollout: int = 30,
31
+ high_cpu: bool = False,
32
+ reward_method: str = "mi",
33
+ lambda_connect: float = 0.5,
34
+ lambda_size: float = 0.05,
35
+ max_nodes: int = 20,
36
+ prune: bool = True,
37
+ seed: int = 0,
38
+ **kwargs,
39
+ ):
40
+ self.num_hops = num_hops
41
+ self.rollout = rollout
42
+ self.high_cpu = high_cpu
43
+ self.reward_method = reward_method
44
+ self.lambda_connect = lambda_connect
45
+ self.lambda_size = lambda_size
46
+ self.max_nodes = max_nodes
47
+ self.prune = prune
48
+ self.seed = seed
49
+
50
+ def explain(
51
+ self,
52
+ backend: Any,
53
+ model: Any,
54
+ data: Any,
55
+ index: int | torch.Tensor | None,
56
+ target_class: int | None = None,
57
+ **kwargs,
58
+ ) -> Explanation:
59
+ if index is None:
60
+ raise ValueError("SubgraphX requiere un nodo de interés (index)")
61
+ model.eval()
62
+ x = backend.node_features(data)
63
+ edge_index = backend.edge_index(data)
64
+ num_nodes = backend.num_nodes(data)
65
+
66
+ ni = int(index[0]) if torch.is_tensor(index) else int(index)
67
+ device = x.device
68
+ baseline = x.mean(dim=0, keepdim=True)
69
+
70
+ logits = backend.forward(model, x, edge_index)
71
+ if target_class is None:
72
+ target_class = int(logits[ni].argmax().item()) if logits.dim() == 2 else 0
73
+
74
+ candidates = self._neighborhood(
75
+ edge_index, num_nodes, ni, self.num_hops, self.max_nodes
76
+ )
77
+ if not candidates:
78
+ candidates = list(range(num_nodes))
79
+ candidates.remove(ni)
80
+
81
+ best_state, _ = self._mcts(
82
+ backend, model, x, edge_index, ni, candidates, target_class, baseline
83
+ )
84
+
85
+ selected = set(best_state)
86
+ selected = self._grow_if_unfaithful(
87
+ backend,
88
+ model,
89
+ x,
90
+ edge_index,
91
+ ni,
92
+ target_class,
93
+ selected,
94
+ candidates,
95
+ baseline,
96
+ )
97
+ if self.prune:
98
+ selected = self._prune(
99
+ backend, model, x, edge_index, ni, target_class, selected, baseline
100
+ )
101
+
102
+ edge_mask = self._edge_mask(edge_index, selected, num_nodes, device)
103
+ node_importance = self._node_importance(edge_index, selected, num_nodes, device)
104
+
105
+ with torch.no_grad():
106
+ x_eff = self._masked_features(x, selected, baseline)
107
+ pred_masked = backend.forward(
108
+ model, x_eff, edge_index, edge_weight=edge_mask
109
+ )
110
+ pred_node = (
111
+ pred_masked[ni] if pred_masked.dim() == 2 else pred_masked.unsqueeze(0)
112
+ )
113
+
114
+ return Explanation(
115
+ node_importance=node_importance.cpu(),
116
+ edge_importance=edge_mask.cpu(),
117
+ feature_importance=None,
118
+ prediction_original=(
119
+ logits[ni].detach().reshape(1, -1).cpu()
120
+ if logits.dim() == 2
121
+ else logits.detach().reshape(1, -1).cpu()
122
+ ),
123
+ prediction_explanation=pred_node.detach().reshape(1, -1).cpu(),
124
+ node_idx=ni,
125
+ target_class=target_class,
126
+ metadata={"selected": sorted(selected)},
127
+ )
128
+
129
+ def _neighborhood(self, edge_index, num_nodes, node, hops, max_nodes) -> list[int]:
130
+ visited = {node}
131
+ frontier = {node}
132
+ for _ in range(hops):
133
+ nxt: set[int] = set()
134
+ for u in frontier:
135
+ mask = (edge_index[0] == u) | (edge_index[1] == u)
136
+ nxt.update(edge_index[:, mask].flatten().tolist())
137
+ frontier = nxt - visited
138
+ visited |= frontier
139
+ ordered = sorted(visited)
140
+ if len(ordered) > max_nodes:
141
+ gated = set(ordered[:max_nodes])
142
+ gated.add(node)
143
+ ordered = sorted(gated)
144
+ return ordered
145
+
146
+ def _mcts(
147
+ self, backend, model, x, edge_index, node, candidates, target_class, baseline
148
+ ):
149
+ root_state = (node,)
150
+ root = MCTSNode(root_state)
151
+ for _ in range(self.rollout):
152
+ leaf, path = self._select(root, node, candidates)
153
+ reward = self._evaluate(
154
+ backend, model, x, edge_index, leaf.state, node, target_class, baseline
155
+ )
156
+ self._backprop(path, reward)
157
+ best_leaf = self._best_child(root)
158
+ return best_leaf.state, best_leaf.reward
159
+
160
+ def _select(self, root, node, candidates):
161
+ current = root
162
+ path = [root]
163
+ while current.children:
164
+ unexplored = [
165
+ n
166
+ for n in candidates
167
+ if n not in current.state and n not in current.children
168
+ ]
169
+ if unexplored:
170
+ nxt_node = unexplored[0]
171
+ new_state = tuple(sorted(current.state | {nxt_node}))
172
+ child = MCTSNode(new_state, parent=current)
173
+ current.children[nxt_node] = child
174
+ path.append(child)
175
+ return child, path
176
+ best_child, best_score = None, -math.inf
177
+ for child in current.children.values():
178
+ uct = (child.value / max(child.visits, 1)) + math.sqrt(
179
+ 2 * math.log(max(root.visits, 1)) / (child.visits + 1)
180
+ )
181
+ if uct > best_score:
182
+ best_score, best_child = uct, child
183
+ if best_child is None or len(current.state) >= len(candidates):
184
+ return current, path
185
+ current = best_child
186
+ path.append(current)
187
+ return current, path
188
+
189
+ def _evaluate(
190
+ self, backend, model, x, edge_index, state, node, target_class, baseline
191
+ ):
192
+ selected = set(state)
193
+ if len(selected) == 0:
194
+ return -1e6
195
+ weight = self._edge_mask(edge_index, selected, x.size(0), x.device)
196
+ x_eff = self._masked_features(x, selected, baseline)
197
+ with torch.no_grad():
198
+ pred = backend.forward(model, x_eff, edge_index, edge_weight=weight)
199
+ if pred.dim() == 2:
200
+ logits_n = pred[node]
201
+ logp = logits_n.log_softmax(-1)[target_class].item()
202
+ if int(logits_n.argmax().item()) != target_class:
203
+ logp -= 20.0
204
+ else:
205
+ logp = float(pred[node])
206
+ components = self._connectivity(edge_index, selected)
207
+ size_penalty = self.lambda_size * (len(selected) - 1)
208
+ return logp - self.lambda_connect * (components - 1) - size_penalty
209
+
210
+ def _grow_if_unfaithful(
211
+ self,
212
+ backend,
213
+ model,
214
+ x,
215
+ edge_index,
216
+ node,
217
+ target_class,
218
+ selected,
219
+ candidates,
220
+ baseline,
221
+ ):
222
+ selected = set(selected)
223
+ if not self._preserves_class(
224
+ backend, model, x, edge_index, node, target_class, selected, baseline
225
+ ):
226
+ remaining = [c for c in candidates if c not in selected]
227
+ added = set()
228
+ while remaining:
229
+ best, best_reward = None, -math.inf
230
+ for u in remaining:
231
+ trial = selected | {u}
232
+ weight = self._edge_mask(edge_index, trial, x.size(0), x.device)
233
+ x_eff = self._masked_features(x, trial, baseline)
234
+ with torch.no_grad():
235
+ pred = backend.forward(
236
+ model, x_eff, edge_index, edge_weight=weight
237
+ )
238
+ if pred.dim() != 2:
239
+ continue
240
+ r = pred[node].log_softmax(-1)[target_class].item()
241
+ if int(pred[node].argmax().item()) != target_class:
242
+ r -= 20.0
243
+ if r > best_reward:
244
+ best_reward, best = r, u
245
+ if best is None:
246
+ break
247
+ selected.add(best)
248
+ remaining.remove(best)
249
+ added.add(best)
250
+ if self._preserves_class(
251
+ backend,
252
+ model,
253
+ x,
254
+ edge_index,
255
+ node,
256
+ target_class,
257
+ selected,
258
+ baseline,
259
+ ):
260
+ break
261
+ return selected
262
+
263
+ def _preserves_class(
264
+ self, backend, model, x, edge_index, node, target_class, selected, baseline
265
+ ) -> bool:
266
+ if len(selected) == 0:
267
+ return False
268
+ weight = self._edge_mask(edge_index, selected, x.size(0), x.device)
269
+ x_eff = self._masked_features(x, selected, baseline)
270
+ with torch.no_grad():
271
+ pred = backend.forward(model, x_eff, edge_index, edge_weight=weight)
272
+ if pred.dim() != 2:
273
+ return True
274
+ return (
275
+ not (node >= pred.size(0))
276
+ and int(pred[node].argmax().item()) == target_class
277
+ )
278
+
279
+ @staticmethod
280
+ def _masked_features(x, selected, baseline):
281
+ sel = torch.zeros(x.size(0), dtype=torch.bool, device=x.device)
282
+ sel[list(selected)] = True
283
+ return baseline + (x - baseline) * sel.unsqueeze(-1)
284
+
285
+ def _prune(
286
+ self, backend, model, x, edge_index, node, target_class, selected, baseline
287
+ ):
288
+ selected = set(selected)
289
+ weight = self._edge_mask(edge_index, selected, x.size(0), x.device)
290
+ x_eff = self._masked_features(x, selected, baseline)
291
+ with torch.no_grad():
292
+ pred = backend.forward(model, x_eff, edge_index, edge_weight=weight)
293
+ if pred.dim() != 2:
294
+ return selected
295
+ if int(pred[node].argmax().item()) != target_class:
296
+ return selected
297
+ keep = sorted(selected)
298
+ if len(keep) > 1:
299
+ ordered = self._removal_order(
300
+ backend, model, x, edge_index, node, target_class, keep, baseline
301
+ )
302
+ for u in ordered:
303
+ trial = set(keep) - {u}
304
+ if u == node or not trial:
305
+ continue
306
+ t_weight = self._edge_mask(edge_index, trial, x.size(0), x.device)
307
+ t_x = self._masked_features(x, trial, baseline)
308
+ with torch.no_grad():
309
+ t_pred = backend.forward(
310
+ model, t_x, edge_index, edge_weight=t_weight
311
+ )
312
+ if int(t_pred[node].argmax().item()) == target_class:
313
+ keep = list(trial)
314
+ return set(keep)
315
+
316
+ def _removal_order(
317
+ self, backend, model, x, edge_index, node, target_class, keep, baseline
318
+ ):
319
+ scores = []
320
+ for u in keep:
321
+ if u == node:
322
+ continue
323
+ trial = set(keep) - {u}
324
+ t_weight = self._edge_mask(edge_index, trial, x.size(0), x.device)
325
+ t_x = self._masked_features(x, trial, baseline)
326
+ with torch.no_grad():
327
+ t_pred = backend.forward(model, t_x, edge_index, edge_weight=t_weight)
328
+ if int(t_pred[node].argmax().item()) != target_class:
329
+ scores.append((u, -math.inf))
330
+ else:
331
+ logp_drop = t_pred[node].log_softmax(-1)[target_class].item()
332
+ scores.append((u, logp_drop))
333
+ return [u for u, _ in sorted(scores, key=lambda t: t[1])]
334
+
335
+ def _backprop(self, path, reward):
336
+ for n in reversed(path):
337
+ n.visits += 1
338
+ n.value += reward
339
+ n.reward = max(n.reward, reward)
340
+
341
+ def _best_child(self, root):
342
+ best = root
343
+ queue = [root]
344
+ while queue:
345
+ cur = queue.pop(0)
346
+ if cur.reward > best.reward:
347
+ best = cur
348
+ queue.extend(cur.children.values())
349
+ return best
350
+
351
+ @staticmethod
352
+ def _edge_mask(edge_index, selected, num_nodes, device):
353
+ src = edge_index[0]
354
+ dst = edge_index[1]
355
+ both = torch.isin(
356
+ src, torch.as_tensor(list(selected), device=device)
357
+ ) & torch.isin(dst, torch.as_tensor(list(selected), device=device))
358
+ return both.to(torch.float32)
359
+
360
+ @staticmethod
361
+ def _node_importance(edge_index, selected, num_nodes, device):
362
+ imp = torch.zeros(num_nodes, dtype=torch.float32, device=device)
363
+ sel = torch.as_tensor(list(selected), device=device)
364
+ for s in sel:
365
+ imp[s] = 1.0
366
+ return imp
367
+
368
+ @staticmethod
369
+ def _connectivity(edge_index, selected):
370
+ sel = set(selected)
371
+ seen = set()
372
+ comps = 0
373
+ adj = {}
374
+ src = edge_index[0].tolist()
375
+ dst = edge_index[1].tolist()
376
+ for u, v in zip(src, dst):
377
+ if u in sel and v in sel:
378
+ adj.setdefault(u, set()).add(v)
379
+ adj.setdefault(v, set()).add(u)
380
+ for s in sel:
381
+ if s in seen:
382
+ continue
383
+ comps += 1
384
+ stack = [s]
385
+ while stack:
386
+ cur = stack.pop()
387
+ if cur in seen:
388
+ continue
389
+ seen.add(cur)
390
+ for nb in adj.get(cur, ()):
391
+ if nb not in seen:
392
+ stack.append(nb)
393
+ return comps