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.
Files changed (60) hide show
  1. {graph_explain-0.7.2 → graph_explain-0.8.0}/LICENSE +1 -1
  2. {graph_explain-0.7.2/src/graph_explain.egg-info → graph_explain-0.8.0}/PKG-INFO +35 -23
  3. {graph_explain-0.7.2 → graph_explain-0.8.0}/README.md +33 -22
  4. {graph_explain-0.7.2 → graph_explain-0.8.0}/pyproject.toml +33 -4
  5. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/__init__.py +1 -1
  6. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/benchmarks/synthetic.py +8 -7
  7. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/benchmark.py +4 -1
  8. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/evaluation.py +1 -1
  9. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/explanation.py +1 -1
  10. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/registry.py +6 -4
  11. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/attention/attention.py +1 -1
  12. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/base.py +1 -3
  13. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/grad_x_input.py +1 -1
  14. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/guided_backprop.py +2 -1
  15. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/integrated_gradients.py +2 -2
  16. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/saliency.py +6 -2
  17. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/gnn_explainer.py +11 -4
  18. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/pg_explainer.py +1 -1
  19. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/subgraphx.py +2 -2
  20. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/deeplift.py +3 -3
  21. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/gnn_lrp.py +3 -3
  22. {graph_explain-0.7.2 → graph_explain-0.8.0/src/graph_explain.egg-info}/PKG-INFO +35 -23
  23. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/requires.txt +1 -0
  24. {graph_explain-0.7.2 → graph_explain-0.8.0}/setup.cfg +0 -0
  25. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/backends/__init__.py +0 -0
  26. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/backends/base.py +0 -0
  27. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/backends/dgl.py +0 -0
  28. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/benchmarks/__init__.py +0 -0
  29. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/cli.py +0 -0
  30. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/__init__.py +0 -0
  31. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/explainer.py +0 -0
  32. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/core/model_utils.py +0 -0
  33. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/__init__.py +0 -0
  34. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/baseline/random_baseline.py +0 -0
  35. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/counterfactual/counterfactual.py +0 -0
  36. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/feature/graph_lime.py +0 -0
  37. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/__init__.py +0 -0
  38. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/__init__.py +0 -0
  39. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/node_mask.py +0 -0
  40. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/narration/__init__.py +0 -0
  41. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/narration/narrator.py +0 -0
  42. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/py.typed +0 -0
  43. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/visualization/__init__.py +0 -0
  44. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/visualization/interactive.py +0 -0
  45. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain/visualization/static.py +0 -0
  46. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/SOURCES.txt +0 -0
  47. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/dependency_links.txt +0 -0
  48. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/entry_points.txt +0 -0
  49. {graph_explain-0.7.2 → graph_explain-0.8.0}/src/graph_explain.egg-info/top_level.txt +0 -0
  50. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_cli.py +0 -0
  51. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_core.py +0 -0
  52. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_counterfactual.py +0 -0
  53. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_dgl_integration.py +0 -0
  54. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_narration.py +0 -0
  55. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase10.py +0 -0
  56. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase2.py +0 -0
  57. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase3.py +0 -0
  58. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase4.py +0 -0
  59. {graph_explain-0.7.2 → graph_explain-0.8.0}/tests/test_phase6.py +0 -0
  60. {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.7.2
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.com/Tzinny-dev/graph-explain#readme
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
+ [![CI](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml/badge.svg)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
57
+ [![docs](https://github.com/Tzinny-dev/graph-explain/actions/workflows/docs.yml/badge.svg)](https://tzinny-dev.github.io/graph-explain/)
58
+ [![coverage](https://tzinny-dev.github.io/graph-explain/_static/badges/coverage.svg)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
59
+ [![mypy](https://img.shields.io/badge/mypy-checked-blue)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
60
+ [![PyPI](https://img.shields.io/pypi/v/graph-explain.svg)](https://pypi.org/project/graph-explain/)
61
+ [![Python](https://img.shields.io/pypi/pyversions/graph-explain.svg)](https://pypi.org/project/graph-explain/)
62
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](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) # BA-Shapes with ground truth
131
- model = GCN(in_channels=data.x.size(1)) # your trained GNN
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(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
139
- show(expl, show_labels=True) # highlight the explanatory subgraph
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(data, model, node=42, methods=None, # None = all
273
- num_perturbations=5, epochs=200)
274
- report_html(results, "bench.html") # self-contained HTML report
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(num_pos=8, num_neg=8, seed=0) # binary y, gt_edge_mask
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
- - [x] Phase 2: PGExplainer, SubgraphX, Integrated Gradients
324
- - [x] Phase 2: interactive visualization (pyvis → HTML)
325
- - [x] Phase 3: full metrics (fidelity±, stability, GEA)
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
+ [![CI](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml/badge.svg)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
4
+ [![docs](https://github.com/Tzinny-dev/graph-explain/actions/workflows/docs.yml/badge.svg)](https://tzinny-dev.github.io/graph-explain/)
5
+ [![coverage](https://tzinny-dev.github.io/graph-explain/_static/badges/coverage.svg)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
6
+ [![mypy](https://img.shields.io/badge/mypy-checked-blue)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
7
+ [![PyPI](https://img.shields.io/pypi/v/graph-explain.svg)](https://pypi.org/project/graph-explain/)
8
+ [![Python](https://img.shields.io/pypi/pyversions/graph-explain.svg)](https://pypi.org/project/graph-explain/)
9
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](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) # BA-Shapes with ground truth
79
- model = GCN(in_channels=data.x.size(1)) # your trained GNN
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(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
87
- show(expl, show_labels=True) # highlight the explanatory subgraph
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(data, model, node=42, methods=None, # None = all
221
- num_perturbations=5, epochs=200)
222
- report_html(results, "bench.html") # self-contained HTML report
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(num_pos=8, num_neg=8, seed=0) # binary y, gt_edge_mask
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
- - [x] Phase 2: PGExplainer, SubgraphX, Integrated Gradients
272
- - [x] Phase 2: interactive visualization (pyvis → HTML)
273
- - [x] Phase 3: full metrics (fidelity±, stability, GEA)
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.2"
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.com/Tzinny-dev/graph-explain#readme"
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"
@@ -35,7 +35,7 @@ from .methods import (
35
35
  from .narration import Narrator, describe, narrate, summarize
36
36
  from .visualization import show, visualize_interactive, visualize_static
37
37
 
38
- __version__ = "0.7.2"
38
+ __version__ = "0.8.0"
39
39
 
40
40
  __all__ = [
41
41
  "AttentionExplainer",
@@ -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
- x = rng.normal(0.0, 1.0, size=(node_count, num_features)).astype(np.float32)
60
+ x_np = rng.normal(0.0, 1.0, size=(node_count, num_features)).astype(np.float32)
60
61
  else:
61
- x = np.zeros((node_count, feat_dim), dtype=np.float32)
62
- x[np.arange(node_count), np.minimum(degrees.astype(np.int64), feat_dim - 1)] = (
63
- 1.0
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(x)
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 None, None
325
+ return [], []
326
326
  node_idx = int(explanation.node_idx)
327
327
  visited = {node_idx}
328
328
  frontier = {node_idx}
@@ -42,7 +42,7 @@ class Explanation:
42
42
  raise ValueError(
43
43
  "La explicación se creó sin backend/backing_data en metadata"
44
44
  )
45
- G = nx.Graph()
45
+ G: nx.Graph = nx.Graph()
46
46
  keep_edges = []
47
47
  if self.edge_importance is not None:
48
48
  edge_index = backend.edge_index(data)
@@ -4,11 +4,13 @@ import inspect
4
4
  from collections.abc import Callable
5
5
  from typing import Any
6
6
 
7
- _ALGORITHMS: dict[str, type] = {}
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
@@ -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 index,
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: int | list[int] | torch.Tensor,
17
+ index: Any = None,
20
18
  target_class: int | None = None,
21
19
  **kwargs,
22
20
  ) -> Any: ...
@@ -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 index,
99
+ node_idx=int(nodes[0]) if nodes.shape[0] == 1 else None,
100
100
  target_class=target_cls,
101
101
  )
102
102
 
@@ -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
- hooks, guided = [], False
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 | list[int] | torch.Tensor,
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 index,
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 | list[int] | torch.Tensor,
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=int(idx[0].item()) if isinstance(index, int) else index,
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
  )
@@ -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=None
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,
@@ -53,7 +53,7 @@ class PGExplainer(ExplanationAlgorithm):
53
53
  backend: Any,
54
54
  model: Any,
55
55
  data: Any,
56
- index: int | torch.Tensor | None,
56
+ index: int | torch.Tensor | None = None,
57
57
  target_class: int | None = None,
58
58
  **kwargs,
59
59
  ) -> Explanation:
@@ -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 index,
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 index,
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.7.2
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.com/Tzinny-dev/graph-explain#readme
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
+ [![CI](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml/badge.svg)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
57
+ [![docs](https://github.com/Tzinny-dev/graph-explain/actions/workflows/docs.yml/badge.svg)](https://tzinny-dev.github.io/graph-explain/)
58
+ [![coverage](https://tzinny-dev.github.io/graph-explain/_static/badges/coverage.svg)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
59
+ [![mypy](https://img.shields.io/badge/mypy-checked-blue)](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
60
+ [![PyPI](https://img.shields.io/pypi/v/graph-explain.svg)](https://pypi.org/project/graph-explain/)
61
+ [![Python](https://img.shields.io/pypi/pyversions/graph-explain.svg)](https://pypi.org/project/graph-explain/)
62
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](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) # BA-Shapes with ground truth
131
- model = GCN(in_channels=data.x.size(1)) # your trained GNN
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(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
139
- show(expl, show_labels=True) # highlight the explanatory subgraph
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(data, model, node=42, methods=None, # None = all
273
- num_perturbations=5, epochs=200)
274
- report_html(results, "bench.html") # self-contained HTML report
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(num_pos=8, num_neg=8, seed=0) # binary y, gt_edge_mask
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
- - [x] Phase 2: PGExplainer, SubgraphX, Integrated Gradients
324
- - [x] Phase 2: interactive visualization (pyvis → HTML)
325
- - [x] Phase 3: full metrics (fidelity±, stability, GEA)
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
 
@@ -12,6 +12,7 @@ pyvis>=0.3
12
12
 
13
13
  [dev]
14
14
  pytest>=7.0
15
+ pytest-cov>=5.0
15
16
  ruff>=0.5
16
17
  build>=1.0
17
18
 
File without changes