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,284 @@
1
+ from __future__ import annotations
2
+
3
+ import torch
4
+
5
+ from ..narration import summarize
6
+ from .evaluation import (
7
+ evaluate_fidelity_minus,
8
+ evaluate_fidelity_plus,
9
+ evaluate_gea,
10
+ evaluate_gea_graph,
11
+ evaluate_sparsity,
12
+ evaluate_stability,
13
+ )
14
+ from .explainer import Explainer
15
+ from .registry import get_algorithm, instantiate
16
+
17
+ DEFAULT_METHODS = [
18
+ "gnn_explainer",
19
+ "pg_explainer",
20
+ "subgraphx",
21
+ "saliency",
22
+ "integrated_gradients",
23
+ "gnn_lrp",
24
+ "deep_lift",
25
+ "grad_x_input",
26
+ "graph_lime",
27
+ "node_mask",
28
+ "guided_backprop",
29
+ "random",
30
+ "counterfactual",
31
+ "attention",
32
+ ]
33
+
34
+ _METRICS = (
35
+ "fidelity_plus",
36
+ "fidelity_minus",
37
+ "gea",
38
+ "sparsity",
39
+ "sparsity_local",
40
+ "stability",
41
+ )
42
+
43
+
44
+ def _has_gat(model) -> bool:
45
+ try:
46
+ from torch_geometric.nn import GATConv
47
+
48
+ return any(isinstance(m, GATConv) for m in model.modules())
49
+ except ImportError:
50
+ return False
51
+
52
+
53
+ def _method_kwargs(epochs, lr, seed, top_k) -> dict:
54
+ kw: dict = {}
55
+ for name, value in (
56
+ ("epochs", epochs),
57
+ ("lr", lr),
58
+ ("seed", seed),
59
+ ("top_k", top_k),
60
+ ):
61
+ if value is not None:
62
+ kw[name] = value
63
+ return kw
64
+
65
+
66
+ def compare(
67
+ data,
68
+ model,
69
+ node: int | None = None,
70
+ target_class: int | None = None,
71
+ backend: str = "pyg",
72
+ methods: list[str] | None = None,
73
+ top_k: int = 5,
74
+ num_perturbations: int = 5,
75
+ noise_std: float = 0.05,
76
+ epochs: int = 200,
77
+ lr: float = 0.01,
78
+ seed: int = 0,
79
+ mask_threshold: float = 0.5,
80
+ stability: bool = True,
81
+ ) -> dict:
82
+ """Runs several explanation methods and compares their metrics.
83
+
84
+ Uses `node` for node-level (required) or `node=None` for graph-level
85
+ (node-only methods are marked `skipped`). Returns a dictionary with one
86
+ entry per method: class, predictions, metrics (fidelity±, GEA, sparsity,
87
+ stability) and the structured `summarize` summary. Non-applicable methods
88
+ or failing metrics are marked as `skipped`/`None` without aborting the
89
+ rest.
90
+ """
91
+ from ..backends import get_backend
92
+
93
+ if methods is None:
94
+ methods = list(DEFAULT_METHODS)
95
+ backend_obj = get_backend(backend)
96
+ model.eval()
97
+
98
+ if node is None:
99
+ stability = False
100
+
101
+ torch.manual_seed(seed)
102
+ results: dict = {}
103
+ ran: list[str] = []
104
+ skipped: dict[str, str] = {}
105
+
106
+ for name in methods:
107
+ cls = get_algorithm(name)
108
+ entry = {
109
+ "method": name,
110
+ "class": cls.__name__,
111
+ "node": node,
112
+ "target_class": target_class,
113
+ "prediction_original": None,
114
+ "prediction_explanation": None,
115
+ "metrics": {m: None for m in _METRICS},
116
+ "summary": None,
117
+ "skipped": None,
118
+ }
119
+ if node is None and not cls.graph_level:
120
+ entry["skipped"] = "node-level only"
121
+ results[name] = entry
122
+ skipped[name] = entry["skipped"]
123
+ continue
124
+ if name == "attention" and not _has_gat(model):
125
+ entry["skipped"] = "requires a model with GATConv layers"
126
+ results[name] = entry
127
+ skipped[name] = entry["skipped"]
128
+ continue
129
+ algo = instantiate(name, **_method_kwargs(epochs, lr, seed, top_k))
130
+ explainer = Explainer(
131
+ algorithm=algo,
132
+ backend=backend_obj,
133
+ mask_threshold=mask_threshold,
134
+ )
135
+ torch.manual_seed(seed)
136
+ try:
137
+ expl = (
138
+ explainer.explain_node(data, model, node, target_class=target_class)
139
+ if node is not None
140
+ else explainer.explain_graph(data, model, target_class=target_class)
141
+ )
142
+ except (ValueError, TypeError) as exc:
143
+ entry["skipped"] = str(exc)
144
+ results[name] = entry
145
+ skipped[name] = str(exc)
146
+ continue
147
+
148
+ entry["prediction_original"] = _fmt(expl.prediction_original)
149
+ entry["prediction_explanation"] = _fmt(expl.prediction_explanation)
150
+ entry["summary"] = summarize(expl, data=data, top_k=top_k)
151
+
152
+ m = entry["metrics"]
153
+ expl_arg = expl
154
+ m["fidelity_plus"] = _safe(
155
+ lambda expl=expl_arg: float(evaluate_fidelity_plus(model, expl))
156
+ )
157
+ m["fidelity_minus"] = _safe(
158
+ lambda expl=expl_arg: float(evaluate_fidelity_minus(model, expl))
159
+ )
160
+ if node is not None:
161
+ m["gea"] = _safe(
162
+ lambda expl=expl_arg: float(evaluate_gea(expl, data=data, top_k=top_k))
163
+ )
164
+ else:
165
+ m["gea"] = _safe(
166
+ lambda expl=expl_arg: float(
167
+ evaluate_gea_graph(expl, data=data, top_k=top_k)
168
+ )
169
+ )
170
+ m["sparsity"] = _safe(lambda expl=expl_arg: float(evaluate_sparsity(expl)))
171
+ m["sparsity_local"] = _safe(
172
+ lambda expl=expl_arg: float(evaluate_sparsity(expl, local=True))
173
+ )
174
+ if stability:
175
+
176
+ def _again(d, name=name):
177
+ algo_r = instantiate(name, **_method_kwargs(epochs, lr, seed, top_k))
178
+ return Explainer(
179
+ algorithm=algo_r,
180
+ backend=backend_obj,
181
+ mask_threshold=mask_threshold,
182
+ ).explain_node(d, model, node)
183
+
184
+ m["stability"] = _safe(
185
+ lambda: float(
186
+ evaluate_stability(
187
+ _again,
188
+ data,
189
+ num_perturbations=num_perturbations,
190
+ noise_std=noise_std,
191
+ top_k=top_k,
192
+ )
193
+ )
194
+ )
195
+ results[name] = entry
196
+ ran.append(name)
197
+
198
+ return {
199
+ "_meta": {
200
+ "node": node,
201
+ "target_class": target_class,
202
+ "backend": backend,
203
+ "methods": ran,
204
+ "skipped": skipped,
205
+ },
206
+ **{name: results[name] for name in methods},
207
+ }
208
+
209
+
210
+ def report_html(results: dict, output_path: str) -> None:
211
+ """Builds a self-contained HTML report (comparative table)."""
212
+ meta = results["_meta"]
213
+ rows = []
214
+ for name, entry in results.items():
215
+ if name.startswith("_"):
216
+ continue
217
+ m = entry["metrics"]
218
+ if entry["skipped"]:
219
+ rows.append(
220
+ f"<tr><td>{name}</td>"
221
+ f"<td colspan='7' class='skip'>not applicable: {entry['skipped']}</td></tr>"
222
+ )
223
+ continue
224
+ cells = "".join(
225
+ f"<td>{'-' if m[k] is None else f'{m[k]:.4f}'}</td>"
226
+ for k in ("fidelity_plus", "fidelity_minus", "gea", "sparsity", "stability")
227
+ )
228
+ rows.append(
229
+ f"<tr><td>{name} <small>({entry['class']})</small></td>{cells}</tr>"
230
+ )
231
+
232
+ body = "\n".join(rows)
233
+ html = f"""<!doctype html>
234
+ <html lang="es">
235
+ <head>
236
+ <meta charset="utf-8">
237
+ <title>Benchmark - graph-explain</title>
238
+ <style>
239
+ body {{ font-family: system-ui, sans-serif; margin: 2rem; }}
240
+ table {{ border-collapse: collapse; width: 100%; max-width: 900px; }}
241
+ th, td {{ border: 1px solid #ccc; padding: 6px 10px; text-align: right; }}
242
+ th {{ background: #f0f0f0; }}
243
+ td:first-child {{ text-align: left; }}
244
+ td.skip {{ text-align: left; color: #888; font-style: italic; }}
245
+ .meta {{ color: #555; margin-bottom: 1rem; }}
246
+ code {{ background: #f4f4f4; padding: 0 4px; }}
247
+ </style>
248
+ </head>
249
+ <body>
250
+ <h1>Comparative explanation benchmark</h1>
251
+ <p class="meta">
252
+ node <code>{meta["node"]}</code> &middot; target class
253
+ <code>{meta["target_class"]}</code> &middot; backend <code>{meta["backend"]}</code>
254
+ </p>
255
+ <table>
256
+ <tr>
257
+ <th>Method</th><th>fid+</th><th>fid-</th><th>GEA</th><th>sparsity</th><th>stability</th>
258
+ </tr>
259
+ {body}
260
+ </table>
261
+ <p class="meta">
262
+ Generated with <code>graph-explain</code>. fid+ = necessity (drop in P(c) after
263
+ removing top-k), fid- = sufficiency, GEA = overlap with ground truth, sparsity =
264
+ global sparsity, stability = mean similarity under perturbations.
265
+ </p>
266
+ </body>
267
+ </html>"""
268
+ with open(output_path, "w", encoding="utf-8") as fh:
269
+ fh.write(html)
270
+
271
+
272
+ def _fmt(value):
273
+ if value is None:
274
+ return None
275
+ if hasattr(value, "tolist"):
276
+ return [round(float(v), 4) for v in value.reshape(-1).tolist()]
277
+ return value
278
+
279
+
280
+ def _safe(fn):
281
+ try:
282
+ return fn()
283
+ except Exception: # noqa: BLE001
284
+ return None
@@ -0,0 +1,391 @@
1
+ from __future__ import annotations
2
+
3
+ import math
4
+
5
+
6
+ def _resolve_target(explanation) -> int:
7
+ target = explanation.target_class
8
+ if target is None:
9
+ import torch
10
+
11
+ pred = explanation.prediction_original
12
+ if torch.is_tensor(pred):
13
+ pred = pred.reshape(-1)
14
+ target = int(pred.argmax().item())
15
+ else:
16
+ target = 0
17
+ return int(target)
18
+
19
+
20
+ def _mask_top_k(explanation, top_k: int | None, keep_ratio: float, kind: str):
21
+ importance = (
22
+ explanation.edge_importance if kind == "edge" else explanation.node_importance
23
+ )
24
+ if importance is None:
25
+ raise ValueError(
26
+ f"La explicación no tiene edge_importance/node_importance para "
27
+ f"fidelidad por {kind}."
28
+ )
29
+ import torch
30
+
31
+ imp = importance.detach().cpu().reshape(-1)
32
+ n = imp.shape[0]
33
+ if top_k is None:
34
+ top_k = round(n * keep_ratio)
35
+ top_k = int(max(0, min(top_k, n)))
36
+ if top_k == 0:
37
+ return torch.zeros(n, dtype=torch.bool)
38
+ idx = imp.argsort(descending=True)[:top_k]
39
+ mask = torch.zeros(n, dtype=torch.bool)
40
+ mask[idx] = True
41
+ return mask
42
+
43
+
44
+ def _context(model, explanation):
45
+ backend = explanation.metadata.get("backend")
46
+ data = explanation.metadata.get("backing_data")
47
+ if backend is None or data is None:
48
+ raise ValueError(
49
+ "Para fidelidad± es necesario que la explicación tenga backend y "
50
+ "backing_data en metadata (úsala a través de Explainer)."
51
+ )
52
+ return backend, data
53
+
54
+
55
+ def evaluate_fidelity_plus(
56
+ model,
57
+ explanation,
58
+ top_k: int | None = None,
59
+ keep_ratio: float = 0.1,
60
+ kind: str = "edge",
61
+ ) -> float:
62
+ """Necessity: drop in P(c) when removing the top-k important elements."""
63
+ target = _resolve_target(explanation)
64
+ backend, data = _context(model, explanation)
65
+ import torch
66
+
67
+ x = backend.node_features(data)
68
+ edge_index = backend.edge_index(data)
69
+ device = x.device
70
+ with torch.no_grad():
71
+ logits_orig = backend.forward(model, x, edge_index)
72
+ node = _explained_node(explanation)
73
+ p_orig = _prob(logits_orig[node], target, softmax=True)
74
+ mask = _mask_top_k(explanation, top_k, keep_ratio, kind)
75
+ if kind == "edge":
76
+ edge_weight = torch.ones(edge_index.size(1), device=device)
77
+ edge_weight[mask] = 0.0
78
+ masked = backend.forward(model, x, edge_index, edge_weight=edge_weight)
79
+ p_masked = _prob(masked[node], target, softmax=True)
80
+ else:
81
+ node_mask = torch.ones(x.size(0), device=device)
82
+ node_mask[mask] = 0.0
83
+ masked = backend.forward(model, x, edge_index, node_mask=node_mask)
84
+ p_masked = _prob(masked[node], target, softmax=True)
85
+ return float(p_orig - p_masked)
86
+
87
+
88
+ def evaluate_fidelity_minus(
89
+ model,
90
+ explanation,
91
+ top_k: int | None = None,
92
+ keep_ratio: float = 0.1,
93
+ kind: str = "edge",
94
+ ) -> float:
95
+ """Sufficiency: P(c) preserved when keeping ONLY the top-k elements."""
96
+ target = _resolve_target(explanation)
97
+ backend, data = _context(model, explanation)
98
+ import torch
99
+
100
+ x = backend.node_features(data)
101
+ edge_index = backend.edge_index(data)
102
+ device = x.device
103
+ with torch.no_grad():
104
+ node = _explained_node(explanation)
105
+ mask = _mask_top_k(explanation, top_k, keep_ratio, kind)
106
+ if kind == "edge":
107
+ edge_weight = torch.zeros(edge_index.size(1), device=device)
108
+ edge_weight[mask] = 1.0
109
+ masked = backend.forward(model, x, edge_index, edge_weight=edge_weight)
110
+ else:
111
+ node_mask = torch.zeros(x.size(0), device=device)
112
+ node_mask[mask] = 1.0
113
+ masked = backend.forward(model, x, edge_index, node_mask=node_mask)
114
+ p_kept = _prob(masked[node], target, softmax=True)
115
+ return float(p_kept)
116
+
117
+
118
+ def _explained_node(explanation) -> int:
119
+ if explanation.node_idx is None:
120
+ return 0
121
+ return int(explanation.node_idx)
122
+
123
+
124
+ def _perturbed(data, perturbation: str, noise_std: float, num_edges: int | None, rng):
125
+ import copy
126
+
127
+ import torch
128
+
129
+ d = copy.deepcopy(data)
130
+ if perturbation == "feature":
131
+ x = d.x
132
+ noise = torch.randn_like(x) * noise_std
133
+ d.x = x + noise
134
+ elif perturbation == "edge" and hasattr(d, "edge_index"):
135
+ ei = d.edge_index
136
+ n = ei.size(1)
137
+ if num_edges is None:
138
+ num_edges = max(1, n // 10)
139
+ num_edges = min(num_edges, n)
140
+ drop = rng.choice(n, size=num_edges, replace=False)
141
+ keep = [i for i in range(n) if i not in set(drop.tolist())]
142
+ d.edge_index = ei[:, keep]
143
+ if hasattr(d, "edge_weight") and d.edge_weight is not None:
144
+ d.edge_weight = d.edge_weight[keep]
145
+ else:
146
+ raise ValueError(f"perturbación desconocida: {perturbation}")
147
+ return d
148
+
149
+
150
+ def evaluate_stability(
151
+ get_explanation,
152
+ data,
153
+ num_perturbations: int = 10,
154
+ perturbation: str = "feature",
155
+ noise_std: float = 0.05,
156
+ num_edges: int | None = None,
157
+ top_k: int | None = None,
158
+ seed: int = 0,
159
+ ) -> float:
160
+ """Stability: mean similarity between explanations under small graph
161
+ perturbations. `get_explanation` receives a Data and returns an
162
+ Explanation."""
163
+ from itertools import pairwise
164
+
165
+ import numpy as np
166
+
167
+ rng = np.random.default_rng(seed)
168
+ exps = [
169
+ get_explanation(_perturbed(data, perturbation, noise_std, num_edges, rng))
170
+ for _ in range(num_perturbations)
171
+ ]
172
+ if not exps:
173
+ return 1.0
174
+ sims = []
175
+ for a, b in pairwise(exps):
176
+ sims.append(_explanation_similarity(a, b, top_k))
177
+ valid = [
178
+ s for s in sims if math.isfinite(s)
179
+ ] # descarta NaN (p.ej. perturbación de aristas)
180
+ return float(np.mean(valid)) if valid else 1.0
181
+
182
+
183
+ def _explanation_similarity(a, b, top_k: int | None) -> float:
184
+ import torch
185
+
186
+ va = _importance_vector(a)
187
+ vb = _importance_vector(b)
188
+ if va is None or vb is None or va.numel() != vb.numel():
189
+ return float("nan")
190
+ if top_k is not None:
191
+ ta = set(va.argsort(descending=True)[:top_k].tolist())
192
+ tb = set(vb.argsort(descending=True)[:top_k].tolist())
193
+ union = ta | tb
194
+ if not union:
195
+ return 1.0
196
+ return float(len(ta & tb) / len(union))
197
+ if va.norm().item() < 1e-9 or vb.norm().item() < 1e-9:
198
+ return float("nan")
199
+ cos = float(torch.nn.functional.cosine_similarity(va, vb, dim=0).item())
200
+ return max(0.0, min(1.0, cos))
201
+
202
+
203
+ def _importance_vector(explanation):
204
+ parts = []
205
+ if explanation.node_importance is not None:
206
+ parts.append(explanation.node_importance.detach().reshape(-1).cpu().float())
207
+ if explanation.edge_importance is not None:
208
+ parts.append(explanation.edge_importance.detach().reshape(-1).cpu().float())
209
+ if not parts:
210
+ return None
211
+ import torch
212
+
213
+ return torch.cat(parts)
214
+
215
+
216
+ def evaluate_gea(explanation, data=None, top_k: int | None = None) -> float:
217
+ """Graph Explanation Accuracy: overlap between the top-k elements of the
218
+ explanation and the benchmark's relevant ground-truth subgraph."""
219
+ if data is None:
220
+ data = explanation.metadata.get("backing_data")
221
+ if data is None:
222
+ raise ValueError("evaluate_gea necesita `data` (o backing_data en metadata).")
223
+ node = _explained_node(explanation)
224
+ gt_nodes, gt_edges = _ground_truth(data, node, explanation.metadata.get("backend"))
225
+ if explanation.edge_importance is not None and gt_edges:
226
+ imp = explanation.edge_importance.detach().reshape(-1)
227
+ k = top_k if top_k is not None else len(gt_edges)
228
+ k = int(max(0, min(k, imp.shape[0])))
229
+ if k == 0:
230
+ return 0.0
231
+ top = set(imp.argsort(descending=True)[:k].tolist())
232
+ return float(len(top & set(gt_edges)) / max(1, k))
233
+ if explanation.node_importance is not None and gt_nodes:
234
+ imp = explanation.node_importance.detach().reshape(-1)
235
+ k = top_k if top_k is not None else len(gt_nodes)
236
+ k = int(max(0, min(k, imp.shape[0])))
237
+ if k == 0:
238
+ return 0.0
239
+ top = set(imp.argsort(descending=True)[:k].tolist())
240
+ return float(len(top & set(gt_nodes)) / max(1, k))
241
+ raise ValueError("evaluate_gea necesita edge_importance o node_importance.")
242
+
243
+
244
+ def _ground_truth(data, node, backend=None):
245
+ try:
246
+ from ..benchmarks.synthetic import ground_truth_edge_ids, ground_truth_nodes
247
+ except ImportError:
248
+ raise ValueError("Ground truth disponible solo con el benchmark sintético.")
249
+ edge_index = backend.edge_index(data) if backend is not None else data.edge_index
250
+ gt_nodes = ground_truth_nodes(data, node)
251
+ gt_edges = ground_truth_edge_ids(data, node, edge_index)
252
+ return gt_nodes, gt_edges
253
+
254
+
255
+ def evaluate_fidelity(explanation, keep_ratio: float = 0.2) -> float:
256
+ if (
257
+ explanation.prediction_original is None
258
+ or explanation.prediction_explanation is None
259
+ ):
260
+ raise ValueError(
261
+ "Para evaluar fidelidad la explicación debe contener "
262
+ "prediction_original y prediction_explanation."
263
+ )
264
+ target = explanation.target_class
265
+ if target is None:
266
+ import torch
267
+
268
+ pred = explanation.prediction_original
269
+ if torch.is_tensor(pred):
270
+ pred = pred.reshape(-1)
271
+ target = int(pred.argmax().item())
272
+ else:
273
+ target = 0
274
+ p_orig = _prob(explanation.prediction_original, target, softmax=True)
275
+ p_expl = _prob(explanation.prediction_explanation, target, softmax=True)
276
+ return float(p_orig - p_expl)
277
+
278
+
279
+ def evaluate_sparsity(explanation, local: bool = False, local_hops: int = 3) -> float:
280
+ node_mask = explanation.node_importance
281
+ edge_mask = explanation.edge_importance
282
+ if node_mask is None and edge_mask is None:
283
+ raise ValueError(
284
+ "Para evaluar esparcidad la explicación necesita node_importance "
285
+ "y/o edge_importance."
286
+ )
287
+ if local and explanation.node_idx is not None:
288
+ node_ids, edge_ids = _local_scope(explanation, local_hops)
289
+ else:
290
+ node_ids = None
291
+ edge_ids = None
292
+ masked = 0.0
293
+ total = 0.0
294
+ threshold = explanation.mask_threshold
295
+ if node_mask is not None:
296
+ total += len(node_ids) if node_ids is not None else node_mask.shape[0]
297
+ if node_ids is not None:
298
+ masked += float((node_mask[node_ids] < threshold).sum())
299
+ else:
300
+ masked += float((node_mask < threshold).sum())
301
+ if edge_mask is not None:
302
+ total += len(edge_ids) if edge_ids is not None else edge_mask.shape[0]
303
+ if edge_ids is not None:
304
+ masked += float((edge_mask[edge_ids] < threshold).sum())
305
+ else:
306
+ masked += float((edge_mask < threshold).sum())
307
+ return 1.0 - masked / total if total > 0 else 1.0
308
+
309
+
310
+ def _local_scope(explanation, hops: int) -> tuple[list[int], list[int]]:
311
+ backend = explanation.metadata.get("backend")
312
+ data = explanation.metadata.get("backing_data")
313
+ edge_index = None
314
+ num_nodes = None
315
+ if backend is not None and data is not None:
316
+ edge_index = backend.edge_index(data)
317
+ num_nodes = backend.num_nodes(data)
318
+ else:
319
+ edge_index = getattr(data, "edge_index", None) if data is not None else None
320
+ num_nodes = getattr(data, "num_nodes", None)
321
+ if edge_index is None:
322
+ edge_index = explanation.metadata.get("edge_index")
323
+ num_nodes = num_nodes or explanation.metadata.get("num_nodes")
324
+ if edge_index is None or num_nodes is None:
325
+ return None, None
326
+ node_idx = int(explanation.node_idx)
327
+ visited = {node_idx}
328
+ frontier = {node_idx}
329
+ for _ in range(hops):
330
+ nxt = set()
331
+ for u in frontier:
332
+ m = (edge_index[0] == u) | (edge_index[1] == u)
333
+ nxt.update(edge_index[:, m].flatten().tolist())
334
+ frontier = nxt - visited
335
+ visited |= frontier
336
+ nodes = sorted(v for v in visited if v < num_nodes)
337
+ edge_ids = [
338
+ i
339
+ for i in range(edge_index.size(1))
340
+ if int(edge_index[0, i]) in visited and int(edge_index[1, i]) in visited
341
+ ]
342
+ return nodes, edge_ids
343
+
344
+
345
+ def evaluate_gea_graph(
346
+ explanation,
347
+ data=None,
348
+ gt_edge_ids: list[int] | None = None,
349
+ top_k: int | None = None,
350
+ ) -> float:
351
+ """Graph Explanation Accuracy (graph-level): overlap of the explanation's
352
+ top-k edges with the dataset's known motif edges (`gt_edge_mask`) of the
353
+ explained graph."""
354
+ if gt_edge_ids is None:
355
+ if data is None:
356
+ data = explanation.metadata.get("backing_data")
357
+ if data is not None:
358
+ from ..benchmarks.synthetic import ground_truth_edges_graph
359
+
360
+ gt_edge_ids = ground_truth_edges_graph(data)
361
+ if not gt_edge_ids:
362
+ raise ValueError(
363
+ "evaluate_gea_graph no tiene ground truth de aristas para este grafo."
364
+ )
365
+ if explanation.edge_importance is None:
366
+ raise ValueError("evaluate_gea_graph necesita edge_importance.")
367
+ imp = explanation.edge_importance.detach().reshape(-1)
368
+ k = top_k if top_k is not None else len(gt_edge_ids)
369
+ k = int(max(0, min(k, imp.shape[0])))
370
+ if k == 0:
371
+ return 0.0
372
+ top = set(imp.argsort(descending=True)[:k].tolist())
373
+ return float(len(top & set(gt_edge_ids)) / max(1, k))
374
+
375
+
376
+ def _prob(pred, class_idx, softmax: bool = False):
377
+ import torch
378
+
379
+ if torch.is_tensor(pred) and pred.dim() > 1:
380
+ pred = pred.reshape(-1)
381
+ if torch.is_tensor(pred) and pred.dim() == 1 and softmax:
382
+ pred = pred.softmax(dim=0)
383
+ if class_idx is None:
384
+ if torch.is_tensor(pred) and pred.numel() == 1:
385
+ return float(pred)
386
+ return float(pred)
387
+ if torch.is_tensor(pred):
388
+ if pred.dim() == 0:
389
+ return float(pred)
390
+ return float(pred[class_idx])
391
+ return float(pred)