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.
Files changed (42) hide show
  1. graph_explain/__init__.py +79 -0
  2. graph_explain/backends/__init__.py +4 -0
  3. graph_explain/backends/base.py +103 -0
  4. graph_explain/backends/dgl.py +121 -0
  5. graph_explain/benchmarks/__init__.py +3 -0
  6. graph_explain/benchmarks/synthetic.py +246 -0
  7. graph_explain/cli.py +459 -0
  8. graph_explain/core/__init__.py +14 -0
  9. graph_explain/core/benchmark.py +284 -0
  10. graph_explain/core/evaluation.py +391 -0
  11. graph_explain/core/explainer.py +83 -0
  12. graph_explain/core/explanation.py +72 -0
  13. graph_explain/core/model_utils.py +44 -0
  14. graph_explain/core/registry.py +55 -0
  15. graph_explain/methods/__init__.py +39 -0
  16. graph_explain/methods/attention/attention.py +147 -0
  17. graph_explain/methods/base.py +25 -0
  18. graph_explain/methods/baseline/random_baseline.py +78 -0
  19. graph_explain/methods/counterfactual/counterfactual.py +304 -0
  20. graph_explain/methods/feature/graph_lime.py +141 -0
  21. graph_explain/methods/gradient/__init__.py +0 -0
  22. graph_explain/methods/gradient/grad_x_input.py +110 -0
  23. graph_explain/methods/gradient/guided_backprop.py +117 -0
  24. graph_explain/methods/gradient/integrated_gradients.py +115 -0
  25. graph_explain/methods/gradient/saliency.py +93 -0
  26. graph_explain/methods/perturbation/__init__.py +0 -0
  27. graph_explain/methods/perturbation/gnn_explainer.py +265 -0
  28. graph_explain/methods/perturbation/node_mask.py +136 -0
  29. graph_explain/methods/perturbation/pg_explainer.py +162 -0
  30. graph_explain/methods/perturbation/subgraphx.py +393 -0
  31. graph_explain/methods/relevance/deeplift.py +262 -0
  32. graph_explain/methods/relevance/gnn_lrp.py +219 -0
  33. graph_explain/narration/__init__.py +3 -0
  34. graph_explain/narration/narrator.py +185 -0
  35. graph_explain/visualization/__init__.py +4 -0
  36. graph_explain/visualization/interactive.py +73 -0
  37. graph_explain/visualization/static.py +90 -0
  38. graph_explain-0.7.0.dist-info/METADATA +332 -0
  39. graph_explain-0.7.0.dist-info/RECORD +42 -0
  40. graph_explain-0.7.0.dist-info/WHEEL +5 -0
  41. graph_explain-0.7.0.dist-info/entry_points.txt +2 -0
  42. 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)