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.
- {graph_explain-0.7.3 → graph_explain-0.8.0}/LICENSE +1 -1
- {graph_explain-0.7.3/src/graph_explain.egg-info → graph_explain-0.8.0}/PKG-INFO +22 -9
- {graph_explain-0.7.3 → graph_explain-0.8.0}/README.md +21 -9
- {graph_explain-0.7.3 → graph_explain-0.8.0}/pyproject.toml +32 -3
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/__init__.py +1 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/benchmarks/synthetic.py +8 -7
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/benchmark.py +4 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/evaluation.py +1 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/explanation.py +1 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/registry.py +6 -4
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/attention/attention.py +1 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/base.py +1 -3
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/grad_x_input.py +1 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/guided_backprop.py +2 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/integrated_gradients.py +2 -2
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/saliency.py +6 -2
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/gnn_explainer.py +11 -4
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/pg_explainer.py +1 -1
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/subgraphx.py +2 -2
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/deeplift.py +3 -3
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/relevance/gnn_lrp.py +3 -3
- {graph_explain-0.7.3 → graph_explain-0.8.0/src/graph_explain.egg-info}/PKG-INFO +22 -9
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/requires.txt +1 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/setup.cfg +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/backends/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/backends/base.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/backends/dgl.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/benchmarks/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/cli.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/explainer.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/core/model_utils.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/baseline/random_baseline.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/counterfactual/counterfactual.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/feature/graph_lime.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/gradient/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/node_mask.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/narration/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/narration/narrator.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/py.typed +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/visualization/__init__.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/visualization/interactive.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain/visualization/static.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/SOURCES.txt +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/dependency_links.txt +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/entry_points.txt +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/src/graph_explain.egg-info/top_level.txt +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_cli.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_core.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_counterfactual.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_dgl_integration.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_narration.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase10.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase2.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase3.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase4.py +0 -0
- {graph_explain-0.7.3 → graph_explain-0.8.0}/tests/test_phase6.py +0 -0
- {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.
|
|
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
|
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
56
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)
|
|
57
60
|
[](https://pypi.org/project/graph-explain/)
|
|
58
61
|
[](https://pypi.org/project/graph-explain/)
|
|
59
62
|
[](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)
|
|
140
|
-
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
|
|
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(
|
|
148
|
-
|
|
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(
|
|
282
|
-
|
|
283
|
-
|
|
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(
|
|
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
|
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
4
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)
|
|
5
7
|
[](https://pypi.org/project/graph-explain/)
|
|
6
8
|
[](https://pypi.org/project/graph-explain/)
|
|
7
9
|
[](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)
|
|
88
|
-
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
|
|
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(
|
|
96
|
-
|
|
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(
|
|
230
|
-
|
|
231
|
-
|
|
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(
|
|
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
|
+
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"
|
|
@@ -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.3 → 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.3 → 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.3 → 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.3 → 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.3 → 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,6 +1,6 @@
|
|
|
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
|
|
@@ -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
|
[](https://github.com/Tzinny-dev/graph-explain/actions/workflows/ci.yml)
|
|
56
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)
|
|
57
60
|
[](https://pypi.org/project/graph-explain/)
|
|
58
61
|
[](https://pypi.org/project/graph-explain/)
|
|
59
62
|
[](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)
|
|
140
|
-
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
|
|
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(
|
|
148
|
-
|
|
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(
|
|
282
|
-
|
|
283
|
-
|
|
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(
|
|
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)
|
|
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.3 → 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.3 → graph_explain-0.8.0}/src/graph_explain/methods/perturbation/__init__.py
RENAMED
|
File without changes
|
{graph_explain-0.7.3 → 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
|