graph-explain 0.7.0__py3-none-any.whl
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/__init__.py +79 -0
- graph_explain/backends/__init__.py +4 -0
- graph_explain/backends/base.py +103 -0
- graph_explain/backends/dgl.py +121 -0
- graph_explain/benchmarks/__init__.py +3 -0
- graph_explain/benchmarks/synthetic.py +246 -0
- graph_explain/cli.py +459 -0
- graph_explain/core/__init__.py +14 -0
- graph_explain/core/benchmark.py +284 -0
- graph_explain/core/evaluation.py +391 -0
- graph_explain/core/explainer.py +83 -0
- graph_explain/core/explanation.py +72 -0
- graph_explain/core/model_utils.py +44 -0
- graph_explain/core/registry.py +55 -0
- graph_explain/methods/__init__.py +39 -0
- graph_explain/methods/attention/attention.py +147 -0
- graph_explain/methods/base.py +25 -0
- graph_explain/methods/baseline/random_baseline.py +78 -0
- graph_explain/methods/counterfactual/counterfactual.py +304 -0
- graph_explain/methods/feature/graph_lime.py +141 -0
- graph_explain/methods/gradient/__init__.py +0 -0
- graph_explain/methods/gradient/grad_x_input.py +110 -0
- graph_explain/methods/gradient/guided_backprop.py +117 -0
- graph_explain/methods/gradient/integrated_gradients.py +115 -0
- graph_explain/methods/gradient/saliency.py +93 -0
- graph_explain/methods/perturbation/__init__.py +0 -0
- graph_explain/methods/perturbation/gnn_explainer.py +265 -0
- graph_explain/methods/perturbation/node_mask.py +136 -0
- graph_explain/methods/perturbation/pg_explainer.py +162 -0
- graph_explain/methods/perturbation/subgraphx.py +393 -0
- graph_explain/methods/relevance/deeplift.py +262 -0
- graph_explain/methods/relevance/gnn_lrp.py +219 -0
- graph_explain/narration/__init__.py +3 -0
- graph_explain/narration/narrator.py +185 -0
- graph_explain/visualization/__init__.py +4 -0
- graph_explain/visualization/interactive.py +73 -0
- graph_explain/visualization/static.py +90 -0
- graph_explain-0.7.0.dist-info/METADATA +332 -0
- graph_explain-0.7.0.dist-info/RECORD +42 -0
- graph_explain-0.7.0.dist-info/WHEEL +5 -0
- graph_explain-0.7.0.dist-info/entry_points.txt +2 -0
- graph_explain-0.7.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,83 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
from ..backends.base import Backend, get_backend
|
|
8
|
+
from ..methods.base import ExplanationAlgorithm
|
|
9
|
+
from .explanation import Explanation
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class Explainer:
|
|
13
|
+
def __init__(
|
|
14
|
+
self,
|
|
15
|
+
algorithm: ExplanationAlgorithm,
|
|
16
|
+
backend: Backend | str = "pyg",
|
|
17
|
+
node_mask_type: str | None = "attributes",
|
|
18
|
+
edge_mask_type: str | None = "object",
|
|
19
|
+
mask_threshold: float = 0.5,
|
|
20
|
+
explanation_type: str = "model",
|
|
21
|
+
**kwargs,
|
|
22
|
+
):
|
|
23
|
+
if isinstance(backend, str):
|
|
24
|
+
backend = get_backend(backend)
|
|
25
|
+
self.backend = backend
|
|
26
|
+
self.algorithm = algorithm
|
|
27
|
+
self.node_mask_type = node_mask_type
|
|
28
|
+
self.edge_mask_type = edge_mask_type
|
|
29
|
+
self.mask_threshold = mask_threshold
|
|
30
|
+
self.explanation_type = explanation_type
|
|
31
|
+
self._extra = kwargs
|
|
32
|
+
|
|
33
|
+
def explain(
|
|
34
|
+
self,
|
|
35
|
+
data: Any,
|
|
36
|
+
model: Any,
|
|
37
|
+
index: int | list[int] | torch.Tensor | None = None,
|
|
38
|
+
target_class: int | None = None,
|
|
39
|
+
**kwargs,
|
|
40
|
+
) -> Explanation:
|
|
41
|
+
explanation = self.algorithm.explain(
|
|
42
|
+
backend=self.backend,
|
|
43
|
+
model=model,
|
|
44
|
+
data=data,
|
|
45
|
+
index=index,
|
|
46
|
+
target_class=target_class,
|
|
47
|
+
node_mask_type=self.node_mask_type,
|
|
48
|
+
edge_mask_type=self.edge_mask_type,
|
|
49
|
+
**kwargs,
|
|
50
|
+
)
|
|
51
|
+
explanation.node_idx = index
|
|
52
|
+
explanation.mask_threshold = self.mask_threshold
|
|
53
|
+
explanation.metadata.setdefault("backend", self.backend)
|
|
54
|
+
explanation.metadata.setdefault("backing_data", data)
|
|
55
|
+
return explanation
|
|
56
|
+
|
|
57
|
+
def explain_node(
|
|
58
|
+
self,
|
|
59
|
+
data: Any,
|
|
60
|
+
model: Any,
|
|
61
|
+
node_idx: int,
|
|
62
|
+
target_class: int | None = None,
|
|
63
|
+
**kwargs,
|
|
64
|
+
) -> Explanation:
|
|
65
|
+
return self.explain(
|
|
66
|
+
data, model, index=node_idx, target_class=target_class, **kwargs
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
def explain_graph(
|
|
70
|
+
self,
|
|
71
|
+
data: Any,
|
|
72
|
+
model: Any,
|
|
73
|
+
target_class: int | None = None,
|
|
74
|
+
**kwargs,
|
|
75
|
+
) -> Explanation:
|
|
76
|
+
return self.explain(
|
|
77
|
+
data, model, index=None, target_class=target_class, **kwargs
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
def __call__(
|
|
81
|
+
self, data: Any, model: Any, index: int | list[int] | None = None, **kwargs
|
|
82
|
+
):
|
|
83
|
+
return self.explain(data, model, index=index, **kwargs)
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import networkx as nx
|
|
7
|
+
|
|
8
|
+
from ..core.evaluation import evaluate_fidelity, evaluate_sparsity
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@dataclass
|
|
12
|
+
class Explanation:
|
|
13
|
+
node_importance: Any | None = None
|
|
14
|
+
edge_importance: Any | None = None
|
|
15
|
+
feature_importance: Any | None = None
|
|
16
|
+
subgraph: Any | None = None
|
|
17
|
+
prediction_original: Any | None = None
|
|
18
|
+
prediction_explanation: Any | None = None
|
|
19
|
+
node_idx: int | None = None
|
|
20
|
+
target_class: int | None = None
|
|
21
|
+
mask_threshold: float = 0.5
|
|
22
|
+
metadata: dict = field(default_factory=dict)
|
|
23
|
+
|
|
24
|
+
def evaluate(self, metrics: list[str] | None = None, **kwargs) -> dict:
|
|
25
|
+
metrics = metrics or ["fidelity", "sparsity"]
|
|
26
|
+
results: dict = {}
|
|
27
|
+
for name in metrics:
|
|
28
|
+
name = name.lower()
|
|
29
|
+
if name == "fidelity":
|
|
30
|
+
results[name] = evaluate_fidelity(self)
|
|
31
|
+
elif name in ("sparsity", "sparsity_ratio"):
|
|
32
|
+
results[name] = evaluate_sparsity(self, **kwargs)
|
|
33
|
+
else:
|
|
34
|
+
raise ValueError(f"Métrica desconocida: {name}")
|
|
35
|
+
return results
|
|
36
|
+
|
|
37
|
+
def to_networkx(self, threshold: float | None = None) -> nx.Graph:
|
|
38
|
+
threshold = threshold if threshold is not None else self.mask_threshold
|
|
39
|
+
backend = self.metadata.get("backend")
|
|
40
|
+
data = self.metadata.get("backing_data")
|
|
41
|
+
if backend is None or data is None:
|
|
42
|
+
raise ValueError(
|
|
43
|
+
"La explicación se creó sin backend/backing_data en metadata"
|
|
44
|
+
)
|
|
45
|
+
G = nx.Graph()
|
|
46
|
+
keep_edges = []
|
|
47
|
+
if self.edge_importance is not None:
|
|
48
|
+
edge_index = backend.edge_index(data)
|
|
49
|
+
for e in range(self.edge_importance.shape[0]):
|
|
50
|
+
if float(self.edge_importance[e]) >= threshold:
|
|
51
|
+
u = int(edge_index[0, e])
|
|
52
|
+
v = int(edge_index[1, e])
|
|
53
|
+
keep_edges.append((u, v, float(self.edge_importance[e])))
|
|
54
|
+
if keep_edges:
|
|
55
|
+
G.add_weighted_edges_from(keep_edges)
|
|
56
|
+
nodes = {n for e in keep_edges for n in e[:2]}
|
|
57
|
+
if self.node_importance is not None:
|
|
58
|
+
for n in range(self.node_importance.shape[0]):
|
|
59
|
+
if float(self.node_importance[n]) >= threshold:
|
|
60
|
+
nodes.add(n)
|
|
61
|
+
G.add_nodes_from(nodes)
|
|
62
|
+
return G
|
|
63
|
+
|
|
64
|
+
def __repr__(self) -> str:
|
|
65
|
+
parts = []
|
|
66
|
+
if self.node_importance is not None:
|
|
67
|
+
parts.append(f"node_importance={tuple(self.node_importance.shape)}")
|
|
68
|
+
if self.edge_importance is not None:
|
|
69
|
+
parts.append(f"edge_importance={tuple(self.edge_importance.shape)}")
|
|
70
|
+
if self.feature_importance is not None:
|
|
71
|
+
parts.append(f"feature_importance={tuple(self.feature_importance.shape)}")
|
|
72
|
+
return f"Explanation({', '.join(parts)})"
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def capture_node_embeddings(
|
|
9
|
+
model: Any,
|
|
10
|
+
backend: Any,
|
|
11
|
+
data: Any,
|
|
12
|
+
layer_index: int = -2,
|
|
13
|
+
) -> torch.Tensor:
|
|
14
|
+
children = list(model.children())
|
|
15
|
+
modules = list(model.modules())
|
|
16
|
+
if not children:
|
|
17
|
+
raise ValueError("El modelo no tiene submódulos para capturar embeddings")
|
|
18
|
+
target = children[layer_index] if layer_index < 0 else modules[layer_index]
|
|
19
|
+
captured: dict[str, torch.Tensor] = {}
|
|
20
|
+
|
|
21
|
+
def hook(_module, _input, output):
|
|
22
|
+
if isinstance(output, tuple):
|
|
23
|
+
output = output[0]
|
|
24
|
+
captured["emb"] = output.detach()
|
|
25
|
+
|
|
26
|
+
handle = target.register_forward_hook(hook)
|
|
27
|
+
try:
|
|
28
|
+
x = backend.node_features(data)
|
|
29
|
+
edge_index = backend.edge_index(data)
|
|
30
|
+
edge_weight = backend.edge_weight(data)
|
|
31
|
+
backend.forward(model, x, edge_index, edge_weight=edge_weight)
|
|
32
|
+
finally:
|
|
33
|
+
handle.remove()
|
|
34
|
+
emb = captured.get("emb")
|
|
35
|
+
if emb is None:
|
|
36
|
+
raise RuntimeError("No se capturaron embeddings del modelo")
|
|
37
|
+
return emb
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def edge_embeddings(
|
|
41
|
+
embeddings: torch.Tensor,
|
|
42
|
+
edge_index: torch.Tensor,
|
|
43
|
+
) -> torch.Tensor:
|
|
44
|
+
return torch.cat([embeddings[edge_index[0]], embeddings[edge_index[1]]], dim=-1)
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import inspect
|
|
4
|
+
from collections.abc import Callable
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
_ALGORITHMS: dict[str, type] = {}
|
|
8
|
+
_ALIASES: dict[str, str] = {}
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _accepted_params(cls: type) -> set[str]:
|
|
12
|
+
try:
|
|
13
|
+
sig = inspect.signature(cls.__init__)
|
|
14
|
+
except (TypeError, ValueError):
|
|
15
|
+
return set()
|
|
16
|
+
names = set()
|
|
17
|
+
for name, p in sig.parameters.items():
|
|
18
|
+
if name in ("self", "kwargs", "args"):
|
|
19
|
+
continue
|
|
20
|
+
if p.kind in (p.POSITIONAL_OR_KEYWORD, p.KEYWORD_ONLY):
|
|
21
|
+
names.add(name)
|
|
22
|
+
return names
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def register(name: str, *aliases: str) -> Callable[[type], type]:
|
|
26
|
+
def decorator(cls: type) -> type:
|
|
27
|
+
_ALGORITHMS[name] = cls
|
|
28
|
+
for alias in aliases:
|
|
29
|
+
_ALIASES[alias] = name
|
|
30
|
+
cls.name = name
|
|
31
|
+
return cls
|
|
32
|
+
|
|
33
|
+
return decorator
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def get_algorithm(name: str) -> type:
|
|
37
|
+
registered = _ALGORITHMS.get(name) or _ALGORITHMS.get(_ALIASES.get(name, ""))
|
|
38
|
+
if registered is None:
|
|
39
|
+
from ..methods import _available_methods
|
|
40
|
+
|
|
41
|
+
raise ValueError(
|
|
42
|
+
f"Algoritmo desconocido: {name}. Disponibles: {sorted(_available_methods())}"
|
|
43
|
+
)
|
|
44
|
+
return registered
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def instantiate(name: str, **kwargs) -> Any:
|
|
48
|
+
cls = get_algorithm(name)
|
|
49
|
+
accepted = _accepted_params(cls)
|
|
50
|
+
filtered = {k: v for k, v in kwargs.items() if k in accepted}
|
|
51
|
+
return cls(**filtered)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _available_methods() -> set[str]:
|
|
55
|
+
return set(_ALGORITHMS.keys()) | set(_ALIASES.keys())
|
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
from .attention.attention import AttentionExplainer
|
|
2
|
+
from .base import ExplanationAlgorithm
|
|
3
|
+
from .baseline.random_baseline import RandomBaseline
|
|
4
|
+
from .counterfactual.counterfactual import Counterfactual
|
|
5
|
+
from .feature.graph_lime import GraphLIME
|
|
6
|
+
from .gradient.grad_x_input import GradXInput
|
|
7
|
+
from .gradient.guided_backprop import GuidedBackprop
|
|
8
|
+
from .gradient.integrated_gradients import IntegratedGradients
|
|
9
|
+
from .gradient.saliency import Saliency
|
|
10
|
+
from .perturbation.gnn_explainer import GNNExplainer
|
|
11
|
+
from .perturbation.node_mask import NodeMask
|
|
12
|
+
from .perturbation.pg_explainer import PGExplainer
|
|
13
|
+
from .perturbation.subgraphx import SubgraphX
|
|
14
|
+
from .relevance.deeplift import DeepLift
|
|
15
|
+
from .relevance.gnn_lrp import GNNGatedLRP
|
|
16
|
+
|
|
17
|
+
__all__ = [
|
|
18
|
+
"AttentionExplainer",
|
|
19
|
+
"Counterfactual",
|
|
20
|
+
"DeepLift",
|
|
21
|
+
"ExplanationAlgorithm",
|
|
22
|
+
"GNNExplainer",
|
|
23
|
+
"GNNGatedLRP",
|
|
24
|
+
"GradXInput",
|
|
25
|
+
"GraphLIME",
|
|
26
|
+
"GuidedBackprop",
|
|
27
|
+
"IntegratedGradients",
|
|
28
|
+
"NodeMask",
|
|
29
|
+
"PGExplainer",
|
|
30
|
+
"RandomBaseline",
|
|
31
|
+
"Saliency",
|
|
32
|
+
"SubgraphX",
|
|
33
|
+
]
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _available_methods():
|
|
37
|
+
from ..core.registry import _available_methods
|
|
38
|
+
|
|
39
|
+
return _available_methods()
|
|
@@ -0,0 +1,147 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
from ...core.explanation import Explanation
|
|
8
|
+
from ...core.registry import register
|
|
9
|
+
from ..base import ExplanationAlgorithm
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _is_gat(module) -> bool:
|
|
13
|
+
try:
|
|
14
|
+
from torch_geometric.nn import GATConv
|
|
15
|
+
|
|
16
|
+
return isinstance(module, GATConv)
|
|
17
|
+
except ImportError:
|
|
18
|
+
return False
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@register("attention", "gat", "attention_explainer")
|
|
22
|
+
class AttentionExplainer(ExplanationAlgorithm):
|
|
23
|
+
"""Explanation based on the attention weights of GAT models.
|
|
24
|
+
|
|
25
|
+
Captures the attention coefficients (pre-softmax) of each `GATConv` during a
|
|
26
|
+
single forward pass and normalizes them with a per-neighbor softmax. Edge
|
|
27
|
+
importance is the mean of the coefficients across attention heads and GAT
|
|
28
|
+
layers; node importance aggregates the attention of the incident edges.
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
head_aggregate: str = "mean",
|
|
34
|
+
layer_aggregate: str = "mean",
|
|
35
|
+
node_aggregate: str = "sum",
|
|
36
|
+
):
|
|
37
|
+
self.head_aggregate = head_aggregate
|
|
38
|
+
self.layer_aggregate = layer_aggregate
|
|
39
|
+
self.node_aggregate = node_aggregate
|
|
40
|
+
|
|
41
|
+
def explain(
|
|
42
|
+
self,
|
|
43
|
+
backend: Any,
|
|
44
|
+
model: Any,
|
|
45
|
+
data: Any,
|
|
46
|
+
index: int | list[int] | torch.Tensor | None = None,
|
|
47
|
+
target_class: int | None = None,
|
|
48
|
+
**kwargs,
|
|
49
|
+
) -> Explanation:
|
|
50
|
+
model.eval()
|
|
51
|
+
x = backend.node_features(data)
|
|
52
|
+
edge_index = backend.edge_index(data)
|
|
53
|
+
num_nodes = int(x.size(0))
|
|
54
|
+
convs = [m for m in model.modules() if _is_gat(m)]
|
|
55
|
+
if not convs:
|
|
56
|
+
raise ValueError(
|
|
57
|
+
"AttentionExplainer requiere un modelo con capas GATConv "
|
|
58
|
+
"(torch_geometric.nn.GATConv)."
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
nodes = self._to_node_ids(index, num_nodes)
|
|
62
|
+
captured = self._capture_attention(model, backend, x, edge_index)
|
|
63
|
+
logits = captured["logits"]
|
|
64
|
+
# GATConv normaliza sobre las aristas + self-loops (añadidos al final)
|
|
65
|
+
alphas = [a[: edge_index.size(1)] for a in captured["alphas"]]
|
|
66
|
+
|
|
67
|
+
if logits.dim() != 2:
|
|
68
|
+
raise ValueError(
|
|
69
|
+
"AttentionExplainer requiere predicciones node-level (logits (N, C))."
|
|
70
|
+
)
|
|
71
|
+
target = target_class
|
|
72
|
+
if target is None:
|
|
73
|
+
target = int(logits[nodes[0]].argmax().item())
|
|
74
|
+
target_cls = max(0, min(int(target), logits.size(1) - 1))
|
|
75
|
+
|
|
76
|
+
per_layer: list[torch.Tensor] = []
|
|
77
|
+
for layer_idx, alpha in enumerate(alphas):
|
|
78
|
+
att = torch.softmax(alpha, dim=0) # (E, heads) normalizada por vecino
|
|
79
|
+
if self.head_aggregate == "mean":
|
|
80
|
+
att = att.mean(dim=1)
|
|
81
|
+
elif self.head_aggregate == "max":
|
|
82
|
+
att = att.max(dim=1).values
|
|
83
|
+
elif self.head_aggregate == "sum":
|
|
84
|
+
att = att.sum(dim=1)
|
|
85
|
+
else:
|
|
86
|
+
raise ValueError(f"head_aggregate desconocido: {self.head_aggregate}")
|
|
87
|
+
weight = 1.0 / len(convs) if self.layer_aggregate == "mean" else 1.0
|
|
88
|
+
per_layer.append(weight * att)
|
|
89
|
+
|
|
90
|
+
edge_importance = torch.stack(per_layer).sum(dim=0).clamp(min=0.0)
|
|
91
|
+
|
|
92
|
+
node_importance = torch.zeros(num_nodes, device=x.device)
|
|
93
|
+
if self.node_aggregate == "sum":
|
|
94
|
+
node_importance = node_importance.index_add(
|
|
95
|
+
0, edge_index[0], edge_importance
|
|
96
|
+
)
|
|
97
|
+
node_importance = node_importance.index_add(
|
|
98
|
+
0, edge_index[1], edge_importance
|
|
99
|
+
)
|
|
100
|
+
elif self.node_aggregate == "in":
|
|
101
|
+
node_importance = node_importance.index_add(
|
|
102
|
+
0, edge_index[1], edge_importance
|
|
103
|
+
)
|
|
104
|
+
elif self.node_aggregate == "out":
|
|
105
|
+
node_importance = node_importance.index_add(
|
|
106
|
+
0, edge_index[0], edge_importance
|
|
107
|
+
)
|
|
108
|
+
else:
|
|
109
|
+
raise ValueError(f"node_aggregate desconocido: {self.node_aggregate}")
|
|
110
|
+
|
|
111
|
+
return Explanation(
|
|
112
|
+
node_importance=node_importance.detach().cpu(),
|
|
113
|
+
edge_importance=edge_importance.detach().cpu(),
|
|
114
|
+
feature_importance=None,
|
|
115
|
+
prediction_original=logits[nodes[0]].detach().cpu(),
|
|
116
|
+
prediction_explanation=None,
|
|
117
|
+
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else index,
|
|
118
|
+
target_class=target_cls,
|
|
119
|
+
)
|
|
120
|
+
|
|
121
|
+
@staticmethod
|
|
122
|
+
def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
|
|
123
|
+
if index is None:
|
|
124
|
+
return torch.zeros(1, dtype=torch.long)
|
|
125
|
+
if isinstance(index, int):
|
|
126
|
+
return torch.tensor([index], dtype=torch.long)
|
|
127
|
+
idx = torch.as_tensor(index, dtype=torch.long)
|
|
128
|
+
return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
|
|
129
|
+
|
|
130
|
+
def _capture_attention(self, model, backend, x, edge_index):
|
|
131
|
+
from torch_geometric import nn as pyg_nn
|
|
132
|
+
|
|
133
|
+
captured: list[torch.Tensor] = []
|
|
134
|
+
real = pyg_nn.conv.gat_conv.softmax
|
|
135
|
+
|
|
136
|
+
def _wrapped(alpha, index, ptr=None, num_nodes=None):
|
|
137
|
+
captured.append(alpha.detach())
|
|
138
|
+
return real(alpha, index, ptr, num_nodes)
|
|
139
|
+
|
|
140
|
+
pyg_nn.conv.gat_conv.softmax = _wrapped
|
|
141
|
+
try:
|
|
142
|
+
with torch.no_grad():
|
|
143
|
+
logits = backend.forward(model, x, edge_index)
|
|
144
|
+
finally:
|
|
145
|
+
pyg_nn.conv.gat_conv.softmax = real
|
|
146
|
+
# un solo softmax por capa GAT en el forward
|
|
147
|
+
return {"logits": logits, "alphas": captured}
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from abc import ABC, abstractmethod
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import torch
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ExplanationAlgorithm(ABC):
|
|
10
|
+
name = "base"
|
|
11
|
+
graph_level = False
|
|
12
|
+
|
|
13
|
+
@abstractmethod
|
|
14
|
+
def explain(
|
|
15
|
+
self,
|
|
16
|
+
backend: Any,
|
|
17
|
+
model: Any,
|
|
18
|
+
data: Any,
|
|
19
|
+
index: int | list[int] | torch.Tensor,
|
|
20
|
+
target_class: int | None = None,
|
|
21
|
+
**kwargs,
|
|
22
|
+
) -> Any: ...
|
|
23
|
+
|
|
24
|
+
def validate(self, backend: Any, model: Any, data: Any) -> None:
|
|
25
|
+
return None
|
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
from ...core.explanation import Explanation
|
|
8
|
+
from ...core.registry import register
|
|
9
|
+
from ..base import ExplanationAlgorithm
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@register("random", "random_baseline", "rand")
|
|
13
|
+
class RandomBaseline(ExplanationAlgorithm):
|
|
14
|
+
"""Random: seed-able uniform random importance baseline.
|
|
15
|
+
|
|
16
|
+
Assigns random importances in [0, 1] to nodes, edges and features with no
|
|
17
|
+
link to the model; useful as a null-scenario reference in comparative
|
|
18
|
+
benchmarks.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
graph_level = True
|
|
22
|
+
|
|
23
|
+
def __init__(self, seed: int | None = 0, **kwargs):
|
|
24
|
+
self.seed = seed
|
|
25
|
+
|
|
26
|
+
def explain(
|
|
27
|
+
self,
|
|
28
|
+
backend: Any,
|
|
29
|
+
model: Any,
|
|
30
|
+
data: Any,
|
|
31
|
+
index: int | torch.Tensor | None = None,
|
|
32
|
+
target_class: int | None = None,
|
|
33
|
+
**kwargs,
|
|
34
|
+
) -> Explanation:
|
|
35
|
+
x = backend.node_features(data)
|
|
36
|
+
edge_index = backend.edge_index(data)
|
|
37
|
+
num_nodes = int(x.size(0))
|
|
38
|
+
num_edges = int(edge_index.size(1))
|
|
39
|
+
|
|
40
|
+
if self.seed is not None:
|
|
41
|
+
torch.manual_seed(self.seed)
|
|
42
|
+
|
|
43
|
+
node_importance = torch.rand(num_nodes).cpu()
|
|
44
|
+
edge_importance = torch.rand(num_edges).cpu()
|
|
45
|
+
feature_importance = torch.rand_like(x).cpu()
|
|
46
|
+
|
|
47
|
+
nodes = self._to_node_ids(index, num_nodes)
|
|
48
|
+
root = int(nodes[0])
|
|
49
|
+
with torch.no_grad():
|
|
50
|
+
logits = backend.forward(model, x, edge_index)
|
|
51
|
+
target = target_class
|
|
52
|
+
if logits.dim() == 2:
|
|
53
|
+
if target is None:
|
|
54
|
+
target = int(logits[root].argmax().item())
|
|
55
|
+
target = max(0, min(int(target), logits.size(1) - 1))
|
|
56
|
+
pred = logits[root].detach().cpu()
|
|
57
|
+
else:
|
|
58
|
+
target = None
|
|
59
|
+
pred = logits.detach().cpu()
|
|
60
|
+
|
|
61
|
+
return Explanation(
|
|
62
|
+
node_importance=node_importance,
|
|
63
|
+
edge_importance=edge_importance,
|
|
64
|
+
feature_importance=feature_importance,
|
|
65
|
+
prediction_original=pred,
|
|
66
|
+
prediction_explanation=None,
|
|
67
|
+
node_idx=root,
|
|
68
|
+
target_class=target,
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
@staticmethod
|
|
72
|
+
def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
|
|
73
|
+
if index is None:
|
|
74
|
+
return torch.zeros(1, dtype=torch.long)
|
|
75
|
+
if isinstance(index, int):
|
|
76
|
+
return torch.tensor([index], dtype=torch.long)
|
|
77
|
+
idx = torch.as_tensor(index, dtype=torch.long)
|
|
78
|
+
return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
|