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,73 @@
1
+ from __future__ import annotations
2
+
3
+ from ..core.explanation import Explanation
4
+
5
+
6
+ def visualize_interactive(
7
+ explanation: Explanation,
8
+ output_path: str = "explanation.html",
9
+ threshold: float | None = None,
10
+ show_labels: bool = True,
11
+ title: str | None = None,
12
+ node_size_scale: float = 120.0,
13
+ ) -> str:
14
+ try:
15
+ from pyvis.network import Network
16
+ except ImportError:
17
+ raise ImportError(
18
+ "visualize_interactive requiere pyvis: pip install pyvis (extras: [interactive])."
19
+ )
20
+
21
+ threshold = threshold if threshold is not None else explanation.mask_threshold
22
+ G = explanation.to_networkx(threshold=threshold)
23
+
24
+ net = Network(
25
+ height="700px",
26
+ width="100%",
27
+ directed=False,
28
+ notebook=False,
29
+ heading=(title or "Graph Explain"),
30
+ select_menu=True,
31
+ filter_menu=False,
32
+ )
33
+
34
+ edge_weights = {(u, v): float(w) for u, v, w in G.edges(data="weight", default=0.0)}
35
+ max_w = max(edge_weights.values(), default=0.0) or 1.0
36
+
37
+ node_imp = {}
38
+ if explanation.node_importance is not None:
39
+ for n in G.nodes():
40
+ node_imp[int(n)] = float(explanation.node_importance[int(n)])
41
+
42
+ for n in G.nodes():
43
+ label = str(n) if show_labels else ""
44
+ if explanation.node_idx is not None and n == int(explanation.node_idx):
45
+ color, size = "#d62728", 1.6 * node_size_scale
46
+ title_attr = f"<b>nodo objetivo {n}</b>"
47
+ else:
48
+ imp = node_imp.get(int(n), 0.0)
49
+ scale = 0.5 + 1.8 * imp
50
+ color = _importance_color(imp)
51
+ size = scale * node_size_scale
52
+ title_attr = f"nodo {n}<br>importancia: {imp:.3f}"
53
+ net.add_node(int(n), label=label, color=color, size=size, title=title_attr)
54
+
55
+ for (u, v), w in edge_weights.items():
56
+ width = 1 + 8 * (w / max_w)
57
+ color = f"rgba(200, 40, 40, {0.3 + 0.7 * w / max_w})"
58
+ net.add_edge(
59
+ int(u), int(v), value=w, width=width, color=color, title=f"peso: {w:.3f}"
60
+ )
61
+
62
+ net.write_html(output_path, open_browser=False, notebook=False)
63
+ return output_path
64
+
65
+
66
+ def _importance_color(value: float) -> str:
67
+ lo = "#dcedc1"
68
+ hi = "#d62728"
69
+ t = max(0.0, min(1.0, value))
70
+ r = int(float(int(lo[1:3], 16)) + t * (int(hi[1:3], 16) - int(lo[1:3], 16)))
71
+ g = int(float(int(lo[3:5], 16)) + t * (int(hi[3:5], 16) - int(lo[3:5], 16)))
72
+ b = int(float(int(lo[5:7], 16)) + t * (int(hi[5:7], 16) - int(lo[5:7], 16)))
73
+ return f"#{(r << 16) | (g << 8) | b:06x}"
@@ -0,0 +1,90 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import matplotlib.pyplot as plt
6
+ import networkx as nx
7
+
8
+ from ..core.explanation import Explanation
9
+
10
+
11
+ def visualize_static(
12
+ explanation: Explanation,
13
+ threshold: float | None = None,
14
+ show_labels: bool = False,
15
+ node_size: int = 400,
16
+ title: str | None = None,
17
+ ax: Any | None = None,
18
+ seed: int = 42,
19
+ cmap: str = "YlOrRd",
20
+ ) -> Any:
21
+ threshold = threshold if threshold is not None else explanation.mask_threshold
22
+ G = explanation.to_networkx(threshold=threshold)
23
+ if explanation.node_idx is not None and explanation.node_idx in G.nodes():
24
+ target = int(explanation.node_idx)
25
+ else:
26
+ target = None
27
+
28
+ if ax is None:
29
+ _, ax = plt.subplots(figsize=(8, 6))
30
+
31
+ pos = nx.spring_layout(G, seed=seed)
32
+ edge_weights = {}
33
+ for u, v, w in G.edges(data="weight", default=0.0):
34
+ edge_weights[(u, v)] = float(w)
35
+ nx.draw_networkx_edges(G, pos, ax=ax, edge_color="#888", alpha=0.6)
36
+
37
+ if target is not None:
38
+ others = [n for n in G.nodes() if n != target]
39
+ if others:
40
+ nx.draw_networkx_nodes(
41
+ G,
42
+ pos,
43
+ nodelist=others,
44
+ node_size=node_size,
45
+ node_color="#aaddff",
46
+ ax=ax,
47
+ node_shape="o",
48
+ )
49
+ nx.draw_networkx_nodes(
50
+ G,
51
+ pos,
52
+ nodelist=[target],
53
+ node_size=node_size * 1.4,
54
+ node_color="#d62728",
55
+ ax=ax,
56
+ )
57
+ else:
58
+ nx.draw_networkx_nodes(G, pos, node_size=node_size, node_color="#aaddff", ax=ax)
59
+
60
+ if explanation.edge_importance is not None:
61
+ vals = list(edge_weights.values())
62
+ if vals:
63
+ vmin, vmax = min(vals), max(vals)
64
+ span = vmax - vmin or 1.0
65
+ cmap_obj = plt.colormaps[cmap]
66
+ for (u, v), w in edge_weights.items():
67
+ t = (w - vmin) / span
68
+ width = 0.5 + 4.0 * t
69
+ nx.draw_networkx_edges(
70
+ G,
71
+ pos,
72
+ edgelist=[(u, v)],
73
+ width=width,
74
+ edge_color=cmap_obj(t),
75
+ ax=ax,
76
+ alpha=0.9,
77
+ )
78
+
79
+ if show_labels:
80
+ nx.draw_networkx_labels(G, pos, ax=ax)
81
+
82
+ ax.set_axis_off()
83
+ if title:
84
+ ax.set_title(title)
85
+ return ax
86
+
87
+
88
+ def show(explanation: Explanation, **kwargs) -> None:
89
+ visualize_static(explanation, **kwargs)
90
+ plt.show()
@@ -0,0 +1,332 @@
1
+ Metadata-Version: 2.4
2
+ Name: graph-explain
3
+ Version: 0.7.0
4
+ Summary: Explainability library for graph-based models (GNN)
5
+ Author: graph-explain contributors
6
+ License-Expression: MIT
7
+ Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
8
+ Classifier: Development Status :: 4 - Beta
9
+ Classifier: Intended Audience :: Science/Research
10
+ Classifier: Operating System :: OS Independent
11
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
12
+ Classifier: Programming Language :: Python :: 3
13
+ Classifier: Programming Language :: Python :: 3 :: Only
14
+ Classifier: Programming Language :: Python :: 3.10
15
+ Classifier: Programming Language :: Python :: 3.11
16
+ Classifier: Programming Language :: Python :: 3.12
17
+ Requires-Python: >=3.10
18
+ Description-Content-Type: text/markdown
19
+ Requires-Dist: numpy>=1.24
20
+ Requires-Dist: networkx>=3.0
21
+ Requires-Dist: matplotlib>=3.6
22
+ Provides-Extra: pyg
23
+ Requires-Dist: torch>=2.0; extra == "pyg"
24
+ Requires-Dist: torch-geometric>=2.5; extra == "pyg"
25
+ Provides-Extra: dgl
26
+ Requires-Dist: torch>=2.0; extra == "dgl"
27
+ Requires-Dist: dgl>=2.0; extra == "dgl"
28
+ Provides-Extra: interactive
29
+ Requires-Dist: plotly>=5.15; extra == "interactive"
30
+ Requires-Dist: pyvis>=0.3; extra == "interactive"
31
+ Provides-Extra: all
32
+ Requires-Dist: torch>=2.0; extra == "all"
33
+ Requires-Dist: torch-geometric>=2.5; extra == "all"
34
+ Requires-Dist: plotly>=5.15; extra == "all"
35
+ Requires-Dist: pyvis>=0.3; extra == "all"
36
+ Provides-Extra: dev
37
+ Requires-Dist: pytest>=7.0; extra == "dev"
38
+ Requires-Dist: ruff>=0.5; extra == "dev"
39
+ Requires-Dist: build>=1.0; extra == "dev"
40
+ Provides-Extra: docs
41
+ Requires-Dist: sphinx>=7.2; extra == "docs"
42
+ Requires-Dist: sphinx-rtd-theme>=2.0; extra == "docs"
43
+ Provides-Extra: publish
44
+ Requires-Dist: twine>=5.0; extra == "publish"
45
+
46
+ # graph-explain
47
+
48
+ Explainability library for graph-based models (Graph Neural Networks).
49
+ Explains a GNN's predictions in terms of **important nodes, edges and subgraphs**,
50
+ with built-in metrics and visualization.
51
+
52
+ ## Features
53
+
54
+ - **Unified API**: a single `Explainer` object for every method.
55
+ - **Node-level and graph-level**: `explain_node(...)` explains a node's
56
+ prediction; `explain_graph(...)` (or CLI without `--node`) explains a whole
57
+ graph with graph-level models (`task_level = "graph"`), including GEA
58
+ graph-level metrics and comparative benchmarking.
59
+ - **Explanation methods**:
60
+ - `GNNExplainer` — soft masks over nodes/edges (perturbation).
61
+ - `PGExplainer` — MLP generating edge masks (inductive, fast at inference).
62
+ - `SubgraphX` — MCTS search for subgraphs that maximize the prediction (high fidelity).
63
+ - `Saliency` — gradient-based importance.
64
+ - `Integrated Gradients` — gradient accumulation vs. a baseline (attribution paths).
65
+ - `GNNGatedLRP` — layer-wise relevance propagation (LRP-0/z+) over GCNs;
66
+ distributes relevance between nodes and edges from the positive contributions
67
+ of each conv/linear layer; supports `GCNConv` + `ReLU` + `Linear`.
68
+ - `DeepLift` — additive rescale rule vs. a (zero) baseline: each feature gets a
69
+ contribution proportional to its effect on the target class; conservative
70
+ (contributions sum ≈ Δ logits); supports `GCNConv` + `ReLU` + `Linear`.
71
+ - `AttentionExplainer` — node/edge importance from a `GATConv` model's
72
+ attention weights (softmax per neighbor, averaged over heads and layers).
73
+ - `GradXInput` — gradient × activation (zero baseline) for nodes and edges.
74
+ - `GraphLIME` — local linear (ridge) regression over the k-hop neighbors'
75
+ features, weighted by similarity to the target node; gives directly
76
+ interpretable feature importance without training.
77
+ - `NodeMask` — node mask learned by optimization (tracking the prediction)
78
+ over the k-hop subgraph, regularized toward sparsity.
79
+ - `GuidedBackprop` — gradients guided by the ReLU mask (positive activations
80
+ only); falls back to standard gradients if the model uses functional ReLUs.
81
+ - `Random` — uniformly random importance baseline (seed-able) for benchmarks.
82
+ - `Counterfactual` — minimal perturbation (edges or features) that changes a
83
+ node's prediction (deterministic greedy search); returns the modified
84
+ elements as importance plus the logits after the change.
85
+ - **Narration**: `describe(expl)` builds a template-based natural-language
86
+ explanation (Spanish by default), and `narrate(expl, llm=...)` lets you plug
87
+ in a generative model (a `prompt -> text` callable) for free-form text.
88
+ - **Metrics**:
89
+ - `evaluate_sparsity` — global or local sparsity (`local=True`, over the k-hop subgraph).
90
+ - `evaluate_fidelity_plus` — **necessity**: drop in `P(c)` when removing the top-k elements.
91
+ - `evaluate_fidelity_minus` — **sufficiency**: `P(c)` preserved when keeping ONLY the top-k.
92
+ - `evaluate_stability` — mean similarity between explanations under feature/edge perturbations.
93
+ - `evaluate_gea` — **Graph Explanation Accuracy**: overlap of the top-k with the ground-truth subgraph (BA-Shapes).
94
+ - **Built-in benchmarks**: BA-Shapes synthetic generator with ground truth and
95
+ `ground_truth_nodes` / `ground_truth_edge_ids` helpers; in addition,
96
+ `build_graph_classification` builds a **graph classification** dataset (house
97
+ motif) with per-graph `gt_edge_mask` for graph-level GEA
98
+ (`evaluate_gea_graph`).
99
+ - **Visualization**: static (matplotlib + networkx) and interactive (pyvis → HTML).
100
+ - **Backends**: PyTorch Geometric and DGL (through an adapter; DGL requires a
101
+ PyTorch version with pre-built graphbolt libraries).
102
+ - **CLI** to explain saved models without writing code, plus a **comparative
103
+ benchmark** of all methods over a node (table, JSON and HTML).
104
+ - **Programmatic comparison**: `compare(...)` to evaluate and compare methods.
105
+
106
+ ## Installation
107
+
108
+ ```bash
109
+ python -m venv .venv && source .venv/bin/activate
110
+ pip install -e .[all]
111
+ ```
112
+
113
+ Optional extras: `pyg` (PyTorch Geometric), `dgl` (DGL backend),
114
+ `interactive` (plotly/pyvis).
115
+
116
+ ## Quick start
117
+
118
+ ```python
119
+ from graph_explain import Explainer, GNNExplainer, Saliency
120
+ from graph_explain.benchmarks.synthetic import build_data
121
+ from graph_explain.visualization import show
122
+
123
+ data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
124
+ model = GCN(in_channels=data.x.size(1)) # your trained GNN
125
+ model.eval()
126
+
127
+ explainer = Explainer(algorithm=GNNExplainer(epochs=150))
128
+ expl = explainer.explain_node(data, model, node_idx=42)
129
+
130
+ print(expl.evaluate(metrics=["fidelity", "sparsity"]))
131
+ print(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
132
+ show(expl, show_labels=True) # highlight the explanatory subgraph
133
+ ```
134
+
135
+ ## Sparsity tuning notes
136
+
137
+ - **Structural models**: perturbation-based explainers (GNNExplainer,
138
+ PGExplainer, SubgraphX) assume the prediction depends on the neighborhood
139
+ structure. A `GCNConv` with `add_self_loops=True` and `bias=True` can predict
140
+ the class from biases/self-loops alone; in that case edge masks collapse to
141
+ zero because edges do not matter. For meaningful demos use
142
+ `GCNConv(..., add_self_loops=False, bias=False)` (see `examples/model.py`).
143
+ - **Benchmark split**: `build_data` splits train/test across **all** nodes
144
+ (including motifs). If the model is trained on class 0 only, it learns to
145
+ ignore structure.
146
+ - **`PGExplainer(temp=...)`**: with `temp=5` the Gumbel-sigmoid sampling
147
+ gradient flattens (~0.05) and the mask collapses to zero. The default is `temp=1.0`.
148
+ - **Local sparsity**: `evaluate_sparsity(expl, local=True)` measures sparsity
149
+ over the explained node's `k-hop` subgraph instead of the whole graph; when
150
+ the mask is counted over the full graph, local explanations get diluted
151
+ (sparsity appears near 1).
152
+
153
+ ## CLI
154
+
155
+ ```bash
156
+ # Save model and data first:
157
+ torch.save(model, "model.pt"); torch.save(data, "data.pt")
158
+
159
+ graph-explain explain \
160
+ --model model.pt --data data.pt \
161
+ --method gnn_explainer --node 42 \
162
+ --plot explicacion.png
163
+ ```
164
+
165
+ ## The `Explanation` object
166
+
167
+ - `node_importance`: importance per node `(num_nodes,)`.
168
+ - `edge_importance`: importance per edge.
169
+ - `feature_importance`: importance per feature (method-dependent).
170
+ - `prediction_original` / `prediction_explanation`: logits for fidelity evaluation.
171
+ - Methods: `evaluate(metrics=[...])`, `to_networkx(threshold=...)`, `visualize_static(...)`.
172
+
173
+ ## Structure
174
+
175
+ ```
176
+ src/graph_explain/
177
+ ├── core/ # Explainer, Explanation, registry, evaluation
178
+ ├── methods/ # gnn_explainer, subgraphx, pg_explainer, saliency, integrated_gradients
179
+ ├── backends/ # Backend API + PyGAdapter + DGLAdapter
180
+ ├── benchmarks/ # BA-Shapes synthetic generator + ground-truth helpers
181
+ └── visualization/ # static plots
182
+
183
+ ```
184
+
185
+ `get_backend(name)` returns `PyGAdapter` or `DGLAdapter`. For DGL, features go
186
+ in `ndata['feat']`, labels in `ndata['label']` and edge weights in `edata['w']`;
187
+ the model must read `g.ndata['feat']` and `g.edata['w']`.
188
+
189
+ **DGL validation against the real library**: DGL 2.1.0 only ships graphbolt C++
190
+ libraries for torch ≤ 2.2.1, so the real integration is tested in an isolated
191
+ virtual machine (`tests/test_dgl_integration.py`, skipped when dgl is not
192
+ available):
193
+
194
+ ```bash
195
+ python3.12 -m venv /tmp/dgl-venv
196
+ /tmp/dgl-venv/bin/pip install torch==2.2.1 --index-url https://download.pytorch.org/whl/cpu \
197
+ dgl==2.1.0 "numpy<2" "scipy<1.14" "pandas" "torchdata==0.7.1" \
198
+ "torch-geometric==2.6.1" setuptools packaging
199
+ cd graph-explain && PYTHONPATH=. /tmp/dgl-venv/bin/python -m pytest tests -q
200
+ ```
201
+
202
+ ## Metrics (phase 3)
203
+
204
+ ````python
205
+ from graph_explain.core.evaluation import (
206
+ evaluate_fidelity_plus, evaluate_fidelity_minus,
207
+ evaluate_stability, evaluate_gea,
208
+ )
209
+
210
+ fp = evaluate_fidelity_plus(model, expl) # necessity: remove top-k elements → P(c) drops
211
+ fm = evaluate_fidelity_minus(model, expl) # sufficiency: keep only top-k → P(c) is preserved
212
+ stab = evaluate_stability(
213
+ lambda d: Explainer(algorithm=GNNExplainer(epochs=40)).explain_node(d, model, node_idx=42),
214
+ data, num_perturbations=5, noise_std=0.02,
215
+ )
216
+ gea = evaluate_gea(expl, data=data) # overlap with the BA-Shapes motif
217
+ ```
218
+
219
+ Example in `examples/example.py`, benchmark with `num_houses=30`: GNNExplainer →
220
+ `fid+ 0.74 / fid- 0.99 / GEA 0.92 / stab 0.85`.
221
+ ````
222
+
223
+ ## CLI (phase 5)
224
+
225
+ The command-line interface covers all methods (including the `lrp`/`gnn_lrp`
226
+ and `cf`/`counterfactual` aliases), metrics, narration and JSON reports:
227
+
228
+ ```bash
229
+ graph-explain --version
230
+
231
+ # Counterfactual explanation for node 42 + narration + JSON report
232
+ graph-explain explain --model model.pt --data data.pt \
233
+ --method counterfactual --node 42 --mode feature \
234
+ --hops 2 --max-steps 10 --describe --json report.json
235
+
236
+ # Normalized GNN-LRP with metrics
237
+ graph-explain explain --model model.pt --data data.pt \
238
+ --method lrp --node 42 --normalize \
239
+ --metrics fidelity_plus,fidelity_minus,gea,stability \
240
+ --top-k 5 --num-perturbations 5
241
+
242
+ # GNNExplainer + static and interactive visualizations
243
+ graph-explain explain --model model.pt --data data.pt \
244
+ --method gnn_explainer --node 42 --epochs 200 \
245
+ --threshold 0.5 --plot expl.png --html expl.html
246
+ ```
247
+
248
+ Main options: `--method`, `--node`, `--target-class`, `--epochs`, `--lr`,
249
+ `--mode` (edge/feature), `--hops`, `--max-steps`, `--eps`, `--steps`,
250
+ `--normalize`, `--backend` (pyg/dgl), `--threshold`, `--top-k`, `--metrics`,
251
+ `--num-perturbations`, `--noise-std`, `--describe`, `--json`, `--output`,
252
+ `--plot`, `--html`. The JSON report includes method, predictions, metrics and
253
+ the structured summary (`summarize`) with top-k nodes/edges.
254
+
255
+ ## Comparative benchmark (phase 7)
256
+
257
+ `compare(data, model, node=...)` runs every method on a node, computes the
258
+ metric battery (fid+ / fid- / GEA / sparsity / stability) and returns a
259
+ structured dict; non-applicable methods (e.g. Attention without `GATConv`) and
260
+ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
261
+
262
+ ```python
263
+ from graph_explain import compare, report_html
264
+
265
+ results = compare(data, model, node=42, methods=None, # None = all
266
+ num_perturbations=5, epochs=200)
267
+ report_html(results, "bench.html") # self-contained HTML report
268
+ ```
269
+
270
+ The CLI ships an equivalent subcommand:
271
+
272
+ ```bash
273
+ graph-explain bench --model model.pt --data data.pt --node 42 \
274
+ --methods all --num-perturbations 5 \
275
+ --json bench.json --html bench.html
276
+ ```
277
+
278
+ Note: `gea` is only defined when the node belongs to a ground-truth subgraph of
279
+ the benchmark (BA-Shapes); otherwise it shows up empty in the table.
280
+
281
+ ## Graph-level (phase 10)
282
+
283
+ Models that predict over whole graphs (`task_level = "graph"`, e.g. GCN +
284
+ global pooling). Without `--node`, the CLI explains the whole graph; methods
285
+ marked with `graph_level`:
286
+
287
+ ```bash
288
+ # Explain a whole graph (graph-level model) + GEA over the motif
289
+ graph-explain explain --model model.pt --data graph.pt \
290
+ --method grad_x_input --metrics fidelity_plus,gea
291
+
292
+ # Graph-level bench (shows skipped methods and only runs applicable ones)
293
+ graph-explain bench --model model.pt --data graph.pt \
294
+ --methods all --no-stability --json bench_graph.json
295
+ ```
296
+
297
+ In Python:
298
+
299
+ ```python
300
+ from graph_explain import Explainer, evaluate_gea_graph
301
+ from graph_explain.benchmarks.synthetic import build_graph_classification
302
+
303
+ graphs = build_graph_classification(num_pos=8, num_neg=8, seed=0) # binary y, gt_edge_mask
304
+ model = ... # GraphGCN (task_level="graph")
305
+
306
+ expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
307
+ print(evaluate_gea_graph(expl, data=graphs[0], top_k=13))
308
+ ```
309
+
310
+ Node-only methods (`GraphLIME`, `NodeMask`, `Attention`, `GNNGatedLRP`,
311
+ `Counterfactual`, `DeepLift`, `PGExplainer`, `SubgraphX`) are marked as
312
+ `skipped` at graph-level.
313
+
314
+ ## Roadmap
315
+
316
+ - [x] Phase 2: PGExplainer, SubgraphX, Integrated Gradients
317
+ - [x] Phase 2: interactive visualization (pyvis → HTML)
318
+ - [x] Phase 3: full metrics (fidelity±, stability, GEA)
319
+ - [x] Phase 3: DGL backend (adapter; integration validated with DGL 2.1 + torch 2.2.1)
320
+ - [x] Phase 4: GNN-LRP (layer-wise relevance for GCNs; validates the house motif in BA-Shapes)
321
+ - [x] Phase 4: counterfactual explanations (minimal edge/feature removal that changes the class)
322
+ - [x] Phase 4: LLM narration (`describe` deterministic + pluggable `narrate` LLM)
323
+ - [x] Phase 5: full CLI (all methods, metrics, narration and JSON export)
324
+ - [x] Phase 6: more methods (DeepLIFT rescale, Attention/GAT, Gradient×Input)
325
+ - [x] Phase 7: comparative benchmark (`compare` + CLI `bench` subcommand, table and JSON/HTML reports)
326
+ - [x] Phase 8: more methods (GraphLIME, NodeMask, GuidedBackprop and Random baseline)
327
+ - [x] Phase 10: graph-level explanations (graph-classification dataset with house
328
+ motif, graph-level GEA, CLI/bench without `--node` and `graph_level` flag)
329
+
330
+ ## License
331
+
332
+ MIT
@@ -0,0 +1,42 @@
1
+ graph_explain/__init__.py,sha256=VLwa0lllOqsat3D6ToBa1yZbtU6deXn58r1tzWiBlb0,1582
2
+ graph_explain/cli.py,sha256=45pvpzSkulc9POJ4XrmHju4xsyksSGArsellln3OXQs,15363
3
+ graph_explain/backends/__init__.py,sha256=2A0h99lJfueSaQRxclPg1u28WckJp2h4cN1K9gkDCmY,145
4
+ graph_explain/backends/base.py,sha256=_wUDTS9CpKnz0hn0dXjDnSRWgPBloe0_WDr3kVErpDQ,3126
5
+ graph_explain/backends/dgl.py,sha256=AfYkSEKIPKhY1fracyAuh5sZfYW9gXRIG0u3GZiBO4k,3965
6
+ graph_explain/benchmarks/__init__.py,sha256=ZP9fQc35CDN_YVKo4qiyUQ9RRtzSjKWNyAWcWbUM9ik,84
7
+ graph_explain/benchmarks/synthetic.py,sha256=TsfIyNWaBmGHBTWVXI7p2wapQvWnnf8hC8ZHbQe6EBM,8033
8
+ graph_explain/core/__init__.py,sha256=GfeO_Zfd8HDeyCjZtQkmoVL-V1p4qneIgSv5-YBpkQc,314
9
+ graph_explain/core/benchmark.py,sha256=AxCwkLJYgzxdEKbrZsOvYVpET33z9Dd19yK6RUVzUx0,8457
10
+ graph_explain/core/evaluation.py,sha256=XAiHDtJvYYlYCgPLDqeAKjRVRpqfJElkXJDXjJopuw0,13917
11
+ graph_explain/core/explainer.py,sha256=gPo_OiFuurAW_uuYpX_jO6qcQLk8Ze-0LGvXzg3ASi0,2421
12
+ graph_explain/core/explanation.py,sha256=E0Op2GodQ6abuzkfDxt0x2ZOwS_ThFbabectFo29dEM,2870
13
+ graph_explain/core/model_utils.py,sha256=M7KQbasF8yDcLcHAr9zg0N7VE8ODdaVI8P3RbxfdGuk,1273
14
+ graph_explain/core/registry.py,sha256=pvUw5naWqGOHnvwjys2SwYQGU23bCfGeDJCpNk17qww,1493
15
+ graph_explain/methods/__init__.py,sha256=NhkWO3e3pnWnprfU5-ii9BEEfKDVLK_Tgr_CNmuCK3A,1157
16
+ graph_explain/methods/base.py,sha256=pvY6pluaLPEg9BUwyEpvOBmESB2np_Skf9xjqk06Pl0,513
17
+ graph_explain/methods/attention/attention.py,sha256=_GWmLXxvgm_azJhNbrreKRtZjpM6jPRWe1qHbnm_pBw,5446
18
+ graph_explain/methods/baseline/random_baseline.py,sha256=IEhDuuogli57u6X-zWzIutIrq4v8J-h_IdMJyOGF0ZY,2485
19
+ graph_explain/methods/counterfactual/counterfactual.py,sha256=Mp8TT_sm1aru1yWuXlXqIvjHZ6y036dSBKfyvT_6RXE,10186
20
+ graph_explain/methods/feature/graph_lime.py,sha256=lgRXE1k78pzC-o6_0yUQBrdhc9bH-xBvrMvf9ZTSATg,4717
21
+ graph_explain/methods/gradient/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
22
+ graph_explain/methods/gradient/grad_x_input.py,sha256=lnvJu9TZ7NFaJOMFZMqWnaROilCpF-CA1YIAlI-CheY,3778
23
+ graph_explain/methods/gradient/guided_backprop.py,sha256=IdFm2ErD7VwIqaJf25HAeGCkmEzyojuOYFXjRkylO_4,4151
24
+ graph_explain/methods/gradient/integrated_gradients.py,sha256=MEZ1DicX59vNs28ZNYLor2VDBocL5gvRByTGbpLqKRU,4092
25
+ graph_explain/methods/gradient/saliency.py,sha256=USQjyTkgfIpc_2NDRgdubXuhp-g3goPAb8Bo1QVPAuE,2866
26
+ graph_explain/methods/perturbation/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
27
+ graph_explain/methods/perturbation/gnn_explainer.py,sha256=FZWpgHKkxIWCZwZkWknRrVwTy3nRfx_D_oLMd5w_42w,9253
28
+ graph_explain/methods/perturbation/node_mask.py,sha256=hvbmhq-EJMpmHbjVqeLIgOWJY6JP-4lHXUi1wGI1AQw,4552
29
+ graph_explain/methods/perturbation/pg_explainer.py,sha256=Eceiih7nwi4x5-c44NB9IpCmJCpuPXm6f69sMdJAtmU,5475
30
+ graph_explain/methods/perturbation/subgraphx.py,sha256=injkUmysN_mlwIkVrGt5FYdvtT6zCa21ZmfjWWeJY0E,13889
31
+ graph_explain/methods/relevance/deeplift.py,sha256=fAkdAcLKfwB8KYWgkO8N7iykrKKzKzzSTLp7sCCfvnU,8941
32
+ graph_explain/methods/relevance/gnn_lrp.py,sha256=HWBvXgywAsvvm1YD2YBrD2hKJw6Bh_qprwqI6xCHR3E,7793
33
+ graph_explain/narration/__init__.py,sha256=mI4DdNcZX2MU1KfZ4WvQwfDNSOZpxTxFc2Pf1iV-fzI,121
34
+ graph_explain/narration/narrator.py,sha256=LObk7dhDlDdGd8KD-KWL6UNt6qEjWWfXTQwQuGVNMxA,6371
35
+ graph_explain/visualization/__init__.py,sha256=DpBZQFdMt92FI4zZKknoq96kA_2jwy6Xg13BT2R-C-0,155
36
+ graph_explain/visualization/interactive.py,sha256=1TEluuQ82FwuYXj56613gyojLLTagXLS_X_q0GK_eDI,2551
37
+ graph_explain/visualization/static.py,sha256=jyfiWzDi63MnXqXpw-MOYPqqaCdOGcCH_ZVEJmJdIGs,2561
38
+ graph_explain-0.7.0.dist-info/METADATA,sha256=WlnXxLQEgaMivvIPebwZg5p1bjHC9XjN6HlBVXAF38Q,14785
39
+ graph_explain-0.7.0.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
40
+ graph_explain-0.7.0.dist-info/entry_points.txt,sha256=MgZ6bnSIGBlcerCsKPVLHikws4RL8p3tPGCCZQdMXHs,57
41
+ graph_explain-0.7.0.dist-info/top_level.txt,sha256=HJLm-M7AXPfhgNC9_dTk2abzDpkSvrb45FIggoRb9lA,14
42
+ graph_explain-0.7.0.dist-info/RECORD,,
@@ -0,0 +1,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (84.0.0)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -0,0 +1,2 @@
1
+ [console_scripts]
2
+ graph-explain = graph_explain.cli:main
@@ -0,0 +1 @@
1
+ graph_explain