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
graph_explain/cli.py ADDED
@@ -0,0 +1,459 @@
1
+ from __future__ import annotations
2
+
3
+ import argparse
4
+ import json
5
+ import sys
6
+ from typing import Any
7
+
8
+ from graph_explain import __version__
9
+ from graph_explain.core.registry import get_algorithm, instantiate
10
+
11
+ _METHODS = [
12
+ "gnn_explainer",
13
+ "gnnexplainer",
14
+ "saliency",
15
+ "gradient",
16
+ "grad",
17
+ "pg_explainer",
18
+ "pgexplainer",
19
+ "subgraphx",
20
+ "subgraph_x",
21
+ "integrated_gradients",
22
+ "ig",
23
+ "gnn_lrp",
24
+ "gnn-lrp",
25
+ "lrp",
26
+ "deep_lift",
27
+ "deeplift",
28
+ "dl",
29
+ "attention",
30
+ "gat",
31
+ "attention_explainer",
32
+ "grad_x_input",
33
+ "gradient_x_input",
34
+ "gx",
35
+ "graph_lime",
36
+ "glime",
37
+ "gl",
38
+ "node_mask",
39
+ "nodemask",
40
+ "nm",
41
+ "guided_backprop",
42
+ "guided-backprop",
43
+ "gbp",
44
+ "random",
45
+ "random_baseline",
46
+ "rand",
47
+ "counterfactual",
48
+ "cf",
49
+ ]
50
+
51
+ _METRICS = ["fidelity", "fidelity_plus", "fidelity_minus", "gea", "stability"]
52
+
53
+
54
+ def build_parser() -> argparse.ArgumentParser:
55
+ parser = argparse.ArgumentParser(
56
+ prog="graph-explain",
57
+ description="Explainability for graph-based models (GNNs).",
58
+ )
59
+ parser.add_argument(
60
+ "--version", action="version", version=f"graph-explain {__version__}"
61
+ )
62
+ sub = parser.add_subparsers(dest="command")
63
+
64
+ explain = sub.add_parser("explain", help="Explain the prediction of a node/graph")
65
+ explain.add_argument("--model", required=True, help="Path to the saved model (.pt)")
66
+ explain.add_argument("--data", required=True, help="Path to the saved Data (.pt)")
67
+ explain.add_argument("--method", default="gnn_explainer", choices=_METHODS)
68
+ explain.add_argument(
69
+ "--node",
70
+ type=int,
71
+ default=None,
72
+ help="Node index to explain (node-level)",
73
+ )
74
+ explain.add_argument("--target-class", type=int, default=None)
75
+ explain.add_argument("--epochs", type=int, default=200)
76
+ explain.add_argument("--lr", type=float, default=None)
77
+ explain.add_argument(
78
+ "--mode",
79
+ default="edge",
80
+ choices=["edge", "feature"],
81
+ help="Counterfactual mode",
82
+ )
83
+ explain.add_argument("--hops", type=int, default=2)
84
+ explain.add_argument("--max-steps", type=int, default=10)
85
+ explain.add_argument("--eps", type=float, default=None)
86
+ explain.add_argument("--steps", type=int, default=50)
87
+ explain.add_argument("--normalize", action="store_true", help="Normalized GNN-LRP")
88
+ explain.add_argument("--backend", default="pyg", choices=["pyg", "dgl"])
89
+ explain.add_argument("--output", default=None, help="Save the explanation to .pt")
90
+ explain.add_argument("--plot", default=None, help="Save visualization to .png/.pdf")
91
+ explain.add_argument(
92
+ "--html",
93
+ default=None,
94
+ help="Save interactive visualization to .html",
95
+ )
96
+ explain.add_argument("--threshold", type=float, default=0.5)
97
+ explain.add_argument(
98
+ "--top-k", type=int, default=5, help="Top-k for GEA/stability/JSON"
99
+ )
100
+ explain.add_argument(
101
+ "--metrics",
102
+ default="",
103
+ help=f"Comma-separated list: {', '.join(_METRICS)}",
104
+ )
105
+ explain.add_argument("--num-perturbations", type=int, default=10)
106
+ explain.add_argument("--noise-std", type=float, default=0.05)
107
+ explain.add_argument("--describe", action="store_true", help="Print the narration")
108
+ explain.add_argument(
109
+ "--json", default=None, help="Export summary + metrics to .json"
110
+ )
111
+
112
+ bench = sub.add_parser("bench", help="Comparative benchmark of methods over a node")
113
+ bench.add_argument("--model", required=True, help="Path to the saved model (.pt)")
114
+ bench.add_argument("--data", required=True, help="Path to the saved Data (.pt)")
115
+ bench.add_argument(
116
+ "--node",
117
+ type=int,
118
+ default=None,
119
+ help="Node index (node-level); omit for graph-level",
120
+ )
121
+ bench.add_argument("--target-class", type=int, default=None)
122
+ bench.add_argument(
123
+ "--methods",
124
+ default="all",
125
+ help=f"Comma-separated methods (or 'all'). Aliases: {', '.join(_METHODS)}",
126
+ )
127
+ bench.add_argument("--backend", default="pyg", choices=["pyg", "dgl"])
128
+ bench.add_argument("--epochs", type=int, default=200)
129
+ bench.add_argument("--lr", type=float, default=None)
130
+ bench.add_argument("--top-k", type=int, default=5)
131
+ bench.add_argument("--num-perturbations", type=int, default=5)
132
+ bench.add_argument("--noise-std", type=float, default=0.05)
133
+ bench.add_argument("--threshold", type=float, default=0.5)
134
+ bench.add_argument("--seed", type=int, default=0)
135
+ bench.add_argument("--no-stability", action="store_true")
136
+ bench.add_argument("--json", default=None, help="Export results to .json")
137
+ bench.add_argument(
138
+ "--html", default=None, help="Export comparative report to .html"
139
+ )
140
+ return parser
141
+
142
+
143
+ def _instantiate(name: str, args: argparse.Namespace):
144
+ kw: dict[str, Any] = {}
145
+ for attr, param in (
146
+ ("epochs", "epochs"),
147
+ ("lr", "lr"),
148
+ ("mode", "mode"),
149
+ ("hops", "hops"),
150
+ ("max_steps", "max_steps"),
151
+ ("eps", "eps"),
152
+ ("steps", "steps"),
153
+ ):
154
+ val = getattr(args, attr)
155
+ if val is not None:
156
+ kw[param] = val
157
+ if "normalize" in dir(args) and args.normalize:
158
+ kw["normalize"] = True
159
+ return instantiate(name, **kw)
160
+
161
+
162
+ def _make_explainer(args: argparse.Namespace):
163
+ from graph_explain import Explainer
164
+
165
+ algorithm = _instantiate(args.method, args)
166
+ return Explainer(
167
+ algorithm=algorithm,
168
+ backend=args.backend,
169
+ mask_threshold=args.threshold,
170
+ )
171
+
172
+
173
+ def _fmt(value: Any) -> Any:
174
+ if hasattr(value, "tolist"):
175
+ return [round(float(v), 4) for v in value.reshape(-1).tolist()]
176
+ return value
177
+
178
+
179
+ def _json_safe(value: Any) -> Any:
180
+ if isinstance(value, dict):
181
+ return {str(k): _json_safe(v) for k, v in value.items()}
182
+ if isinstance(value, (list, tuple)):
183
+ return [_json_safe(v) for v in value]
184
+ import torch
185
+
186
+ if torch.is_tensor(value):
187
+ return _json_safe(value.tolist())
188
+ if isinstance(value, (int, float, str, bool)) or value is None:
189
+ return value
190
+ return str(value)
191
+
192
+
193
+ def _eval_metric(name: str, args, model, data, explanation) -> float | None:
194
+ from graph_explain.core.evaluation import (
195
+ evaluate_fidelity_minus,
196
+ evaluate_fidelity_plus,
197
+ evaluate_gea,
198
+ evaluate_stability,
199
+ )
200
+
201
+ try:
202
+ if name in ("fidelity", "fidelity_plus"):
203
+ return float(evaluate_fidelity_plus(model, explanation))
204
+ if name == "fidelity_minus":
205
+ return float(evaluate_fidelity_minus(model, explanation))
206
+ if name == "gea":
207
+ if getattr(model, "task_level", "node") == "graph":
208
+ from graph_explain.core.evaluation import evaluate_gea_graph
209
+
210
+ return float(
211
+ evaluate_gea_graph(explanation, data=data, top_k=args.top_k)
212
+ )
213
+ return float(evaluate_gea(explanation, data=data, top_k=args.top_k))
214
+ if name == "stability":
215
+ if args.node is None:
216
+ raise ValueError("stability requires --node")
217
+
218
+ def _again(d):
219
+ return _make_explainer(args).explain_node(d, model, args.node)
220
+
221
+ return float(
222
+ evaluate_stability(
223
+ _again,
224
+ data,
225
+ num_perturbations=args.num_perturbations,
226
+ noise_std=args.noise_std,
227
+ top_k=args.top_k,
228
+ )
229
+ )
230
+ except Exception as exc: # noqa: BLE001
231
+ print(f" * metric {name} unavailable: {exc}", file=sys.stderr)
232
+ return None
233
+ raise ValueError(f"unknown metric: {name}")
234
+
235
+
236
+ def _cmd_explain(args: argparse.Namespace) -> int:
237
+ import torch
238
+
239
+ from graph_explain.core.registry import get_algorithm
240
+
241
+ model = torch.load(args.model, map_location="cpu", weights_only=False)
242
+ data = torch.load(args.data, map_location="cpu", weights_only=False)
243
+ model.eval()
244
+
245
+ task = getattr(model, "task_level", "node")
246
+ if args.node is None:
247
+ if task == "graph":
248
+ if not get_algorithm(args.method).graph_level:
249
+ print(
250
+ f"Error: {args.method} does not support graph-level explanations "
251
+ "(node-level only).",
252
+ file=sys.stderr,
253
+ )
254
+ return 2
255
+ index = None
256
+ else:
257
+ print(
258
+ "To explain a node you must pass --node (per-node explanation).",
259
+ file=sys.stderr,
260
+ )
261
+ return 2
262
+ else:
263
+ index = args.node
264
+
265
+ algorithm = _instantiate(args.method, args)
266
+ explainer = _make_explainer(args)
267
+ try:
268
+ explanation = explainer.explain(
269
+ data,
270
+ model,
271
+ index=index,
272
+ target_class=args.target_class,
273
+ )
274
+ except ValueError as exc:
275
+ print(f"Error with {args.method}: {exc}", file=sys.stderr)
276
+ return 2
277
+ algorithm_class = get_algorithm(args.method)
278
+
279
+ print(f"Method: {args.method} ({algorithm_class.__name__})")
280
+ print(f"Original prediction: {_fmt(explanation.prediction_original)}")
281
+ if explanation.prediction_explanation is not None:
282
+ print(
283
+ f"Prediction after explanation: {_fmt(explanation.prediction_explanation)}"
284
+ )
285
+
286
+ metrics: dict[str, float | None] = {}
287
+ if args.metrics:
288
+ for name in args.metrics.split(","):
289
+ name = name.strip()
290
+ if not name:
291
+ continue
292
+ metrics[name] = _eval_metric(name, args, model, data, explanation)
293
+ print(f"Metrics: {metrics}")
294
+
295
+ if args.describe:
296
+ from graph_explain import describe
297
+
298
+ print(f"Narration: {describe(explanation, data=data, top_k=args.top_k)}")
299
+
300
+ if args.json:
301
+ from graph_explain import summarize
302
+
303
+ report = {
304
+ "version": __version__,
305
+ "method": algorithm.name,
306
+ "backend": args.backend,
307
+ "node": args.node,
308
+ "target_class": args.target_class,
309
+ "threshold": args.threshold,
310
+ "prediction_original": _fmt(explanation.prediction_original),
311
+ "prediction_explanation": _fmt(explanation.prediction_explanation),
312
+ "metrics": metrics,
313
+ "summary": summarize(explanation, data=data, top_k=args.top_k),
314
+ }
315
+ with open(args.json, "w", encoding="utf-8") as fh:
316
+ json.dump(_json_safe(report), fh, indent=2, ensure_ascii=False)
317
+ print(f"JSON report saved to {args.json}")
318
+
319
+ if args.output:
320
+ torch.save(explanation, args.output)
321
+ print(f"Explanation saved to {args.output}")
322
+ if args.plot:
323
+ from graph_explain.visualization import visualize_static
324
+
325
+ visualize_static(explanation, threshold=args.threshold)
326
+ import matplotlib.pyplot as plt
327
+
328
+ plt.savefig(args.plot, bbox_inches="tight")
329
+ print(f"Visualization saved to {args.plot}")
330
+ if args.html:
331
+ from graph_explain.visualization import visualize_interactive
332
+
333
+ visualize_interactive(
334
+ explanation, output_path=args.html, threshold=args.threshold
335
+ )
336
+ print(f"Interactive visualization saved to {args.html}")
337
+ return 0
338
+
339
+
340
+ def _cmd_bench(args: argparse.Namespace) -> int:
341
+ import torch
342
+
343
+ from graph_explain.core.benchmark import DEFAULT_METHODS, compare, report_html
344
+
345
+ model = torch.load(args.model, map_location="cpu", weights_only=False)
346
+ data = torch.load(args.data, map_location="cpu", weights_only=False)
347
+ model.eval()
348
+
349
+ task = getattr(model, "task_level", "node")
350
+ if args.node is None and task != "graph":
351
+ print("For node-level you must pass --node.", file=sys.stderr)
352
+ return 2
353
+
354
+ if args.methods.strip().lower() == "all":
355
+ methods = list(DEFAULT_METHODS)
356
+ else:
357
+ methods = []
358
+ for item in args.methods.split(","):
359
+ item = item.strip()
360
+ if not item:
361
+ continue
362
+ try:
363
+ methods.append(get_algorithm(item).name)
364
+ except ValueError as exc:
365
+ print(f"Error: {exc}", file=sys.stderr)
366
+ return 2
367
+
368
+ methods = list(dict.fromkeys(methods))
369
+
370
+ print(
371
+ f"Benchmark {'over node ' + str(args.node) if args.node is not None else 'graph-level'}"
372
+ f" ({len(methods)} methods) - backend {args.backend}\n"
373
+ )
374
+ results = compare(
375
+ data,
376
+ model,
377
+ node=args.node,
378
+ target_class=args.target_class,
379
+ backend=args.backend,
380
+ methods=methods,
381
+ top_k=args.top_k,
382
+ num_perturbations=args.num_perturbations,
383
+ noise_std=args.noise_std,
384
+ epochs=args.epochs,
385
+ lr=args.lr,
386
+ seed=args.seed,
387
+ mask_threshold=args.threshold,
388
+ stability=not args.no_stability,
389
+ )
390
+
391
+ headers = ("Method", "fid+", "fid-", "GEA", "sparsity", "stab")
392
+ widths = [len(h) for h in headers]
393
+ rows: list[tuple[Any, ...]] = []
394
+ for name, entry in results.items():
395
+ if name.startswith("_"):
396
+ continue
397
+ if entry["skipped"]:
398
+ rows.append((name, "not applicable", "", "", "", ""))
399
+ continue
400
+ m = entry["metrics"]
401
+ rows.append(
402
+ (
403
+ name,
404
+ "" if m["fidelity_plus"] is None else f"{m['fidelity_plus']:.3f}",
405
+ "" if m["fidelity_minus"] is None else f"{m['fidelity_minus']:.3f}",
406
+ "" if m["gea"] is None else f"{m['gea']:.3f}",
407
+ "" if m["sparsity"] is None else f"{m['sparsity']:.3f}",
408
+ "" if m["stability"] is None else f"{m['stability']:.3f}",
409
+ )
410
+ )
411
+
412
+ widths[0] = max(widths[0], max((len(r[0]) for r in rows), default=0))
413
+ for i in range(1, len(headers)):
414
+ widths[i] = max(widths[i], max((len(r[i]) for r in rows if r[i]), default=0))
415
+
416
+ for i, h in enumerate(headers):
417
+ widths[i] = max(widths[i], len(h))
418
+ line = " ".join(h.ljust(widths[i]) for i, h in enumerate(headers))
419
+ print(line)
420
+ print(" ".join("-" * w for w in widths))
421
+ for row in rows:
422
+ print(" ".join(str(c).ljust(widths[i]) for i, c in enumerate(row)))
423
+
424
+ skipped = results["_meta"]["skipped"]
425
+ if skipped:
426
+ print("\nSkipped:")
427
+ for name, reason in skipped.items():
428
+ print(f" {name}: {reason}")
429
+
430
+ if args.json:
431
+ out = {
432
+ name: entry for name, entry in results.items() if not name.startswith("_")
433
+ }
434
+ out["_meta"] = results["_meta"]
435
+ out["_meta"]["version"] = __version__
436
+ with open(args.json, "w", encoding="utf-8") as fh:
437
+ json.dump(_json_safe(out), fh, indent=2, ensure_ascii=False)
438
+ print(f"\nResults saved to JSON at {args.json}")
439
+ if args.html:
440
+ report_html(results, args.html)
441
+ print(f"HTML report saved to {args.html}")
442
+ return 0
443
+
444
+
445
+ def main(argv: list[str] | None = None) -> int:
446
+ parser = build_parser()
447
+ args = parser.parse_args(argv)
448
+ if args.command is None:
449
+ parser.print_help()
450
+ return 1
451
+ if args.command == "explain":
452
+ return _cmd_explain(args)
453
+ if args.command == "bench":
454
+ return _cmd_bench(args)
455
+ return 1
456
+
457
+
458
+ if __name__ == "__main__":
459
+ sys.exit(main())
@@ -0,0 +1,14 @@
1
+ from .benchmark import compare, report_html
2
+ from .explainer import Explainer
3
+ from .explanation import Explanation
4
+ from .registry import get_algorithm, instantiate, register
5
+
6
+ __all__ = [
7
+ "Explainer",
8
+ "Explanation",
9
+ "compare",
10
+ "get_algorithm",
11
+ "instantiate",
12
+ "register",
13
+ "report_html",
14
+ ]