graph-explain 0.7.2__tar.gz → 0.8.0__tar.gz
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-0.7.2 → graph_explain-0.8.0}/LICENSE +1 -1
- {graph_explain-0.7.2/src/graph_explain.egg-info → graph_explain-0.8.0}/PKG-INFO +35 -23
- {graph_explain-0.7.2 → graph_explain-0.8.0}/README.md +33 -22
- {graph_explain-0.7.2 → graph_explain-0.8.0}/pyproject.toml +33 -4
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/__init__.py +1 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/benchmarks/synthetic.py +8 -7
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/benchmark.py +4 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/evaluation.py +1 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/explanation.py +1 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/registry.py +6 -4
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/attention/attention.py +1 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/base.py +1 -3
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/grad_x_input.py +1 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/guided_backprop.py +2 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/integrated_gradients.py +2 -2
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/saliency.py +6 -2
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/gnn_explainer.py +11 -4
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/pg_explainer.py +1 -1
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/subgraphx.py +2 -2
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/deeplift.py +3 -3
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/gnn_lrp.py +3 -3
- {graph_explain-0.7.2 → graph_explain-0.8.0/src/graph_explain.egg-info}/PKG-INFO +35 -23
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/requires.txt +1 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/setup.cfg +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/backends/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/backends/base.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/backends/dgl.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/benchmarks/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/cli.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/explainer.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/model_utils.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/baseline/random_baseline.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/counterfactual/counterfactual.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/feature/graph_lime.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/node_mask.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/narration/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/narration/narrator.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/py.typed +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/visualization/__init__.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/visualization/interactive.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/visualization/static.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/SOURCES.txt +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/dependency_links.txt +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/entry_points.txt +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/top_level.txt +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_cli.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_core.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_counterfactual.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_dgl_integration.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_narration.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase10.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase2.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase3.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase4.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase6.py +0 -0
- {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase8.py +0 -0
|
@@ -18,4 +18,4 @@ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
|
18
18
|
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
19
|
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
20
|
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
-
SOFTWARE.
|
|
21
|
+
SOFTWARE.
|
|
@@ -1,12 +1,12 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: graph-explain
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.8.0
|
|
4
4
|
Summary: Explainability library for graph-based models (GNN)
|
|
5
5
|
Author: graph-explain contributors
|
|
6
6
|
License-Expression: MIT
|
|
7
7
|
Project-URL: Homepage, https://github.com/Tzinny-dev/graph-explain
|
|
8
8
|
Project-URL: Repository, https://github.com/Tzinny-dev/graph-explain
|
|
9
|
-
Project-URL: Documentation, https://github.
|
|
9
|
+
Project-URL: Documentation, https://tzinny-dev.github.io/graph-explain/
|
|
10
10
|
Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
|
|
11
11
|
Classifier: Development Status :: 4 - Beta
|
|
12
12
|
Classifier: Intended Audience :: Science/Research
|
|
@@ -41,6 +41,7 @@ Requires-Dist: plotly>=5.15; extra == "all"
|
|
|
41
41
|
Requires-Dist: pyvis>=0.3; extra == "all"
|
|
42
42
|
Provides-Extra: dev
|
|
43
43
|
Requires-Dist: pytest>=7.0; extra == "dev"
|
|
44
|
+
Requires-Dist: pytest-cov>=5.0; extra == "dev"
|
|
44
45
|
Requires-Dist: ruff>=0.5; extra == "dev"
|
|
45
46
|
Requires-Dist: build>=1.0; extra == "dev"
|
|
46
47
|
Provides-Extra: docs
|
|
@@ -52,6 +53,14 @@ Dynamic: license-file
|
|
|
52
53
|
|
|
53
54
|
# graph-explain
|
|
54
55
|
|
|
56
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
57
|
+
[](https://tzinny-dev.github.io/graph-explain/)
|
|
58
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
59
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
60
|
+
[](https://pypi.org/project/graph-explain/)
|
|
61
|
+
[](https://pypi.org/project/graph-explain/)
|
|
62
|
+
[](https://github.com/Tzinny-dev/graph-explain/blob/main/LICENSE)
|
|
63
|
+
|
|
55
64
|
Explainability library for graph-based models (Graph Neural Networks).
|
|
56
65
|
Explains a GNN's predictions in terms of **important nodes, edges and subgraphs**,
|
|
57
66
|
with built-in metrics and visualization.
|
|
@@ -120,6 +129,9 @@ pip install -e .[all]
|
|
|
120
129
|
Optional extras: `pyg` (PyTorch Geometric), `dgl` (DGL backend),
|
|
121
130
|
`interactive` (plotly/pyvis).
|
|
122
131
|
|
|
132
|
+
> 📚 **Documentation**: full API reference and guides at
|
|
133
|
+
> https://tzinny-dev.github.io/graph-explain/
|
|
134
|
+
|
|
123
135
|
## Quick start
|
|
124
136
|
|
|
125
137
|
```python
|
|
@@ -127,16 +139,18 @@ from graph_explain import Explainer, GNNExplainer, Saliency
|
|
|
127
139
|
from graph_explain.benchmarks.synthetic import build_data
|
|
128
140
|
from graph_explain.visualization import show
|
|
129
141
|
|
|
130
|
-
data = build_data(base_nodes=300, num_houses=80)
|
|
131
|
-
model = GCN(in_channels=data.x.size(1))
|
|
142
|
+
data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
|
|
143
|
+
model = GCN(in_channels=data.x.size(1)) # your trained GNN
|
|
132
144
|
model.eval()
|
|
133
145
|
|
|
134
146
|
explainer = Explainer(algorithm=GNNExplainer(epochs=150))
|
|
135
147
|
expl = explainer.explain_node(data, model, node_idx=42)
|
|
136
148
|
|
|
137
149
|
print(expl.evaluate(metrics=["fidelity", "sparsity"]))
|
|
138
|
-
print(
|
|
139
|
-
|
|
150
|
+
print(
|
|
151
|
+
expl.evaluate(metrics=["sparsity"], local=True)
|
|
152
|
+
) # sparsity over the node's k-hop subgraph
|
|
153
|
+
show(expl, show_labels=True) # highlight the explanatory subgraph
|
|
140
154
|
```
|
|
141
155
|
|
|
142
156
|
## Sparsity tuning notes
|
|
@@ -269,9 +283,15 @@ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
|
|
|
269
283
|
```python
|
|
270
284
|
from graph_explain import compare, report_html
|
|
271
285
|
|
|
272
|
-
results = compare(
|
|
273
|
-
|
|
274
|
-
|
|
286
|
+
results = compare(
|
|
287
|
+
data,
|
|
288
|
+
model,
|
|
289
|
+
node=42,
|
|
290
|
+
methods=None, # None = all
|
|
291
|
+
num_perturbations=5,
|
|
292
|
+
epochs=200,
|
|
293
|
+
)
|
|
294
|
+
report_html(results, "bench.html") # self-contained HTML report
|
|
275
295
|
```
|
|
276
296
|
|
|
277
297
|
The CLI ships an equivalent subcommand:
|
|
@@ -307,7 +327,9 @@ In Python:
|
|
|
307
327
|
from graph_explain import Explainer, evaluate_gea_graph
|
|
308
328
|
from graph_explain.benchmarks.synthetic import build_graph_classification
|
|
309
329
|
|
|
310
|
-
graphs = build_graph_classification(
|
|
330
|
+
graphs = build_graph_classification(
|
|
331
|
+
num_pos=8, num_neg=8, seed=0
|
|
332
|
+
) # binary y, gt_edge_mask
|
|
311
333
|
model = ... # GraphGCN (task_level="graph")
|
|
312
334
|
|
|
313
335
|
expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
|
|
@@ -320,19 +342,9 @@ Node-only methods (`GraphLIME`, `NodeMask`, `Attention`, `GNNGatedLRP`,
|
|
|
320
342
|
|
|
321
343
|
## Roadmap
|
|
322
344
|
|
|
323
|
-
|
|
324
|
-
|
|
325
|
-
-
|
|
326
|
-
- [x] Phase 3: DGL backend (adapter; integration validated with DGL 2.1 + torch 2.2.1)
|
|
327
|
-
- [x] Phase 4: GNN-LRP (layer-wise relevance for GCNs; validates the house motif in BA-Shapes)
|
|
328
|
-
- [x] Phase 4: counterfactual explanations (minimal edge/feature removal that changes the class)
|
|
329
|
-
- [x] Phase 4: LLM narration (`describe` deterministic + pluggable `narrate` LLM)
|
|
330
|
-
- [x] Phase 5: full CLI (all methods, metrics, narration and JSON export)
|
|
331
|
-
- [x] Phase 6: more methods (DeepLIFT rescale, Attention/GAT, Gradient×Input)
|
|
332
|
-
- [x] Phase 7: comparative benchmark (`compare` + CLI `bench` subcommand, table and JSON/HTML reports)
|
|
333
|
-
- [x] Phase 8: more methods (GraphLIME, NodeMask, GuidedBackprop and Random baseline)
|
|
334
|
-
- [x] Phase 10: graph-level explanations (graph-classification dataset with house
|
|
335
|
-
motif, graph-level GEA, CLI/bench without `--node` and `graph_level` flag)
|
|
345
|
+
Planned ideas (real-world Datasets, more methods, robustness metrics) are
|
|
346
|
+
tracked as issues in the repository — see
|
|
347
|
+
https://github.com/Tzinny-dev/graph-explain/issues
|
|
336
348
|
|
|
337
349
|
## License
|
|
338
350
|
|
|
@@ -1,5 +1,13 @@
|
|
|
1
1
|
# graph-explain
|
|
2
2
|
|
|
3
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
4
|
+
[](https://tzinny-dev.github.io/graph-explain/)
|
|
5
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
6
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
7
|
+
[](https://pypi.org/project/graph-explain/)
|
|
8
|
+
[](https://pypi.org/project/graph-explain/)
|
|
9
|
+
[](https://github.com/Tzinny-dev/graph-explain/blob/main/LICENSE)
|
|
10
|
+
|
|
3
11
|
Explainability library for graph-based models (Graph Neural Networks).
|
|
4
12
|
Explains a GNN's predictions in terms of **important nodes, edges and subgraphs**,
|
|
5
13
|
with built-in metrics and visualization.
|
|
@@ -68,6 +76,9 @@ pip install -e .[all]
|
|
|
68
76
|
Optional extras: `pyg` (PyTorch Geometric), `dgl` (DGL backend),
|
|
69
77
|
`interactive` (plotly/pyvis).
|
|
70
78
|
|
|
79
|
+
> 📚 **Documentation**: full API reference and guides at
|
|
80
|
+
> https://tzinny-dev.github.io/graph-explain/
|
|
81
|
+
|
|
71
82
|
## Quick start
|
|
72
83
|
|
|
73
84
|
```python
|
|
@@ -75,16 +86,18 @@ from graph_explain import Explainer, GNNExplainer, Saliency
|
|
|
75
86
|
from graph_explain.benchmarks.synthetic import build_data
|
|
76
87
|
from graph_explain.visualization import show
|
|
77
88
|
|
|
78
|
-
data = build_data(base_nodes=300, num_houses=80)
|
|
79
|
-
model = GCN(in_channels=data.x.size(1))
|
|
89
|
+
data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
|
|
90
|
+
model = GCN(in_channels=data.x.size(1)) # your trained GNN
|
|
80
91
|
model.eval()
|
|
81
92
|
|
|
82
93
|
explainer = Explainer(algorithm=GNNExplainer(epochs=150))
|
|
83
94
|
expl = explainer.explain_node(data, model, node_idx=42)
|
|
84
95
|
|
|
85
96
|
print(expl.evaluate(metrics=["fidelity", "sparsity"]))
|
|
86
|
-
print(
|
|
87
|
-
|
|
97
|
+
print(
|
|
98
|
+
expl.evaluate(metrics=["sparsity"], local=True)
|
|
99
|
+
) # sparsity over the node's k-hop subgraph
|
|
100
|
+
show(expl, show_labels=True) # highlight the explanatory subgraph
|
|
88
101
|
```
|
|
89
102
|
|
|
90
103
|
## Sparsity tuning notes
|
|
@@ -217,9 +230,15 @@ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
|
|
|
217
230
|
```python
|
|
218
231
|
from graph_explain import compare, report_html
|
|
219
232
|
|
|
220
|
-
results = compare(
|
|
221
|
-
|
|
222
|
-
|
|
233
|
+
results = compare(
|
|
234
|
+
data,
|
|
235
|
+
model,
|
|
236
|
+
node=42,
|
|
237
|
+
methods=None, # None = all
|
|
238
|
+
num_perturbations=5,
|
|
239
|
+
epochs=200,
|
|
240
|
+
)
|
|
241
|
+
report_html(results, "bench.html") # self-contained HTML report
|
|
223
242
|
```
|
|
224
243
|
|
|
225
244
|
The CLI ships an equivalent subcommand:
|
|
@@ -255,7 +274,9 @@ In Python:
|
|
|
255
274
|
from graph_explain import Explainer, evaluate_gea_graph
|
|
256
275
|
from graph_explain.benchmarks.synthetic import build_graph_classification
|
|
257
276
|
|
|
258
|
-
graphs = build_graph_classification(
|
|
277
|
+
graphs = build_graph_classification(
|
|
278
|
+
num_pos=8, num_neg=8, seed=0
|
|
279
|
+
) # binary y, gt_edge_mask
|
|
259
280
|
model = ... # GraphGCN (task_level="graph")
|
|
260
281
|
|
|
261
282
|
expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
|
|
@@ -268,20 +289,10 @@ Node-only methods (`GraphLIME`, `NodeMask`, `Attention`, `GNNGatedLRP`,
|
|
|
268
289
|
|
|
269
290
|
## Roadmap
|
|
270
291
|
|
|
271
|
-
|
|
272
|
-
|
|
273
|
-
-
|
|
274
|
-
- [x] Phase 3: DGL backend (adapter; integration validated with DGL 2.1 + torch 2.2.1)
|
|
275
|
-
- [x] Phase 4: GNN-LRP (layer-wise relevance for GCNs; validates the house motif in BA-Shapes)
|
|
276
|
-
- [x] Phase 4: counterfactual explanations (minimal edge/feature removal that changes the class)
|
|
277
|
-
- [x] Phase 4: LLM narration (`describe` deterministic + pluggable `narrate` LLM)
|
|
278
|
-
- [x] Phase 5: full CLI (all methods, metrics, narration and JSON export)
|
|
279
|
-
- [x] Phase 6: more methods (DeepLIFT rescale, Attention/GAT, Gradient×Input)
|
|
280
|
-
- [x] Phase 7: comparative benchmark (`compare` + CLI `bench` subcommand, table and JSON/HTML reports)
|
|
281
|
-
- [x] Phase 8: more methods (GraphLIME, NodeMask, GuidedBackprop and Random baseline)
|
|
282
|
-
- [x] Phase 10: graph-level explanations (graph-classification dataset with house
|
|
283
|
-
motif, graph-level GEA, CLI/bench without `--node` and `graph_level` flag)
|
|
292
|
+
Planned ideas (real-world Datasets, more methods, robustness metrics) are
|
|
293
|
+
tracked as issues in the repository — see
|
|
294
|
+
https://github.com/Tzinny-dev/graph-explain/issues
|
|
284
295
|
|
|
285
296
|
## License
|
|
286
297
|
|
|
287
|
-
MIT
|
|
298
|
+
MIT
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "graph-explain"
|
|
7
|
-
version = "0.
|
|
7
|
+
version = "0.8.0"
|
|
8
8
|
description = "Explainability library for graph-based models (GNN)"
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.10"
|
|
@@ -41,14 +41,14 @@ dependencies = [
|
|
|
41
41
|
[project.urls]
|
|
42
42
|
Homepage = "https://github.com/Tzinny-dev/graph-explain"
|
|
43
43
|
Repository = "https://github.com/Tzinny-dev/graph-explain"
|
|
44
|
-
Documentation = "https://github.
|
|
44
|
+
Documentation = "https://tzinny-dev.github.io/graph-explain/"
|
|
45
45
|
|
|
46
46
|
[project.optional-dependencies]
|
|
47
47
|
pyg = ["torch>=2.0", "torch-geometric>=2.5"]
|
|
48
48
|
dgl = ["torch>=2.0", "dgl>=2.0"]
|
|
49
49
|
interactive = ["plotly>=5.15", "pyvis>=0.3"]
|
|
50
50
|
all = ["torch>=2.0", "torch-geometric>=2.5", "dgl>=2.0", "plotly>=5.15", "pyvis>=0.3"]
|
|
51
|
-
dev = ["pytest>=7.0", "ruff>=0.5", "build>=1.0"]
|
|
51
|
+
dev = ["pytest>=7.0", "pytest-cov>=5.0", "ruff>=0.5", "build>=1.0"]
|
|
52
52
|
docs = ["sphinx>=7.2", "sphinx-rtd-theme>=2.0"]
|
|
53
53
|
publish = ["twine>=5.0"]
|
|
54
54
|
|
|
@@ -61,9 +61,38 @@ where = ["src"]
|
|
|
61
61
|
[tool.setuptools.package-data]
|
|
62
62
|
graph_explain = ["py.typed"]
|
|
63
63
|
|
|
64
|
+
[tool.mypy]
|
|
65
|
+
mypy_path = "src"
|
|
66
|
+
files = ["src/graph_explain"]
|
|
67
|
+
ignore_missing_imports = true
|
|
68
|
+
check_untyped_defs = true
|
|
69
|
+
warn_unused_ignores = false
|
|
70
|
+
exclude = ["build/", "docs/"]
|
|
71
|
+
|
|
64
72
|
[tool.pytest.ini_options]
|
|
65
73
|
testpaths = ["tests"]
|
|
74
|
+
# Coverage runs in the dedicated `coverage` CI job (see .github/workflows/ci.yml).
|
|
75
|
+
# Locally: `python -m pytest --cov=graph_explain` picks up this config.
|
|
76
|
+
[tool.coverage.run]
|
|
77
|
+
branch = true
|
|
78
|
+
source = ["graph_explain"]
|
|
79
|
+
[tool.coverage.report]
|
|
80
|
+
fail_under = 80
|
|
81
|
+
show_missing = true
|
|
82
|
+
precision = 0
|
|
83
|
+
exclude_also = [
|
|
84
|
+
"if TYPE_CHECKING:",
|
|
85
|
+
'if __name__ == "__main__":',
|
|
86
|
+
'except ImportError:',
|
|
87
|
+
'raise NotImplementedError',
|
|
88
|
+
]
|
|
89
|
+
[tool.coverage.json]
|
|
90
|
+
pretty_print = false
|
|
91
|
+
[tool.coverage.html]
|
|
92
|
+
directory = "docs/_build/coverage"
|
|
93
|
+
[tool.coverage.paths]
|
|
94
|
+
source = ["src/graph_explain/", "*/site-packages/graph_explain/"]
|
|
66
95
|
|
|
67
96
|
[tool.ruff]
|
|
68
97
|
line-length = 88
|
|
69
|
-
target-version = "py310"
|
|
98
|
+
target-version = "py310"
|
|
@@ -3,6 +3,7 @@ from __future__ import annotations
|
|
|
3
3
|
import networkx as nx
|
|
4
4
|
import numpy as np
|
|
5
5
|
import torch
|
|
6
|
+
from numpy.typing import NDArray
|
|
6
7
|
|
|
7
8
|
|
|
8
9
|
def _house_motif(offset: int, anchor_in_motif: int = 3):
|
|
@@ -56,23 +57,23 @@ def ba_shapes(
|
|
|
56
57
|
max_deg = int(degrees.max()) + 1
|
|
57
58
|
feat_dim = max(max_deg, num_features)
|
|
58
59
|
if feature_style == "random":
|
|
59
|
-
|
|
60
|
+
x_np = rng.normal(0.0, 1.0, size=(node_count, num_features)).astype(np.float32)
|
|
60
61
|
else:
|
|
61
|
-
|
|
62
|
-
|
|
63
|
-
1
|
|
64
|
-
|
|
62
|
+
x_np = np.zeros((node_count, feat_dim), dtype=np.float32)
|
|
63
|
+
x_np[
|
|
64
|
+
np.arange(node_count), np.minimum(degrees.astype(np.int64), feat_dim - 1)
|
|
65
|
+
] = 1.0
|
|
65
66
|
|
|
66
67
|
edge_index = torch.tensor(np.array(g.edges(), dtype=np.int64).T, dtype=torch.long)
|
|
67
68
|
edge_index = torch.cat([edge_index, edge_index.flip(0)], dim=1)
|
|
68
69
|
y = torch.zeros(node_count, dtype=torch.long)
|
|
69
70
|
for n, l in labels.items():
|
|
70
71
|
y[n] = l
|
|
71
|
-
x = torch.from_numpy(
|
|
72
|
+
x = torch.from_numpy(x_np)
|
|
72
73
|
|
|
73
74
|
house_anchors = torch.from_numpy(anchors)
|
|
74
75
|
|
|
75
|
-
perm = rng.permutation(node_count)
|
|
76
|
+
perm: NDArray[np.integer] = rng.permutation(node_count)
|
|
76
77
|
train_mask = torch.zeros(node_count, dtype=torch.bool)
|
|
77
78
|
test_mask = torch.zeros(node_count, dtype=torch.bool)
|
|
78
79
|
split = int(0.3 * node_count)
|
|
@@ -1,5 +1,7 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
3
5
|
import torch
|
|
4
6
|
|
|
5
7
|
from ..narration import summarize
|
|
@@ -105,7 +107,7 @@ def compare(
|
|
|
105
107
|
|
|
106
108
|
for name in methods:
|
|
107
109
|
cls = get_algorithm(name)
|
|
108
|
-
entry = {
|
|
110
|
+
entry: dict[str, Any] = {
|
|
109
111
|
"method": name,
|
|
110
112
|
"class": cls.__name__,
|
|
111
113
|
"node": node,
|
|
@@ -172,6 +174,7 @@ def compare(
|
|
|
172
174
|
lambda expl=expl_arg: float(evaluate_sparsity(expl, local=True))
|
|
173
175
|
)
|
|
174
176
|
if stability:
|
|
177
|
+
assert node is not None # stability is disabled for graph-level
|
|
175
178
|
|
|
176
179
|
def _again(d, name=name):
|
|
177
180
|
algo_r = instantiate(name, **_method_kwargs(epochs, lr, seed, top_k))
|
|
@@ -322,7 +322,7 @@ def _local_scope(explanation, hops: int) -> tuple[list[int], list[int]]:
|
|
|
322
322
|
edge_index = explanation.metadata.get("edge_index")
|
|
323
323
|
num_nodes = num_nodes or explanation.metadata.get("num_nodes")
|
|
324
324
|
if edge_index is None or num_nodes is None:
|
|
325
|
-
return
|
|
325
|
+
return [], []
|
|
326
326
|
node_idx = int(explanation.node_idx)
|
|
327
327
|
visited = {node_idx}
|
|
328
328
|
frontier = {node_idx}
|
|
@@ -4,11 +4,13 @@ import inspect
|
|
|
4
4
|
from collections.abc import Callable
|
|
5
5
|
from typing import Any
|
|
6
6
|
|
|
7
|
-
|
|
7
|
+
from ..methods.base import ExplanationAlgorithm
|
|
8
|
+
|
|
9
|
+
_ALGORITHMS: dict[str, type[ExplanationAlgorithm]] = {}
|
|
8
10
|
_ALIASES: dict[str, str] = {}
|
|
9
11
|
|
|
10
12
|
|
|
11
|
-
def _accepted_params(cls: type) -> set[str]:
|
|
13
|
+
def _accepted_params(cls: type[ExplanationAlgorithm]) -> set[str]:
|
|
12
14
|
try:
|
|
13
15
|
sig = inspect.signature(cls.__init__)
|
|
14
16
|
except (TypeError, ValueError):
|
|
@@ -23,7 +25,7 @@ def _accepted_params(cls: type) -> set[str]:
|
|
|
23
25
|
|
|
24
26
|
|
|
25
27
|
def register(name: str, *aliases: str) -> Callable[[type], type]:
|
|
26
|
-
def decorator(cls: type) -> type:
|
|
28
|
+
def decorator(cls: type[ExplanationAlgorithm]) -> type[ExplanationAlgorithm]:
|
|
27
29
|
_ALGORITHMS[name] = cls
|
|
28
30
|
for alias in aliases:
|
|
29
31
|
_ALIASES[alias] = name
|
|
@@ -33,7 +35,7 @@ def register(name: str, *aliases: str) -> Callable[[type], type]:
|
|
|
33
35
|
return decorator
|
|
34
36
|
|
|
35
37
|
|
|
36
|
-
def get_algorithm(name: str) -> type:
|
|
38
|
+
def get_algorithm(name: str) -> type[ExplanationAlgorithm]:
|
|
37
39
|
registered = _ALGORITHMS.get(name) or _ALGORITHMS.get(_ALIASES.get(name, ""))
|
|
38
40
|
if registered is None:
|
|
39
41
|
from ..methods import _available_methods
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/attention/attention.py
RENAMED
|
@@ -114,7 +114,7 @@ class AttentionExplainer(ExplanationAlgorithm):
|
|
|
114
114
|
feature_importance=None,
|
|
115
115
|
prediction_original=logits[nodes[0]].detach().cpu(),
|
|
116
116
|
prediction_explanation=None,
|
|
117
|
-
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else
|
|
117
|
+
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else None,
|
|
118
118
|
target_class=target_cls,
|
|
119
119
|
)
|
|
120
120
|
|
|
@@ -3,8 +3,6 @@ from __future__ import annotations
|
|
|
3
3
|
from abc import ABC, abstractmethod
|
|
4
4
|
from typing import Any
|
|
5
5
|
|
|
6
|
-
import torch
|
|
7
|
-
|
|
8
6
|
|
|
9
7
|
class ExplanationAlgorithm(ABC):
|
|
10
8
|
name = "base"
|
|
@@ -16,7 +14,7 @@ class ExplanationAlgorithm(ABC):
|
|
|
16
14
|
backend: Any,
|
|
17
15
|
model: Any,
|
|
18
16
|
data: Any,
|
|
19
|
-
index:
|
|
17
|
+
index: Any = None,
|
|
20
18
|
target_class: int | None = None,
|
|
21
19
|
**kwargs,
|
|
22
20
|
) -> Any: ...
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/grad_x_input.py
RENAMED
|
@@ -96,7 +96,7 @@ class GradXInput(ExplanationAlgorithm):
|
|
|
96
96
|
),
|
|
97
97
|
prediction_original=logits[nodes[0]].detach().cpu(),
|
|
98
98
|
prediction_explanation=None,
|
|
99
|
-
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else
|
|
99
|
+
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else None,
|
|
100
100
|
target_class=target_cls,
|
|
101
101
|
)
|
|
102
102
|
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/guided_backprop.py
RENAMED
|
@@ -54,7 +54,8 @@ class GuidedBackprop(ExplanationAlgorithm):
|
|
|
54
54
|
target = max(0, min(int(target), logits.size(1) - 1))
|
|
55
55
|
|
|
56
56
|
relus = [m for m in model.modules() if isinstance(m, nn.ReLU)]
|
|
57
|
-
|
|
57
|
+
guided = False
|
|
58
|
+
hooks: list[torch.utils.hooks.RemovableHandle] = []
|
|
58
59
|
if relus and self.fallback_to_gradient:
|
|
59
60
|
guided = True
|
|
60
61
|
for module in relus:
|
|
@@ -29,7 +29,7 @@ class IntegratedGradients(ExplanationAlgorithm):
|
|
|
29
29
|
backend: Any,
|
|
30
30
|
model: Any,
|
|
31
31
|
data: Any,
|
|
32
|
-
index: int |
|
|
32
|
+
index: int | torch.Tensor | None = None,
|
|
33
33
|
target_class: int | None = None,
|
|
34
34
|
**kwargs,
|
|
35
35
|
) -> Explanation:
|
|
@@ -94,7 +94,7 @@ class IntegratedGradients(ExplanationAlgorithm):
|
|
|
94
94
|
feature_importance=feature_importance.cpu(),
|
|
95
95
|
prediction_original=logits[idx[0]].detach().reshape(1, -1).cpu(),
|
|
96
96
|
prediction_explanation=None,
|
|
97
|
-
node_idx=int(idx[0].item()) if isinstance(index, int) else
|
|
97
|
+
node_idx=int(idx[0].item()) if isinstance(index, int) else None,
|
|
98
98
|
target_class=target_class,
|
|
99
99
|
)
|
|
100
100
|
|
|
@@ -28,7 +28,7 @@ class Saliency(ExplanationAlgorithm):
|
|
|
28
28
|
backend: Any,
|
|
29
29
|
model: Any,
|
|
30
30
|
data: Any,
|
|
31
|
-
index: int |
|
|
31
|
+
index: int | torch.Tensor | None = None,
|
|
32
32
|
target_class: int | None = None,
|
|
33
33
|
**kwargs,
|
|
34
34
|
) -> Explanation:
|
|
@@ -88,6 +88,10 @@ class Saliency(ExplanationAlgorithm):
|
|
|
88
88
|
edge_importance=None,
|
|
89
89
|
prediction_original=prediction_original.cpu(),
|
|
90
90
|
prediction_explanation=None,
|
|
91
|
-
node_idx=
|
|
91
|
+
node_idx=(
|
|
92
|
+
int(idx[0].item())
|
|
93
|
+
if isinstance(index, int)
|
|
94
|
+
else (None if index is None else int(index[0]))
|
|
95
|
+
),
|
|
92
96
|
target_class=target_class,
|
|
93
97
|
)
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/gnn_explainer.py
RENAMED
|
@@ -38,7 +38,7 @@ class GNNExplainer(ExplanationAlgorithm):
|
|
|
38
38
|
backend: Any,
|
|
39
39
|
model: Any,
|
|
40
40
|
data: Any,
|
|
41
|
-
index: int | torch.Tensor,
|
|
41
|
+
index: int | torch.Tensor | None = None,
|
|
42
42
|
target_class: int | None = None,
|
|
43
43
|
**kwargs,
|
|
44
44
|
) -> Explanation:
|
|
@@ -148,15 +148,22 @@ class GNNExplainer(ExplanationAlgorithm):
|
|
|
148
148
|
node_full = self._scatter_node(node_mask, sub_nodes, full_num_nodes)
|
|
149
149
|
edge_full = self._scatter_edge(edge_mask, sub_edge_mask, full_edge_count)
|
|
150
150
|
|
|
151
|
+
node_idx = None
|
|
152
|
+
if index is not None:
|
|
153
|
+
if torch.is_tensor(index):
|
|
154
|
+
node_idx = int(index[0])
|
|
155
|
+
elif isinstance(index, int):
|
|
156
|
+
node_idx = int(index)
|
|
157
|
+
else:
|
|
158
|
+
node_idx = int(index[0])
|
|
159
|
+
|
|
151
160
|
return Explanation(
|
|
152
161
|
node_importance=node_full,
|
|
153
162
|
edge_importance=edge_full,
|
|
154
163
|
feature_importance=None,
|
|
155
164
|
prediction_original=pred_orig.cpu(),
|
|
156
165
|
prediction_explanation=pred_masked.cpu(),
|
|
157
|
-
node_idx=
|
|
158
|
-
if sub_graph_level
|
|
159
|
-
else (int(index[0]) if torch.is_tensor(index) else int(index)),
|
|
166
|
+
node_idx=node_idx,
|
|
160
167
|
target_class=target_class,
|
|
161
168
|
metadata={
|
|
162
169
|
"sub_nodes": sub_nodes,
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/subgraphx.py
RENAMED
|
@@ -52,7 +52,7 @@ class SubgraphX(ExplanationAlgorithm):
|
|
|
52
52
|
backend: Any,
|
|
53
53
|
model: Any,
|
|
54
54
|
data: Any,
|
|
55
|
-
index: int | torch.Tensor | None,
|
|
55
|
+
index: int | torch.Tensor | None = None,
|
|
56
56
|
target_class: int | None = None,
|
|
57
57
|
**kwargs,
|
|
58
58
|
) -> Explanation:
|
|
@@ -370,7 +370,7 @@ class SubgraphX(ExplanationAlgorithm):
|
|
|
370
370
|
sel = set(selected)
|
|
371
371
|
seen = set()
|
|
372
372
|
comps = 0
|
|
373
|
-
adj = {}
|
|
373
|
+
adj: dict[int, set[int]] = {}
|
|
374
374
|
src = edge_index[0].tolist()
|
|
375
375
|
dst = edge_index[1].tolist()
|
|
376
376
|
for u, v in zip(src, dst):
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
-
from typing import Any
|
|
3
|
+
from typing import Any, cast
|
|
4
4
|
|
|
5
5
|
import torch
|
|
6
6
|
from torch import nn
|
|
@@ -117,7 +117,7 @@ class DeepLift(ExplanationAlgorithm):
|
|
|
117
117
|
),
|
|
118
118
|
prediction_original=logits[nodes[0]].detach().cpu(),
|
|
119
119
|
prediction_explanation=None,
|
|
120
|
-
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else
|
|
120
|
+
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else None,
|
|
121
121
|
target_class=target_cls,
|
|
122
122
|
)
|
|
123
123
|
|
|
@@ -242,7 +242,7 @@ class DeepLift(ExplanationAlgorithm):
|
|
|
242
242
|
else torch.ones(ei.size(1), device=x.device)
|
|
243
243
|
)
|
|
244
244
|
src, dst = ei[0], ei[1]
|
|
245
|
-
W = conv.lin.weight # (out, in)
|
|
245
|
+
W = cast(Any, conv.lin).weight # (out, in)
|
|
246
246
|
|
|
247
247
|
mul_agg = mul @ W # (N, F_in)
|
|
248
248
|
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
-
from typing import Any
|
|
3
|
+
from typing import Any, cast
|
|
4
4
|
|
|
5
5
|
import torch
|
|
6
6
|
from torch import nn
|
|
@@ -127,7 +127,7 @@ class GNNGatedLRP(ExplanationAlgorithm):
|
|
|
127
127
|
),
|
|
128
128
|
prediction_original=logits[nodes[0]].detach().cpu(),
|
|
129
129
|
prediction_explanation=None,
|
|
130
|
-
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else
|
|
130
|
+
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else None,
|
|
131
131
|
target_class=target_cls,
|
|
132
132
|
)
|
|
133
133
|
|
|
@@ -210,7 +210,7 @@ class GNNGatedLRP(ExplanationAlgorithm):
|
|
|
210
210
|
agg.index_add_(0, dst, norm[:, None] * x[src])
|
|
211
211
|
agg_pos = agg.clamp(min=0)
|
|
212
212
|
|
|
213
|
-
r_agg = self._linear_lrp(conv.lin.weight, agg_pos, r, eps) # (N, F)
|
|
213
|
+
r_agg = self._linear_lrp(cast(Any, conv.lin).weight, agg_pos, r, eps) # (N, F)
|
|
214
214
|
|
|
215
215
|
frac = msg / agg_pos[dst].clamp(min=eps) # (E, F)
|
|
216
216
|
r_msg = (r_agg[dst] * frac).sum(dim=-1) # (E,)
|
|
@@ -1,12 +1,12 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: graph-explain
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.8.0
|
|
4
4
|
Summary: Explainability library for graph-based models (GNN)
|
|
5
5
|
Author: graph-explain contributors
|
|
6
6
|
License-Expression: MIT
|
|
7
7
|
Project-URL: Homepage, https://github.com/Tzinny-dev/graph-explain
|
|
8
8
|
Project-URL: Repository, https://github.com/Tzinny-dev/graph-explain
|
|
9
|
-
Project-URL: Documentation, https://github.
|
|
9
|
+
Project-URL: Documentation, https://tzinny-dev.github.io/graph-explain/
|
|
10
10
|
Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
|
|
11
11
|
Classifier: Development Status :: 4 - Beta
|
|
12
12
|
Classifier: Intended Audience :: Science/Research
|
|
@@ -41,6 +41,7 @@ Requires-Dist: plotly>=5.15; extra == "all"
|
|
|
41
41
|
Requires-Dist: pyvis>=0.3; extra == "all"
|
|
42
42
|
Provides-Extra: dev
|
|
43
43
|
Requires-Dist: pytest>=7.0; extra == "dev"
|
|
44
|
+
Requires-Dist: pytest-cov>=5.0; extra == "dev"
|
|
44
45
|
Requires-Dist: ruff>=0.5; extra == "dev"
|
|
45
46
|
Requires-Dist: build>=1.0; extra == "dev"
|
|
46
47
|
Provides-Extra: docs
|
|
@@ -52,6 +53,14 @@ Dynamic: license-file
|
|
|
52
53
|
|
|
53
54
|
# graph-explain
|
|
54
55
|
|
|
56
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
57
|
+
[](https://tzinny-dev.github.io/graph-explain/)
|
|
58
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
59
|
+
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
60
|
+
[](https://pypi.org/project/graph-explain/)
|
|
61
|
+
[](https://pypi.org/project/graph-explain/)
|
|
62
|
+
[](https://github.com/Tzinny-dev/graph-explain/blob/main/LICENSE)
|
|
63
|
+
|
|
55
64
|
Explainability library for graph-based models (Graph Neural Networks).
|
|
56
65
|
Explains a GNN's predictions in terms of **important nodes, edges and subgraphs**,
|
|
57
66
|
with built-in metrics and visualization.
|
|
@@ -120,6 +129,9 @@ pip install -e .[all]
|
|
|
120
129
|
Optional extras: `pyg` (PyTorch Geometric), `dgl` (DGL backend),
|
|
121
130
|
`interactive` (plotly/pyvis).
|
|
122
131
|
|
|
132
|
+
> 📚 **Documentation**: full API reference and guides at
|
|
133
|
+
> https://tzinny-dev.github.io/graph-explain/
|
|
134
|
+
|
|
123
135
|
## Quick start
|
|
124
136
|
|
|
125
137
|
```python
|
|
@@ -127,16 +139,18 @@ from graph_explain import Explainer, GNNExplainer, Saliency
|
|
|
127
139
|
from graph_explain.benchmarks.synthetic import build_data
|
|
128
140
|
from graph_explain.visualization import show
|
|
129
141
|
|
|
130
|
-
data = build_data(base_nodes=300, num_houses=80)
|
|
131
|
-
model = GCN(in_channels=data.x.size(1))
|
|
142
|
+
data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
|
|
143
|
+
model = GCN(in_channels=data.x.size(1)) # your trained GNN
|
|
132
144
|
model.eval()
|
|
133
145
|
|
|
134
146
|
explainer = Explainer(algorithm=GNNExplainer(epochs=150))
|
|
135
147
|
expl = explainer.explain_node(data, model, node_idx=42)
|
|
136
148
|
|
|
137
149
|
print(expl.evaluate(metrics=["fidelity", "sparsity"]))
|
|
138
|
-
print(
|
|
139
|
-
|
|
150
|
+
print(
|
|
151
|
+
expl.evaluate(metrics=["sparsity"], local=True)
|
|
152
|
+
) # sparsity over the node's k-hop subgraph
|
|
153
|
+
show(expl, show_labels=True) # highlight the explanatory subgraph
|
|
140
154
|
```
|
|
141
155
|
|
|
142
156
|
## Sparsity tuning notes
|
|
@@ -269,9 +283,15 @@ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
|
|
|
269
283
|
```python
|
|
270
284
|
from graph_explain import compare, report_html
|
|
271
285
|
|
|
272
|
-
results = compare(
|
|
273
|
-
|
|
274
|
-
|
|
286
|
+
results = compare(
|
|
287
|
+
data,
|
|
288
|
+
model,
|
|
289
|
+
node=42,
|
|
290
|
+
methods=None, # None = all
|
|
291
|
+
num_perturbations=5,
|
|
292
|
+
epochs=200,
|
|
293
|
+
)
|
|
294
|
+
report_html(results, "bench.html") # self-contained HTML report
|
|
275
295
|
```
|
|
276
296
|
|
|
277
297
|
The CLI ships an equivalent subcommand:
|
|
@@ -307,7 +327,9 @@ In Python:
|
|
|
307
327
|
from graph_explain import Explainer, evaluate_gea_graph
|
|
308
328
|
from graph_explain.benchmarks.synthetic import build_graph_classification
|
|
309
329
|
|
|
310
|
-
graphs = build_graph_classification(
|
|
330
|
+
graphs = build_graph_classification(
|
|
331
|
+
num_pos=8, num_neg=8, seed=0
|
|
332
|
+
) # binary y, gt_edge_mask
|
|
311
333
|
model = ... # GraphGCN (task_level="graph")
|
|
312
334
|
|
|
313
335
|
expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
|
|
@@ -320,19 +342,9 @@ Node-only methods (`GraphLIME`, `NodeMask`, `Attention`, `GNNGatedLRP`,
|
|
|
320
342
|
|
|
321
343
|
## Roadmap
|
|
322
344
|
|
|
323
|
-
|
|
324
|
-
|
|
325
|
-
-
|
|
326
|
-
- [x] Phase 3: DGL backend (adapter; integration validated with DGL 2.1 + torch 2.2.1)
|
|
327
|
-
- [x] Phase 4: GNN-LRP (layer-wise relevance for GCNs; validates the house motif in BA-Shapes)
|
|
328
|
-
- [x] Phase 4: counterfactual explanations (minimal edge/feature removal that changes the class)
|
|
329
|
-
- [x] Phase 4: LLM narration (`describe` deterministic + pluggable `narrate` LLM)
|
|
330
|
-
- [x] Phase 5: full CLI (all methods, metrics, narration and JSON export)
|
|
331
|
-
- [x] Phase 6: more methods (DeepLIFT rescale, Attention/GAT, Gradient×Input)
|
|
332
|
-
- [x] Phase 7: comparative benchmark (`compare` + CLI `bench` subcommand, table and JSON/HTML reports)
|
|
333
|
-
- [x] Phase 8: more methods (GraphLIME, NodeMask, GuidedBackprop and Random baseline)
|
|
334
|
-
- [x] Phase 10: graph-level explanations (graph-classification dataset with house
|
|
335
|
-
motif, graph-level GEA, CLI/bench without `--node` and `graph_level` flag)
|
|
345
|
+
Planned ideas (real-world Datasets, more methods, robustness metrics) are
|
|
346
|
+
tracked as issues in the repository — see
|
|
347
|
+
https://github.com/Tzinny-dev/graph-explain/issues
|
|
336
348
|
|
|
337
349
|
## License
|
|
338
350
|
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/baseline/random_baseline.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/__init__.py
RENAMED
|
File without changes
|
{graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/node_mask.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|