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.
- graph_explain/__init__.py +79 -0
- graph_explain/backends/__init__.py +4 -0
- graph_explain/backends/base.py +103 -0
- graph_explain/backends/dgl.py +121 -0
- graph_explain/benchmarks/__init__.py +3 -0
- graph_explain/benchmarks/synthetic.py +246 -0
- graph_explain/cli.py +459 -0
- graph_explain/core/__init__.py +14 -0
- graph_explain/core/benchmark.py +284 -0
- graph_explain/core/evaluation.py +391 -0
- graph_explain/core/explainer.py +83 -0
- graph_explain/core/explanation.py +72 -0
- graph_explain/core/model_utils.py +44 -0
- graph_explain/core/registry.py +55 -0
- graph_explain/methods/__init__.py +39 -0
- graph_explain/methods/attention/attention.py +147 -0
- graph_explain/methods/base.py +25 -0
- graph_explain/methods/baseline/random_baseline.py +78 -0
- graph_explain/methods/counterfactual/counterfactual.py +304 -0
- graph_explain/methods/feature/graph_lime.py +141 -0
- graph_explain/methods/gradient/__init__.py +0 -0
- graph_explain/methods/gradient/grad_x_input.py +110 -0
- graph_explain/methods/gradient/guided_backprop.py +117 -0
- graph_explain/methods/gradient/integrated_gradients.py +115 -0
- graph_explain/methods/gradient/saliency.py +93 -0
- graph_explain/methods/perturbation/__init__.py +0 -0
- graph_explain/methods/perturbation/gnn_explainer.py +265 -0
- graph_explain/methods/perturbation/node_mask.py +136 -0
- graph_explain/methods/perturbation/pg_explainer.py +162 -0
- graph_explain/methods/perturbation/subgraphx.py +393 -0
- graph_explain/methods/relevance/deeplift.py +262 -0
- graph_explain/methods/relevance/gnn_lrp.py +219 -0
- graph_explain/narration/__init__.py +3 -0
- graph_explain/narration/narrator.py +185 -0
- graph_explain/visualization/__init__.py +4 -0
- graph_explain/visualization/interactive.py +73 -0
- graph_explain/visualization/static.py +90 -0
- graph_explain-0.7.0.dist-info/METADATA +332 -0
- graph_explain-0.7.0.dist-info/RECORD +42 -0
- graph_explain-0.7.0.dist-info/WHEEL +5 -0
- graph_explain-0.7.0.dist-info/entry_points.txt +2 -0
- 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> · target class
|
|
253
|
+
<code>{meta["target_class"]}</code> · 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)
|