graph-explain 0.7.3__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.3 → graph_explain-0.8.0}/LICENSE +1 -1
  2. {graph_explain-0.7.3/src/graph_explain.egg-info → graph_explain-0.8.0}/PKG-INFO +22 -9
  3. {graph_explain-0.7.3 → graph_explain-0.8.0}/README.md +21 -9
  4. {graph_explain-0.7.3 → graph_explain-0.8.0}/pyproject.toml +32 -3
  5. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/__init__.py +1 -1
  6. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/benchmarks/synthetic.py +8 -7
  7. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/benchmark.py +4 -1
  8. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/evaluation.py +1 -1
  9. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/explanation.py +1 -1
  10. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/registry.py +6 -4
  11. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/attention/attention.py +1 -1
  12. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/base.py +1 -3
  13. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/grad_x_input.py +1 -1
  14. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/guided_backprop.py +2 -1
  15. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/integrated_gradients.py +2 -2
  16. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/saliency.py +6 -2
  17. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/gnn_explainer.py +11 -4
  18. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/pg_explainer.py +1 -1
  19. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/subgraphx.py +2 -2
  20. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/deeplift.py +3 -3
  21. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/gnn_lrp.py +3 -3
  22. {graph_explain-0.7.3 → graph_explain-0.8.0/src/graph_explain.egg-info}/PKG-INFO +22 -9
  23. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/requires.txt +1 -0
  24. {graph_explain-0.7.3 → graph_explain-0.8.0}/setup.cfg +0 -0
  25. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/backends/__init__.py +0 -0
  26. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/backends/base.py +0 -0
  27. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/backends/dgl.py +0 -0
  28. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/benchmarks/__init__.py +0 -0
  29. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/cli.py +0 -0
  30. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/__init__.py +0 -0
  31. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/explainer.py +0 -0
  32. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/model_utils.py +0 -0
  33. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/__init__.py +0 -0
  34. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/baseline/random_baseline.py +0 -0
  35. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/counterfactual/counterfactual.py +0 -0
  36. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/feature/graph_lime.py +0 -0
  37. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/__init__.py +0 -0
  38. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/__init__.py +0 -0
  39. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/node_mask.py +0 -0
  40. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/narration/__init__.py +0 -0
  41. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/narration/narrator.py +0 -0
  42. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/py.typed +0 -0
  43. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/visualization/__init__.py +0 -0
  44. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/visualization/interactive.py +0 -0
  45. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/visualization/static.py +0 -0
  46. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/SOURCES.txt +0 -0
  47. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/dependency_links.txt +0 -0
  48. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/entry_points.txt +0 -0
  49. {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/top_level.txt +0 -0
  50. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_cli.py +0 -0
  51. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_core.py +0 -0
  52. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_counterfactual.py +0 -0
  53. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_dgl_integration.py +0 -0
  54. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_narration.py +0 -0
  55. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase10.py +0 -0
  56. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase2.py +0 -0
  57. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase3.py +0 -0
  58. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase4.py +0 -0
  59. {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase6.py +0 -0
  60. {graph_explain-0.7.3 → 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,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: graph-explain
3
- Version: 0.7.3
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
@@ -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
@@ -54,6 +55,8 @@ Dynamic: license-file
54
55
 
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)
56
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)
57
60
  [![PyPI](https://img.shields.io/pypi/v/graph-explain.svg)](https://pypi.org/project/graph-explain/)
58
61
  [![Python](https://img.shields.io/pypi/pyversions/graph-explain.svg)](https://pypi.org/project/graph-explain/)
59
62
  [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://github.com/Tzinny-dev/graph-explain/blob/main/LICENSE)
@@ -136,16 +139,18 @@ from graph_explain import Explainer, GNNExplainer, Saliency
136
139
  from graph_explain.benchmarks.synthetic import build_data
137
140
  from graph_explain.visualization import show
138
141
 
139
- data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
140
- 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
141
144
  model.eval()
142
145
 
143
146
  explainer = Explainer(algorithm=GNNExplainer(epochs=150))
144
147
  expl = explainer.explain_node(data, model, node_idx=42)
145
148
 
146
149
  print(expl.evaluate(metrics=["fidelity", "sparsity"]))
147
- print(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
148
- 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
149
154
  ```
150
155
 
151
156
  ## Sparsity tuning notes
@@ -278,9 +283,15 @@ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
278
283
  ```python
279
284
  from graph_explain import compare, report_html
280
285
 
281
- results = compare(data, model, node=42, methods=None, # None = all
282
- num_perturbations=5, epochs=200)
283
- 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
284
295
  ```
285
296
 
286
297
  The CLI ships an equivalent subcommand:
@@ -316,7 +327,9 @@ In Python:
316
327
  from graph_explain import Explainer, evaluate_gea_graph
317
328
  from graph_explain.benchmarks.synthetic import build_graph_classification
318
329
 
319
- 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
320
333
  model = ... # GraphGCN (task_level="graph")
321
334
 
322
335
  expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
@@ -2,6 +2,8 @@
2
2
 
3
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
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)
5
7
  [![PyPI](https://img.shields.io/pypi/v/graph-explain.svg)](https://pypi.org/project/graph-explain/)
6
8
  [![Python](https://img.shields.io/pypi/pyversions/graph-explain.svg)](https://pypi.org/project/graph-explain/)
7
9
  [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://github.com/Tzinny-dev/graph-explain/blob/main/LICENSE)
@@ -84,16 +86,18 @@ from graph_explain import Explainer, GNNExplainer, Saliency
84
86
  from graph_explain.benchmarks.synthetic import build_data
85
87
  from graph_explain.visualization import show
86
88
 
87
- data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
88
- 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
89
91
  model.eval()
90
92
 
91
93
  explainer = Explainer(algorithm=GNNExplainer(epochs=150))
92
94
  expl = explainer.explain_node(data, model, node_idx=42)
93
95
 
94
96
  print(expl.evaluate(metrics=["fidelity", "sparsity"]))
95
- print(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
96
- 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
97
101
  ```
98
102
 
99
103
  ## Sparsity tuning notes
@@ -226,9 +230,15 @@ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
226
230
  ```python
227
231
  from graph_explain import compare, report_html
228
232
 
229
- results = compare(data, model, node=42, methods=None, # None = all
230
- num_perturbations=5, epochs=200)
231
- 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
232
242
  ```
233
243
 
234
244
  The CLI ships an equivalent subcommand:
@@ -264,7 +274,9 @@ In Python:
264
274
  from graph_explain import Explainer, evaluate_gea_graph
265
275
  from graph_explain.benchmarks.synthetic import build_graph_classification
266
276
 
267
- 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
268
280
  model = ... # GraphGCN (task_level="graph")
269
281
 
270
282
  expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
@@ -283,4 +295,4 @@ https://github.com/Tzinny-dev/graph-explain/issues
283
295
 
284
296
  ## License
285
297
 
286
- 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.3"
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"
@@ -48,7 +48,7 @@ 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.3"
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,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: graph-explain
3
- Version: 0.7.3
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
@@ -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
@@ -54,6 +55,8 @@ Dynamic: license-file
54
55
 
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)
56
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)
57
60
  [![PyPI](https://img.shields.io/pypi/v/graph-explain.svg)](https://pypi.org/project/graph-explain/)
58
61
  [![Python](https://img.shields.io/pypi/pyversions/graph-explain.svg)](https://pypi.org/project/graph-explain/)
59
62
  [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://github.com/Tzinny-dev/graph-explain/blob/main/LICENSE)
@@ -136,16 +139,18 @@ from graph_explain import Explainer, GNNExplainer, Saliency
136
139
  from graph_explain.benchmarks.synthetic import build_data
137
140
  from graph_explain.visualization import show
138
141
 
139
- data = build_data(base_nodes=300, num_houses=80) # BA-Shapes with ground truth
140
- 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
141
144
  model.eval()
142
145
 
143
146
  explainer = Explainer(algorithm=GNNExplainer(epochs=150))
144
147
  expl = explainer.explain_node(data, model, node_idx=42)
145
148
 
146
149
  print(expl.evaluate(metrics=["fidelity", "sparsity"]))
147
- print(expl.evaluate(metrics=["sparsity"], local=True)) # sparsity over the node's k-hop subgraph
148
- 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
149
154
  ```
150
155
 
151
156
  ## Sparsity tuning notes
@@ -278,9 +283,15 @@ meaningless metrics are marked as `skipped`/`None` without aborting the rest:
278
283
  ```python
279
284
  from graph_explain import compare, report_html
280
285
 
281
- results = compare(data, model, node=42, methods=None, # None = all
282
- num_perturbations=5, epochs=200)
283
- 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
284
295
  ```
285
296
 
286
297
  The CLI ships an equivalent subcommand:
@@ -316,7 +327,9 @@ In Python:
316
327
  from graph_explain import Explainer, evaluate_gea_graph
317
328
  from graph_explain.benchmarks.synthetic import build_graph_classification
318
329
 
319
- 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
320
333
  model = ... # GraphGCN (task_level="graph")
321
334
 
322
335
  expl = Explainer(algorithm=GradXInput()).explain_graph(graphs[0], model)
@@ -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