graph-explain 0.7.0__tar.gz → 0.7.2__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- graph_explain-0.7.2/LICENSE +21 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/PKG-INFO +8 -1
- {graph_explain-0.7.0 → graph_explain-0.7.2}/pyproject.toml +12 -2
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/__init__.py +1 -1
- graph_explain-0.7.2/src/graph_explain/narration/narrator.py +246 -0
- graph_explain-0.7.2/src/graph_explain/py.typed +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/PKG-INFO +8 -1
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/SOURCES.txt +2 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/requires.txt +2 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_narration.py +46 -0
- graph_explain-0.7.0/src/graph_explain/narration/narrator.py +0 -185
- {graph_explain-0.7.0 → graph_explain-0.7.2}/README.md +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/setup.cfg +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/backends/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/backends/base.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/backends/dgl.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/benchmarks/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/benchmarks/synthetic.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/cli.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/benchmark.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/evaluation.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/explainer.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/explanation.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/model_utils.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/registry.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/attention/attention.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/base.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/baseline/random_baseline.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/counterfactual/counterfactual.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/feature/graph_lime.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/grad_x_input.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/guided_backprop.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/integrated_gradients.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/saliency.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/gnn_explainer.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/node_mask.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/pg_explainer.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/subgraphx.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/relevance/deeplift.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/relevance/gnn_lrp.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/narration/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/visualization/__init__.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/visualization/interactive.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/visualization/static.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/dependency_links.txt +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/entry_points.txt +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/top_level.txt +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_cli.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_core.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_counterfactual.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_dgl_integration.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase10.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase2.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase3.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase4.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase6.py +0 -0
- {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase8.py +0 -0
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 graph-explain contributors
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
|
@@ -1,9 +1,12 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: graph-explain
|
|
3
|
-
Version: 0.7.
|
|
3
|
+
Version: 0.7.2
|
|
4
4
|
Summary: Explainability library for graph-based models (GNN)
|
|
5
5
|
Author: graph-explain contributors
|
|
6
6
|
License-Expression: MIT
|
|
7
|
+
Project-URL: Homepage, https://github.com/Tzinny-dev/graph-explain
|
|
8
|
+
Project-URL: Repository, https://github.com/Tzinny-dev/graph-explain
|
|
9
|
+
Project-URL: Documentation, https://github.com/Tzinny-dev/graph-explain#readme
|
|
7
10
|
Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
|
|
8
11
|
Classifier: Development Status :: 4 - Beta
|
|
9
12
|
Classifier: Intended Audience :: Science/Research
|
|
@@ -16,9 +19,11 @@ Classifier: Programming Language :: Python :: 3.11
|
|
|
16
19
|
Classifier: Programming Language :: Python :: 3.12
|
|
17
20
|
Requires-Python: >=3.10
|
|
18
21
|
Description-Content-Type: text/markdown
|
|
22
|
+
License-File: LICENSE
|
|
19
23
|
Requires-Dist: numpy>=1.24
|
|
20
24
|
Requires-Dist: networkx>=3.0
|
|
21
25
|
Requires-Dist: matplotlib>=3.6
|
|
26
|
+
Requires-Dist: torch>=2.0
|
|
22
27
|
Provides-Extra: pyg
|
|
23
28
|
Requires-Dist: torch>=2.0; extra == "pyg"
|
|
24
29
|
Requires-Dist: torch-geometric>=2.5; extra == "pyg"
|
|
@@ -31,6 +36,7 @@ Requires-Dist: pyvis>=0.3; extra == "interactive"
|
|
|
31
36
|
Provides-Extra: all
|
|
32
37
|
Requires-Dist: torch>=2.0; extra == "all"
|
|
33
38
|
Requires-Dist: torch-geometric>=2.5; extra == "all"
|
|
39
|
+
Requires-Dist: dgl>=2.0; extra == "all"
|
|
34
40
|
Requires-Dist: plotly>=5.15; extra == "all"
|
|
35
41
|
Requires-Dist: pyvis>=0.3; extra == "all"
|
|
36
42
|
Provides-Extra: dev
|
|
@@ -42,6 +48,7 @@ Requires-Dist: sphinx>=7.2; extra == "docs"
|
|
|
42
48
|
Requires-Dist: sphinx-rtd-theme>=2.0; extra == "docs"
|
|
43
49
|
Provides-Extra: publish
|
|
44
50
|
Requires-Dist: twine>=5.0; extra == "publish"
|
|
51
|
+
Dynamic: license-file
|
|
45
52
|
|
|
46
53
|
# graph-explain
|
|
47
54
|
|
|
@@ -4,11 +4,12 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "graph-explain"
|
|
7
|
-
version = "0.7.
|
|
7
|
+
version = "0.7.2"
|
|
8
8
|
description = "Explainability library for graph-based models (GNN)"
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.10"
|
|
11
11
|
license = "MIT"
|
|
12
|
+
license-files = ["LICENSE"]
|
|
12
13
|
authors = [{ name = "graph-explain contributors" }]
|
|
13
14
|
keywords = [
|
|
14
15
|
"gnn",
|
|
@@ -34,13 +35,19 @@ dependencies = [
|
|
|
34
35
|
"numpy>=1.24",
|
|
35
36
|
"networkx>=3.0",
|
|
36
37
|
"matplotlib>=3.6",
|
|
38
|
+
"torch>=2.0",
|
|
37
39
|
]
|
|
38
40
|
|
|
41
|
+
[project.urls]
|
|
42
|
+
Homepage = "https://github.com/Tzinny-dev/graph-explain"
|
|
43
|
+
Repository = "https://github.com/Tzinny-dev/graph-explain"
|
|
44
|
+
Documentation = "https://github.com/Tzinny-dev/graph-explain#readme"
|
|
45
|
+
|
|
39
46
|
[project.optional-dependencies]
|
|
40
47
|
pyg = ["torch>=2.0", "torch-geometric>=2.5"]
|
|
41
48
|
dgl = ["torch>=2.0", "dgl>=2.0"]
|
|
42
49
|
interactive = ["plotly>=5.15", "pyvis>=0.3"]
|
|
43
|
-
all = ["torch>=2.0", "torch-geometric>=2.5", "plotly>=5.15", "pyvis>=0.3"]
|
|
50
|
+
all = ["torch>=2.0", "torch-geometric>=2.5", "dgl>=2.0", "plotly>=5.15", "pyvis>=0.3"]
|
|
44
51
|
dev = ["pytest>=7.0", "ruff>=0.5", "build>=1.0"]
|
|
45
52
|
docs = ["sphinx>=7.2", "sphinx-rtd-theme>=2.0"]
|
|
46
53
|
publish = ["twine>=5.0"]
|
|
@@ -51,6 +58,9 @@ graph-explain = "graph_explain.cli:main"
|
|
|
51
58
|
[tool.setuptools.packages.find]
|
|
52
59
|
where = ["src"]
|
|
53
60
|
|
|
61
|
+
[tool.setuptools.package-data]
|
|
62
|
+
graph_explain = ["py.typed"]
|
|
63
|
+
|
|
54
64
|
[tool.pytest.ini_options]
|
|
55
65
|
testpaths = ["tests"]
|
|
56
66
|
|
|
@@ -0,0 +1,246 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
from collections.abc import Callable
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
import torch
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _top_values(importance, k: int) -> list[tuple[int, float]]:
|
|
11
|
+
imp = importance.detach().reshape(-1)
|
|
12
|
+
n = int(imp.numel())
|
|
13
|
+
k = max(1, min(int(k), n))
|
|
14
|
+
idx = imp.argsort(descending=True)[:k]
|
|
15
|
+
return [(int(i), float(imp[i])) for i in idx.tolist()]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _data_context(explanation, data: Any | None):
|
|
19
|
+
if data is None:
|
|
20
|
+
data = explanation.metadata.get("backing_data")
|
|
21
|
+
backend = explanation.metadata.get("backend")
|
|
22
|
+
return backend, data
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _labels(explanation, data: Any | None) -> Any | None:
|
|
26
|
+
backend, data = _data_context(explanation, data)
|
|
27
|
+
if backend is None or data is None:
|
|
28
|
+
return None
|
|
29
|
+
try:
|
|
30
|
+
return backend.node_labels(data)
|
|
31
|
+
except Exception: # noqa: BLE001
|
|
32
|
+
return None
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _edge_index(data, backend) -> Any | None:
|
|
36
|
+
if backend is not None and data is not None:
|
|
37
|
+
try:
|
|
38
|
+
return backend.edge_index(data)
|
|
39
|
+
except Exception: # noqa: BLE001
|
|
40
|
+
return None
|
|
41
|
+
return None
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
_SUPPORTED_LANGS = ("es", "en")
|
|
45
|
+
|
|
46
|
+
_TEMPLATES: dict[str, dict[str, str]] = {
|
|
47
|
+
"es": {
|
|
48
|
+
"graph_head": "Explicación a nivel de grafo.",
|
|
49
|
+
"node_head": "Explicación del nodo {node}.",
|
|
50
|
+
"target": " La clase objetivo es {target}.",
|
|
51
|
+
"correct": " La predicción del modelo es correcta.",
|
|
52
|
+
"incorrect": " La predicción del modelo difiere de la etiqueta real.",
|
|
53
|
+
"nodes": "Los nodos más relevantes son {nodes}.",
|
|
54
|
+
"edges": " Las aristas más relevantes son {edges}.",
|
|
55
|
+
"node_item": "nodo {i} (importancia {v:.3f})",
|
|
56
|
+
"edge_item": "arista {u}-{v} ({w:.3f})",
|
|
57
|
+
"edge_item_float": "arista con relevancia {v:.3f}",
|
|
58
|
+
"cf_intro": " Se necesitaron {n} cambios (aristas/features eliminadas) "
|
|
59
|
+
"para cambiar la predicción de {change}.",
|
|
60
|
+
"cf_change": "la clase {orig} a la clase {new}",
|
|
61
|
+
"cf_change_solo": "la clase {orig}",
|
|
62
|
+
"nodata": "Sin datos suficientes para describir la explicación.",
|
|
63
|
+
"prompt": "Eres un asistente que explica predicciones de GNNs en lenguaje "
|
|
64
|
+
"natural. Dado este resumen de una explicación (JSON), escribe un párrafo "
|
|
65
|
+
"breve en español (2-4 oraciones) describiendo qué hace el modelo y qué "
|
|
66
|
+
"evidencia respalda su predicción. Resumen:\n",
|
|
67
|
+
"llm_fallback": "[LLM no disponible: {exc}]",
|
|
68
|
+
},
|
|
69
|
+
"en": {
|
|
70
|
+
"graph_head": "Graph-level explanation.",
|
|
71
|
+
"node_head": "Explanation of node {node}.",
|
|
72
|
+
"target": " The target class is {target}.",
|
|
73
|
+
"correct": " The model prediction is correct.",
|
|
74
|
+
"incorrect": " The model prediction differs from the true label.",
|
|
75
|
+
"nodes": "The most relevant nodes are {nodes}.",
|
|
76
|
+
"edges": " The most relevant edges are {edges}.",
|
|
77
|
+
"node_item": "node {i} (importance {v:.3f})",
|
|
78
|
+
"edge_item": "edge {u}-{v} ({w:.3f})",
|
|
79
|
+
"edge_item_float": "edge with relevance {v:.3f}",
|
|
80
|
+
"cf_intro": " {n} changes (removed edges/features) were needed to change the "
|
|
81
|
+
"prediction from {change}.",
|
|
82
|
+
"cf_change": "class {orig} to class {new}",
|
|
83
|
+
"cf_change_solo": "class {orig}",
|
|
84
|
+
"nodata": "Not enough data to describe the explanation.",
|
|
85
|
+
"prompt": "You are an assistant that explains GNN predictions in natural "
|
|
86
|
+
"language. Given this JSON summary of an explanation, write a brief "
|
|
87
|
+
"paragraph (2-4 sentences) in English describing what the model does and "
|
|
88
|
+
"what evidence supports its prediction. Summary:\n",
|
|
89
|
+
"llm_fallback": "[LLM unavailable: {exc}]",
|
|
90
|
+
},
|
|
91
|
+
}
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _templates(lang: str) -> dict[str, str]:
|
|
95
|
+
if lang not in _SUPPORTED_LANGS:
|
|
96
|
+
raise ValueError(f"lang must be one of {_SUPPORTED_LANGS}, got {lang!r}")
|
|
97
|
+
return _TEMPLATES[lang]
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def summarize(explanation, data: Any | None = None, top_k: int = 5) -> dict[str, Any]:
|
|
101
|
+
"""Structured summary of an explanation (for narration or JSON)."""
|
|
102
|
+
backend, data = _data_context(explanation, data)
|
|
103
|
+
node = explanation.node_idx
|
|
104
|
+
target = explanation.target_class
|
|
105
|
+
pred = None
|
|
106
|
+
if explanation.prediction_original is not None:
|
|
107
|
+
p = explanation.prediction_original
|
|
108
|
+
if torch.is_tensor(p):
|
|
109
|
+
pred = int(p.reshape(-1).argmax().item())
|
|
110
|
+
labels = _labels(explanation, data)
|
|
111
|
+
|
|
112
|
+
true_label = None
|
|
113
|
+
if labels is not None and node is not None:
|
|
114
|
+
try:
|
|
115
|
+
true_label = int(labels[node].item())
|
|
116
|
+
except Exception: # noqa: BLE001
|
|
117
|
+
true_label = None
|
|
118
|
+
|
|
119
|
+
summary: dict[str, Any] = {
|
|
120
|
+
"node": None if node is None else int(node),
|
|
121
|
+
"target_class": target,
|
|
122
|
+
"predicted_class": pred,
|
|
123
|
+
"true_class": true_label,
|
|
124
|
+
"correct": (
|
|
125
|
+
None if pred is None or true_label is None else bool(pred == true_label)
|
|
126
|
+
),
|
|
127
|
+
"important_nodes": (
|
|
128
|
+
_top_values(explanation.node_importance, top_k)
|
|
129
|
+
if explanation.node_importance is not None
|
|
130
|
+
else []
|
|
131
|
+
),
|
|
132
|
+
"important_edges": [],
|
|
133
|
+
"counterfactual": bool(explanation.metadata.get("counterfactual", False)),
|
|
134
|
+
}
|
|
135
|
+
if explanation.edge_importance is not None:
|
|
136
|
+
ei = _edge_index(data, backend)
|
|
137
|
+
top = _top_values(explanation.edge_importance, top_k)
|
|
138
|
+
if ei is None:
|
|
139
|
+
summary["important_edges"] = [v for _, v in top]
|
|
140
|
+
else:
|
|
141
|
+
summary["important_edges"] = [
|
|
142
|
+
(int(ei[0, i]), int(ei[1, i]), v) for i, v in top
|
|
143
|
+
]
|
|
144
|
+
if summary["counterfactual"]:
|
|
145
|
+
summary["original_class"] = explanation.metadata.get("original_class")
|
|
146
|
+
return summary
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
def describe(
|
|
150
|
+
explanation, data: Any | None = None, top_k: int = 5, lang: str = "es"
|
|
151
|
+
) -> str:
|
|
152
|
+
"""Deterministic template-based narration of an explanation.
|
|
153
|
+
|
|
154
|
+
Args:
|
|
155
|
+
lang: Template language, ``"es"`` (default) or ``"en"``.
|
|
156
|
+
"""
|
|
157
|
+
_T = _templates(lang)
|
|
158
|
+
s = summarize(explanation, data, top_k)
|
|
159
|
+
node = s["node"]
|
|
160
|
+
target = (
|
|
161
|
+
s["target_class"] if s["target_class"] is not None else s["predicted_class"]
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
if node is None:
|
|
165
|
+
head = _T["graph_head"]
|
|
166
|
+
else:
|
|
167
|
+
head = _T["node_head"].format(node=node)
|
|
168
|
+
if target is not None:
|
|
169
|
+
head += _T["target"].format(target=target)
|
|
170
|
+
if s["correct"] is True:
|
|
171
|
+
head += _T["correct"]
|
|
172
|
+
elif s["correct"] is False:
|
|
173
|
+
head += _T["incorrect"]
|
|
174
|
+
|
|
175
|
+
nodes_txt = ", ".join(
|
|
176
|
+
_T["node_item"].format(i=i, v=v) for i, v in s["important_nodes"]
|
|
177
|
+
)
|
|
178
|
+
tail = _T["nodes"].format(nodes=nodes_txt) if nodes_txt else ""
|
|
179
|
+
|
|
180
|
+
if s["important_edges"]:
|
|
181
|
+
pieces = []
|
|
182
|
+
for e in s["important_edges"]:
|
|
183
|
+
if len(e) == 3:
|
|
184
|
+
u, v, w = e
|
|
185
|
+
pieces.append(_T["edge_item"].format(u=u, v=v, w=w))
|
|
186
|
+
else:
|
|
187
|
+
pieces.append(_T["edge_item_float"].format(v=float(e)))
|
|
188
|
+
tail += _T["edges"].format(edges=", ".join(pieces))
|
|
189
|
+
|
|
190
|
+
if s["counterfactual"]:
|
|
191
|
+
n = len(s["important_edges"])
|
|
192
|
+
new_class = s["predicted_class"]
|
|
193
|
+
if new_class is not None:
|
|
194
|
+
change = _T["cf_change"].format(orig=s["original_class"], new=new_class)
|
|
195
|
+
else:
|
|
196
|
+
change = _T["cf_change_solo"].format(orig=s["original_class"])
|
|
197
|
+
tail += _T["cf_intro"].format(n=n, change=change)
|
|
198
|
+
if tail:
|
|
199
|
+
head += " " + tail.strip()
|
|
200
|
+
return head.strip() or _T["nodata"]
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
def _prompt(summary: dict[str, Any], lang: str = "es") -> str:
|
|
204
|
+
_T = _templates(lang)
|
|
205
|
+
return _T["prompt"] + json.dumps(summary, ensure_ascii=False, indent=2)
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def narrate(
|
|
209
|
+
explanation,
|
|
210
|
+
llm: Callable[[str], str] | None = None,
|
|
211
|
+
data: Any | None = None,
|
|
212
|
+
top_k: int = 5,
|
|
213
|
+
lang: str = "es",
|
|
214
|
+
) -> str:
|
|
215
|
+
"""Narrates an explanation. With `llm` (a `prompt -> text` callable) it uses the
|
|
216
|
+
generative model's output; otherwise it falls back to deterministic
|
|
217
|
+
template-based narration."""
|
|
218
|
+
_T = _templates(lang)
|
|
219
|
+
summary = summarize(explanation, data, top_k)
|
|
220
|
+
deterministic = describe(explanation, data, top_k, lang=lang)
|
|
221
|
+
if llm is None:
|
|
222
|
+
return deterministic
|
|
223
|
+
try:
|
|
224
|
+
return llm(_prompt(summary, lang)).strip()
|
|
225
|
+
except Exception as exc: # noqa: BLE001
|
|
226
|
+
return f"{deterministic}\n\n{_T['llm_fallback'].format(exc=exc)}"
|
|
227
|
+
|
|
228
|
+
|
|
229
|
+
class Narrator:
|
|
230
|
+
"""Reusable narrator; lets you inject the LLM just once."""
|
|
231
|
+
|
|
232
|
+
def __init__(
|
|
233
|
+
self,
|
|
234
|
+
llm: Callable[[str], str] | None = None,
|
|
235
|
+
top_k: int = 5,
|
|
236
|
+
lang: str = "es",
|
|
237
|
+
):
|
|
238
|
+
self.llm = llm
|
|
239
|
+
self.top_k = top_k
|
|
240
|
+
self.lang = lang
|
|
241
|
+
|
|
242
|
+
def describe(self, explanation, data: Any | None = None) -> str:
|
|
243
|
+
return describe(explanation, data, self.top_k, lang=self.lang)
|
|
244
|
+
|
|
245
|
+
def narrate(self, explanation, data: Any | None = None) -> str:
|
|
246
|
+
return narrate(explanation, self.llm, data, self.top_k, lang=self.lang)
|
|
File without changes
|
|
@@ -1,9 +1,12 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: graph-explain
|
|
3
|
-
Version: 0.7.
|
|
3
|
+
Version: 0.7.2
|
|
4
4
|
Summary: Explainability library for graph-based models (GNN)
|
|
5
5
|
Author: graph-explain contributors
|
|
6
6
|
License-Expression: MIT
|
|
7
|
+
Project-URL: Homepage, https://github.com/Tzinny-dev/graph-explain
|
|
8
|
+
Project-URL: Repository, https://github.com/Tzinny-dev/graph-explain
|
|
9
|
+
Project-URL: Documentation, https://github.com/Tzinny-dev/graph-explain#readme
|
|
7
10
|
Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
|
|
8
11
|
Classifier: Development Status :: 4 - Beta
|
|
9
12
|
Classifier: Intended Audience :: Science/Research
|
|
@@ -16,9 +19,11 @@ Classifier: Programming Language :: Python :: 3.11
|
|
|
16
19
|
Classifier: Programming Language :: Python :: 3.12
|
|
17
20
|
Requires-Python: >=3.10
|
|
18
21
|
Description-Content-Type: text/markdown
|
|
22
|
+
License-File: LICENSE
|
|
19
23
|
Requires-Dist: numpy>=1.24
|
|
20
24
|
Requires-Dist: networkx>=3.0
|
|
21
25
|
Requires-Dist: matplotlib>=3.6
|
|
26
|
+
Requires-Dist: torch>=2.0
|
|
22
27
|
Provides-Extra: pyg
|
|
23
28
|
Requires-Dist: torch>=2.0; extra == "pyg"
|
|
24
29
|
Requires-Dist: torch-geometric>=2.5; extra == "pyg"
|
|
@@ -31,6 +36,7 @@ Requires-Dist: pyvis>=0.3; extra == "interactive"
|
|
|
31
36
|
Provides-Extra: all
|
|
32
37
|
Requires-Dist: torch>=2.0; extra == "all"
|
|
33
38
|
Requires-Dist: torch-geometric>=2.5; extra == "all"
|
|
39
|
+
Requires-Dist: dgl>=2.0; extra == "all"
|
|
34
40
|
Requires-Dist: plotly>=5.15; extra == "all"
|
|
35
41
|
Requires-Dist: pyvis>=0.3; extra == "all"
|
|
36
42
|
Provides-Extra: dev
|
|
@@ -42,6 +48,7 @@ Requires-Dist: sphinx>=7.2; extra == "docs"
|
|
|
42
48
|
Requires-Dist: sphinx-rtd-theme>=2.0; extra == "docs"
|
|
43
49
|
Provides-Extra: publish
|
|
44
50
|
Requires-Dist: twine>=5.0; extra == "publish"
|
|
51
|
+
Dynamic: license-file
|
|
45
52
|
|
|
46
53
|
# graph-explain
|
|
47
54
|
|
|
@@ -74,3 +74,49 @@ class TestNarration:
|
|
|
74
74
|
_, expl = self._setup()
|
|
75
75
|
narrator = Narrator(llm=None)
|
|
76
76
|
assert narrator.describe(expl) == describe(expl)
|
|
77
|
+
|
|
78
|
+
def test_describe_english(self):
|
|
79
|
+
data, expl = self._setup()
|
|
80
|
+
text = describe(expl, data=data, lang="en")
|
|
81
|
+
assert "node" in text
|
|
82
|
+
assert "class" in text
|
|
83
|
+
assert "most relevant" in text
|
|
84
|
+
assert text == describe(expl, data=data, lang="en") # determinista
|
|
85
|
+
|
|
86
|
+
def test_english_vs_spanish(self):
|
|
87
|
+
data, expl = self._setup()
|
|
88
|
+
es = describe(expl, data=data, lang="es")
|
|
89
|
+
en = describe(expl, data=data, lang="en")
|
|
90
|
+
assert es != en
|
|
91
|
+
assert "nodo" in es and "node" in en
|
|
92
|
+
|
|
93
|
+
def test_narrate_with_llm_english(self):
|
|
94
|
+
_, expl = self._setup()
|
|
95
|
+
|
|
96
|
+
def fake_llm(prompt):
|
|
97
|
+
return "CUSTOM NARRATION"
|
|
98
|
+
|
|
99
|
+
text = narrate(expl, llm=fake_llm, lang="en")
|
|
100
|
+
assert text == "CUSTOM NARRATION"
|
|
101
|
+
|
|
102
|
+
def test_narrate_fallback_english(self):
|
|
103
|
+
_, expl = self._setup()
|
|
104
|
+
|
|
105
|
+
def broken_llm(_prompt):
|
|
106
|
+
raise RuntimeError("no api")
|
|
107
|
+
|
|
108
|
+
text = narrate(expl, llm=broken_llm, lang="en")
|
|
109
|
+
assert "LLM unavailable" in text
|
|
110
|
+
assert "Explanation of node" in text
|
|
111
|
+
|
|
112
|
+
def test_invalid_lang(self):
|
|
113
|
+
import pytest
|
|
114
|
+
|
|
115
|
+
_, expl = self._setup()
|
|
116
|
+
with pytest.raises(ValueError):
|
|
117
|
+
describe(expl, lang="fr")
|
|
118
|
+
|
|
119
|
+
def test_narrator_class_english(self):
|
|
120
|
+
_, expl = self._setup()
|
|
121
|
+
narrator = Narrator(llm=None, lang="en")
|
|
122
|
+
assert narrator.describe(expl) == describe(expl, lang="en")
|
|
@@ -1,185 +0,0 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
|
|
3
|
-
import json
|
|
4
|
-
from collections.abc import Callable
|
|
5
|
-
from typing import Any
|
|
6
|
-
|
|
7
|
-
import torch
|
|
8
|
-
|
|
9
|
-
|
|
10
|
-
def _top_values(importance, k: int) -> list[tuple[int, float]]:
|
|
11
|
-
imp = importance.detach().reshape(-1)
|
|
12
|
-
n = int(imp.numel())
|
|
13
|
-
k = max(1, min(int(k), n))
|
|
14
|
-
idx = imp.argsort(descending=True)[:k]
|
|
15
|
-
return [(int(i), float(imp[i])) for i in idx.tolist()]
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
def _data_context(explanation, data: Any | None):
|
|
19
|
-
if data is None:
|
|
20
|
-
data = explanation.metadata.get("backing_data")
|
|
21
|
-
backend = explanation.metadata.get("backend")
|
|
22
|
-
return backend, data
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
def _labels(explanation, data: Any | None) -> Any | None:
|
|
26
|
-
backend, data = _data_context(explanation, data)
|
|
27
|
-
if backend is None or data is None:
|
|
28
|
-
return None
|
|
29
|
-
try:
|
|
30
|
-
return backend.node_labels(data)
|
|
31
|
-
except Exception: # noqa: BLE001
|
|
32
|
-
return None
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
def _edge_index(data, backend) -> Any | None:
|
|
36
|
-
if backend is not None and data is not None:
|
|
37
|
-
try:
|
|
38
|
-
return backend.edge_index(data)
|
|
39
|
-
except Exception: # noqa: BLE001
|
|
40
|
-
return None
|
|
41
|
-
return None
|
|
42
|
-
|
|
43
|
-
|
|
44
|
-
def summarize(explanation, data: Any | None = None, top_k: int = 5) -> dict[str, Any]:
|
|
45
|
-
"""Structured summary of an explanation (for narration or JSON)."""
|
|
46
|
-
backend, data = _data_context(explanation, data)
|
|
47
|
-
node = explanation.node_idx
|
|
48
|
-
target = explanation.target_class
|
|
49
|
-
pred = None
|
|
50
|
-
if explanation.prediction_original is not None:
|
|
51
|
-
p = explanation.prediction_original
|
|
52
|
-
if torch.is_tensor(p):
|
|
53
|
-
pred = int(p.reshape(-1).argmax().item())
|
|
54
|
-
labels = _labels(explanation, data)
|
|
55
|
-
|
|
56
|
-
true_label = None
|
|
57
|
-
if labels is not None and node is not None:
|
|
58
|
-
try:
|
|
59
|
-
true_label = int(labels[node].item())
|
|
60
|
-
except Exception: # noqa: BLE001
|
|
61
|
-
true_label = None
|
|
62
|
-
|
|
63
|
-
summary: dict[str, Any] = {
|
|
64
|
-
"node": None if node is None else int(node),
|
|
65
|
-
"target_class": target,
|
|
66
|
-
"predicted_class": pred,
|
|
67
|
-
"true_class": true_label,
|
|
68
|
-
"correct": (
|
|
69
|
-
None if pred is None or true_label is None else bool(pred == true_label)
|
|
70
|
-
),
|
|
71
|
-
"important_nodes": (
|
|
72
|
-
_top_values(explanation.node_importance, top_k)
|
|
73
|
-
if explanation.node_importance is not None
|
|
74
|
-
else []
|
|
75
|
-
),
|
|
76
|
-
"important_edges": [],
|
|
77
|
-
"counterfactual": bool(explanation.metadata.get("counterfactual", False)),
|
|
78
|
-
}
|
|
79
|
-
if explanation.edge_importance is not None:
|
|
80
|
-
ei = _edge_index(data, backend)
|
|
81
|
-
top = _top_values(explanation.edge_importance, top_k)
|
|
82
|
-
if ei is None:
|
|
83
|
-
summary["important_edges"] = [v for _, v in top]
|
|
84
|
-
else:
|
|
85
|
-
summary["important_edges"] = [
|
|
86
|
-
(int(ei[0, i]), int(ei[1, i]), v) for i, v in top
|
|
87
|
-
]
|
|
88
|
-
if summary["counterfactual"]:
|
|
89
|
-
summary["original_class"] = explanation.metadata.get("original_class")
|
|
90
|
-
return summary
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
def describe(explanation, data: Any | None = None, top_k: int = 5) -> str:
|
|
94
|
-
"""Deterministic template-based narration of an explanation (Spanish by default)."""
|
|
95
|
-
s = summarize(explanation, data, top_k)
|
|
96
|
-
node = s["node"]
|
|
97
|
-
target = (
|
|
98
|
-
s["target_class"] if s["target_class"] is not None else s["predicted_class"]
|
|
99
|
-
)
|
|
100
|
-
|
|
101
|
-
if node is None:
|
|
102
|
-
head = "Explicación a nivel de grafo."
|
|
103
|
-
else:
|
|
104
|
-
head = f"Explicación del nodo {node}."
|
|
105
|
-
if target is not None:
|
|
106
|
-
head += f" La clase objetivo es {target}."
|
|
107
|
-
if s["correct"] is True:
|
|
108
|
-
head += " La predicción del modelo es correcta."
|
|
109
|
-
elif s["correct"] is False:
|
|
110
|
-
head += " La predicción del modelo difiere de la etiqueta real."
|
|
111
|
-
|
|
112
|
-
nodes_txt = ", ".join(
|
|
113
|
-
f"nodo {i} (importancia {v:.3f})" for i, v in s["important_nodes"]
|
|
114
|
-
)
|
|
115
|
-
tail = f"Los nodos más relevantes son {nodes_txt}." if nodes_txt else ""
|
|
116
|
-
|
|
117
|
-
if s["important_edges"]:
|
|
118
|
-
pieces = []
|
|
119
|
-
for e in s["important_edges"]:
|
|
120
|
-
if len(e) == 3:
|
|
121
|
-
u, v, w = e
|
|
122
|
-
pieces.append(f"arista {u}-{v} ({w:.3f})")
|
|
123
|
-
else:
|
|
124
|
-
pieces.append(f"arista con relevancia {float(e):.3f}")
|
|
125
|
-
tail += " Las aristas más relevantes son " + ", ".join(pieces) + "."
|
|
126
|
-
|
|
127
|
-
if s["counterfactual"]:
|
|
128
|
-
n = len(s["important_edges"])
|
|
129
|
-
new_class = s["predicted_class"]
|
|
130
|
-
change = (
|
|
131
|
-
f"la clase {s['original_class']} a la clase {new_class}"
|
|
132
|
-
if new_class is not None
|
|
133
|
-
else f"la clase {s['original_class']}"
|
|
134
|
-
)
|
|
135
|
-
tail += (
|
|
136
|
-
f" Se necesitaron {n} cambios"
|
|
137
|
-
f" (aristas/features eliminadas) para cambiar la predicción de"
|
|
138
|
-
f" {change}."
|
|
139
|
-
)
|
|
140
|
-
if tail:
|
|
141
|
-
head += " " + tail.strip()
|
|
142
|
-
return head.strip() or "Sin datos suficientes para describir la explicación."
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
def _prompt(summary: dict[str, Any]) -> str:
|
|
146
|
-
return (
|
|
147
|
-
"Eres un asistente que explica predicciones de GNNs en lenguaje natural. "
|
|
148
|
-
"Dado este resumen de una explicación (JSON), escribe un párrafo breve "
|
|
149
|
-
"en español (2-4 oraciones) describiendo qué hace el modelo y qué "
|
|
150
|
-
"evidencia respalda su predicción. Resumen:\n"
|
|
151
|
-
+ json.dumps(summary, ensure_ascii=False, indent=2)
|
|
152
|
-
)
|
|
153
|
-
|
|
154
|
-
|
|
155
|
-
def narrate(
|
|
156
|
-
explanation,
|
|
157
|
-
llm: Callable[[str], str] | None = None,
|
|
158
|
-
data: Any | None = None,
|
|
159
|
-
top_k: int = 5,
|
|
160
|
-
) -> str:
|
|
161
|
-
"""Narrates an explanation. With `llm` (a `prompt -> text` callable) it uses the
|
|
162
|
-
generative model's output; otherwise it falls back to deterministic
|
|
163
|
-
template-based narration."""
|
|
164
|
-
summary = summarize(explanation, data, top_k)
|
|
165
|
-
deterministic = describe(explanation, data, top_k)
|
|
166
|
-
if llm is None:
|
|
167
|
-
return deterministic
|
|
168
|
-
try:
|
|
169
|
-
return llm(_prompt(summary)).strip()
|
|
170
|
-
except Exception as exc: # noqa: BLE001
|
|
171
|
-
return f"{deterministic}\n\n[LLM no disponible: {exc}]"
|
|
172
|
-
|
|
173
|
-
|
|
174
|
-
class Narrator:
|
|
175
|
-
"""Reusable narrator; lets you inject the LLM just once."""
|
|
176
|
-
|
|
177
|
-
def __init__(self, llm: Callable[[str], str] | None = None, top_k: int = 5):
|
|
178
|
-
self.llm = llm
|
|
179
|
-
self.top_k = top_k
|
|
180
|
-
|
|
181
|
-
def describe(self, explanation, data: Any | None = None) -> str:
|
|
182
|
-
return describe(explanation, data, self.top_k)
|
|
183
|
-
|
|
184
|
-
def narrate(self, explanation, data: Any | None = None) -> str:
|
|
185
|
-
return narrate(explanation, self.llm, data, self.top_k)
|
|
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
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/attention/attention.py
RENAMED
|
File without changes
|
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/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.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/grad_x_input.py
RENAMED
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/guided_backprop.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/__init__.py
RENAMED
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/gnn_explainer.py
RENAMED
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/node_mask.py
RENAMED
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/pg_explainer.py
RENAMED
|
File without changes
|
{graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/subgraphx.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
|