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,115 @@
|
|
|
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("integrated_gradients", "ig")
|
|
13
|
+
class IntegratedGradients(ExplanationAlgorithm):
|
|
14
|
+
graph_level = True
|
|
15
|
+
|
|
16
|
+
def __init__(
|
|
17
|
+
self,
|
|
18
|
+
steps: int = 50,
|
|
19
|
+
method: str = "riemann",
|
|
20
|
+
edge_grads: bool = True,
|
|
21
|
+
**kwargs,
|
|
22
|
+
):
|
|
23
|
+
self.steps = steps
|
|
24
|
+
self.method = method
|
|
25
|
+
self.edge_grads = edge_grads
|
|
26
|
+
|
|
27
|
+
def explain(
|
|
28
|
+
self,
|
|
29
|
+
backend: Any,
|
|
30
|
+
model: Any,
|
|
31
|
+
data: Any,
|
|
32
|
+
index: int | list[int] | torch.Tensor,
|
|
33
|
+
target_class: int | None = None,
|
|
34
|
+
**kwargs,
|
|
35
|
+
) -> Explanation:
|
|
36
|
+
model.eval()
|
|
37
|
+
x = backend.node_features(data)
|
|
38
|
+
edge_index = backend.edge_index(data)
|
|
39
|
+
|
|
40
|
+
if index is None:
|
|
41
|
+
idx = torch.zeros(1, dtype=torch.long, device=x.device)
|
|
42
|
+
elif isinstance(index, int):
|
|
43
|
+
idx = torch.tensor([index], device=x.device)
|
|
44
|
+
else:
|
|
45
|
+
idx = torch.as_tensor(index, device=x.device)
|
|
46
|
+
|
|
47
|
+
with torch.no_grad():
|
|
48
|
+
logits = backend.forward(model, x, edge_index)
|
|
49
|
+
if target_class is None and logits.dim() == 2:
|
|
50
|
+
target_class = int(logits[idx[0]].argmax().item())
|
|
51
|
+
|
|
52
|
+
baseline = torch.zeros_like(x)
|
|
53
|
+
alphas, weights = self._alphas(self.steps, self.method, device=x.device)
|
|
54
|
+
|
|
55
|
+
ig_feat = torch.zeros_like(x)
|
|
56
|
+
ig_edge = None
|
|
57
|
+
compute_edge = self.edge_grads and backend.supports_edge_weight(model)
|
|
58
|
+
if compute_edge:
|
|
59
|
+
ig_edge = torch.zeros(
|
|
60
|
+
edge_index.size(1), dtype=torch.float32, device=x.device
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
for alpha, w in zip(alphas, weights):
|
|
64
|
+
x_step = (baseline + alpha * (x - baseline)).requires_grad_(True)
|
|
65
|
+
ew = None
|
|
66
|
+
if compute_edge:
|
|
67
|
+
ew = (
|
|
68
|
+
torch.ones(edge_index.size(1), device=x.device) * (1.0 - alpha)
|
|
69
|
+
).requires_grad_(True)
|
|
70
|
+
out = backend.forward(model, x_step, edge_index, edge_weight=ew)
|
|
71
|
+
if out.dim() == 2:
|
|
72
|
+
score = out[idx, target_class].sum()
|
|
73
|
+
else:
|
|
74
|
+
score = out[idx].sum()
|
|
75
|
+
model.zero_grad()
|
|
76
|
+
score.backward()
|
|
77
|
+
ig_feat = ig_feat + w * x_step.grad
|
|
78
|
+
if compute_edge and ew is not None and ew.grad is not None:
|
|
79
|
+
ig_edge = ig_edge + w * ew.grad
|
|
80
|
+
|
|
81
|
+
ig_feat = ig_feat * (x - baseline)
|
|
82
|
+
if compute_edge and ig_edge is not None:
|
|
83
|
+
ig_edge = ig_edge * (-1.0)
|
|
84
|
+
|
|
85
|
+
grad = ig_feat.detach()
|
|
86
|
+
node_importance = grad.abs().sum(dim=-1)
|
|
87
|
+
feature_importance = grad
|
|
88
|
+
|
|
89
|
+
return Explanation(
|
|
90
|
+
node_importance=node_importance.cpu(),
|
|
91
|
+
edge_importance=(
|
|
92
|
+
ig_edge.detach().abs().cpu() if ig_edge is not None else None
|
|
93
|
+
),
|
|
94
|
+
feature_importance=feature_importance.cpu(),
|
|
95
|
+
prediction_original=logits[idx[0]].detach().reshape(1, -1).cpu(),
|
|
96
|
+
prediction_explanation=None,
|
|
97
|
+
node_idx=int(idx[0].item()) if isinstance(index, int) else index,
|
|
98
|
+
target_class=target_class,
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
def _alphas(self, steps: int, method: str, device):
|
|
102
|
+
if method in ("riemann", "left"):
|
|
103
|
+
alphas = torch.arange(0.0, 1.0, 1.0 / steps, device=device)
|
|
104
|
+
return alphas, torch.full_like(alphas, 1.0 / steps)
|
|
105
|
+
if method == "right":
|
|
106
|
+
alphas = torch.arange(1.0 / steps, 1.0 + 1e-6, 1.0 / steps, device=device)
|
|
107
|
+
return alphas, torch.full_like(alphas, 1.0 / steps)
|
|
108
|
+
if method == "gausslegendre":
|
|
109
|
+
from numpy.polynomial.legendre import leggauss
|
|
110
|
+
|
|
111
|
+
xs, ws = leggauss(steps)
|
|
112
|
+
alphas = torch.as_tensor((xs + 1) / 2, device=device, dtype=torch.float32)
|
|
113
|
+
weights = torch.as_tensor(ws / 2, device=device, dtype=torch.float32)
|
|
114
|
+
return alphas, weights
|
|
115
|
+
raise ValueError(f"method desconocido: {method}")
|
|
@@ -0,0 +1,93 @@
|
|
|
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("saliency", "gradient", "grad")
|
|
13
|
+
class Saliency(ExplanationAlgorithm):
|
|
14
|
+
graph_level = True
|
|
15
|
+
|
|
16
|
+
def __init__(
|
|
17
|
+
self,
|
|
18
|
+
absolute: bool = True,
|
|
19
|
+
aggregate: str = "sum",
|
|
20
|
+
node_mask_type: str | None = None,
|
|
21
|
+
):
|
|
22
|
+
self.absolute = absolute
|
|
23
|
+
self.aggregate = aggregate
|
|
24
|
+
self.node_mask_type = node_mask_type
|
|
25
|
+
|
|
26
|
+
def explain(
|
|
27
|
+
self,
|
|
28
|
+
backend: Any,
|
|
29
|
+
model: Any,
|
|
30
|
+
data: Any,
|
|
31
|
+
index: int | list[int] | torch.Tensor,
|
|
32
|
+
target_class: int | None = None,
|
|
33
|
+
**kwargs,
|
|
34
|
+
) -> Explanation:
|
|
35
|
+
model.eval()
|
|
36
|
+
x = backend.node_features(data).detach().clone().requires_grad_(True)
|
|
37
|
+
edge_index = backend.edge_index(data)
|
|
38
|
+
edge_weight = backend.edge_weight(data)
|
|
39
|
+
|
|
40
|
+
out = backend.forward(model, x, edge_index, edge_weight=edge_weight)
|
|
41
|
+
logits = out
|
|
42
|
+
|
|
43
|
+
if index is None:
|
|
44
|
+
idx = torch.zeros(1, dtype=torch.long, device=x.device)
|
|
45
|
+
elif isinstance(index, int):
|
|
46
|
+
idx = torch.tensor([index], device=x.device)
|
|
47
|
+
else:
|
|
48
|
+
idx = torch.as_tensor(index, device=x.device)
|
|
49
|
+
|
|
50
|
+
if target_class is None and logits.dim() == 2:
|
|
51
|
+
target_class = int(logits[idx].argmax(dim=-1)[0].item())
|
|
52
|
+
|
|
53
|
+
if logits.dim() == 2:
|
|
54
|
+
score = logits[idx, target_class].sum()
|
|
55
|
+
else:
|
|
56
|
+
score = logits[idx].sum()
|
|
57
|
+
|
|
58
|
+
model.zero_grad()
|
|
59
|
+
score.backward()
|
|
60
|
+
|
|
61
|
+
grad = x.grad
|
|
62
|
+
if grad is None:
|
|
63
|
+
raise RuntimeError("No se obtuvieron gradientes del modelo.")
|
|
64
|
+
|
|
65
|
+
if self.absolute:
|
|
66
|
+
grad = grad.abs()
|
|
67
|
+
if self.aggregate == "sum":
|
|
68
|
+
node_importance = grad.sum(dim=-1)
|
|
69
|
+
elif self.aggregate == "mean":
|
|
70
|
+
node_importance = grad.mean(dim=-1)
|
|
71
|
+
elif self.aggregate == "max":
|
|
72
|
+
node_importance = grad.max(dim=-1).values
|
|
73
|
+
else:
|
|
74
|
+
raise ValueError(f"aggregate desconocido: {self.aggregate}")
|
|
75
|
+
|
|
76
|
+
prediction_original = logits[idx].detach()
|
|
77
|
+
feature_importance = (
|
|
78
|
+
grad.detach() if self.node_mask_type == "attributes" else None
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
return Explanation(
|
|
82
|
+
node_importance=node_importance.detach().cpu(),
|
|
83
|
+
feature_importance=(
|
|
84
|
+
feature_importance.detach().cpu()
|
|
85
|
+
if feature_importance is not None
|
|
86
|
+
else None
|
|
87
|
+
),
|
|
88
|
+
edge_importance=None,
|
|
89
|
+
prediction_original=prediction_original.cpu(),
|
|
90
|
+
prediction_explanation=None,
|
|
91
|
+
node_idx=int(idx[0].item()) if isinstance(index, int) else index,
|
|
92
|
+
target_class=target_class,
|
|
93
|
+
)
|
|
File without changes
|
|
@@ -0,0 +1,265 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
import torch.nn.functional as F
|
|
7
|
+
|
|
8
|
+
from ...core.explanation import Explanation
|
|
9
|
+
from ...core.registry import register
|
|
10
|
+
from ..base import ExplanationAlgorithm
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@register("gnn_explainer", "gnnexplainer")
|
|
14
|
+
class GNNExplainer(ExplanationAlgorithm):
|
|
15
|
+
graph_level = True
|
|
16
|
+
|
|
17
|
+
def __init__(
|
|
18
|
+
self,
|
|
19
|
+
epochs: int = 200,
|
|
20
|
+
lr: float = 0.01,
|
|
21
|
+
edge_entropy: float = 0.001,
|
|
22
|
+
node_entropy: float = 0.001,
|
|
23
|
+
node_mask_type: str | None = "attributes",
|
|
24
|
+
edge_mask_type: str | None = "object",
|
|
25
|
+
prints: int = 20,
|
|
26
|
+
**kwargs,
|
|
27
|
+
):
|
|
28
|
+
self.epochs = epochs
|
|
29
|
+
self.lr = lr
|
|
30
|
+
self.edge_entropy = edge_entropy
|
|
31
|
+
self.node_entropy = node_entropy
|
|
32
|
+
self.node_mask_type = node_mask_type
|
|
33
|
+
self.edge_mask_type = edge_mask_type
|
|
34
|
+
self.prints = prints
|
|
35
|
+
|
|
36
|
+
def explain(
|
|
37
|
+
self,
|
|
38
|
+
backend: Any,
|
|
39
|
+
model: Any,
|
|
40
|
+
data: Any,
|
|
41
|
+
index: int | torch.Tensor,
|
|
42
|
+
target_class: int | None = None,
|
|
43
|
+
**kwargs,
|
|
44
|
+
) -> Explanation:
|
|
45
|
+
node_mask_type = kwargs.get("node_mask_type", self.node_mask_type)
|
|
46
|
+
edge_mask_type = kwargs.get("edge_mask_type", self.edge_mask_type)
|
|
47
|
+
|
|
48
|
+
model.eval()
|
|
49
|
+
x = backend.node_features(data)
|
|
50
|
+
edge_index = backend.edge_index(data)
|
|
51
|
+
num_nodes = backend.num_nodes(data)
|
|
52
|
+
|
|
53
|
+
sub_nodes, sub_edge_index, mapping, sub_edge_mask = self._extract_subgraph(
|
|
54
|
+
backend, data, index, edge_index
|
|
55
|
+
)
|
|
56
|
+
device = x.device
|
|
57
|
+
x_sub = x[sub_nodes].to(device)
|
|
58
|
+
sub_graph_level = index is None
|
|
59
|
+
|
|
60
|
+
node_mask = None
|
|
61
|
+
if node_mask_type is not None and not sub_graph_level:
|
|
62
|
+
node_mask = torch.nn.Parameter(torch.randn(x_sub.size(0), device=device))
|
|
63
|
+
|
|
64
|
+
edge_mask = None
|
|
65
|
+
if edge_mask_type is not None:
|
|
66
|
+
edge_mask = torch.nn.Parameter(
|
|
67
|
+
torch.randn(sub_edge_index.size(1), device=device)
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
params = [p for p in (node_mask, edge_mask) if p is not None]
|
|
71
|
+
optimizer = torch.optim.Adam(params, lr=self.lr)
|
|
72
|
+
|
|
73
|
+
with torch.no_grad():
|
|
74
|
+
orig_logits = backend.forward(model, x_sub, sub_edge_index)
|
|
75
|
+
if sub_graph_level:
|
|
76
|
+
if target_class is None:
|
|
77
|
+
target_class = (
|
|
78
|
+
int(orig_logits[0].argmax().item()) if orig_logits.dim() == 2 else 0
|
|
79
|
+
)
|
|
80
|
+
tgt_idx = 0
|
|
81
|
+
else:
|
|
82
|
+
ni = int(mapping.item() if torch.is_tensor(mapping) else mapping)
|
|
83
|
+
if target_class is None and orig_logits.dim() == 2:
|
|
84
|
+
target_class = int(orig_logits[ni].argmax().item())
|
|
85
|
+
tgt_idx = ni
|
|
86
|
+
|
|
87
|
+
for epoch in range(self.epochs):
|
|
88
|
+
optimizer.zero_grad()
|
|
89
|
+
mask_node = None
|
|
90
|
+
if node_mask is not None:
|
|
91
|
+
mask_node = torch.sigmoid(node_mask)
|
|
92
|
+
eweight = None
|
|
93
|
+
if edge_mask is not None:
|
|
94
|
+
eweight = torch.sigmoid(edge_mask)
|
|
95
|
+
|
|
96
|
+
pred = backend.forward(
|
|
97
|
+
model, x_sub, sub_edge_index, edge_weight=eweight, node_mask=mask_node
|
|
98
|
+
)
|
|
99
|
+
loss = self._loss(
|
|
100
|
+
pred, tgt_idx, sub_graph_level, target_class, node_mask, edge_mask
|
|
101
|
+
)
|
|
102
|
+
loss.backward()
|
|
103
|
+
optimizer.step()
|
|
104
|
+
|
|
105
|
+
logits = backend.forward(model, x_sub, sub_edge_index)
|
|
106
|
+
|
|
107
|
+
if sub_graph_level:
|
|
108
|
+
target_class = target_class or (
|
|
109
|
+
int(logits.argmax(-1)[0].item()) if logits.dim() == 2 else 0
|
|
110
|
+
)
|
|
111
|
+
ni = 0
|
|
112
|
+
pred_orig = logits[0].detach()
|
|
113
|
+
final_mask_node = None
|
|
114
|
+
final_mask_edge = None
|
|
115
|
+
if node_mask is not None:
|
|
116
|
+
final_mask_node = torch.sigmoid(node_mask)
|
|
117
|
+
if edge_mask is not None:
|
|
118
|
+
final_mask_edge = torch.sigmoid(edge_mask)
|
|
119
|
+
pred_masked = backend.forward(
|
|
120
|
+
model,
|
|
121
|
+
x_sub,
|
|
122
|
+
sub_edge_index,
|
|
123
|
+
edge_weight=final_mask_edge,
|
|
124
|
+
node_mask=final_mask_node,
|
|
125
|
+
)[0].detach()
|
|
126
|
+
else:
|
|
127
|
+
ni = int(mapping.item() if torch.is_tensor(mapping) else mapping)
|
|
128
|
+
if target_class is None and logits.dim() == 2:
|
|
129
|
+
target_class = int(logits[ni].argmax().item())
|
|
130
|
+
pred_orig = logits[ni].detach()
|
|
131
|
+
final_mask_node = None
|
|
132
|
+
final_mask_edge = None
|
|
133
|
+
if node_mask is not None:
|
|
134
|
+
final_mask_node = torch.sigmoid(node_mask)
|
|
135
|
+
if edge_mask is not None:
|
|
136
|
+
final_mask_edge = torch.sigmoid(edge_mask)
|
|
137
|
+
pred_masked = backend.forward(
|
|
138
|
+
model,
|
|
139
|
+
x_sub,
|
|
140
|
+
sub_edge_index,
|
|
141
|
+
edge_weight=final_mask_edge,
|
|
142
|
+
node_mask=final_mask_node,
|
|
143
|
+
)[ni].detach()
|
|
144
|
+
|
|
145
|
+
full_num_nodes = num_nodes
|
|
146
|
+
full_edge_count = edge_index.size(1)
|
|
147
|
+
|
|
148
|
+
node_full = self._scatter_node(node_mask, sub_nodes, full_num_nodes)
|
|
149
|
+
edge_full = self._scatter_edge(edge_mask, sub_edge_mask, full_edge_count)
|
|
150
|
+
|
|
151
|
+
return Explanation(
|
|
152
|
+
node_importance=node_full,
|
|
153
|
+
edge_importance=edge_full,
|
|
154
|
+
feature_importance=None,
|
|
155
|
+
prediction_original=pred_orig.cpu(),
|
|
156
|
+
prediction_explanation=pred_masked.cpu(),
|
|
157
|
+
node_idx=None
|
|
158
|
+
if sub_graph_level
|
|
159
|
+
else (int(index[0]) if torch.is_tensor(index) else int(index)),
|
|
160
|
+
target_class=target_class,
|
|
161
|
+
metadata={
|
|
162
|
+
"sub_nodes": sub_nodes,
|
|
163
|
+
"sub_edge_index": sub_edge_index,
|
|
164
|
+
"sub_edge_mask": sub_edge_mask,
|
|
165
|
+
},
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
@staticmethod
|
|
169
|
+
def _scatter_node(
|
|
170
|
+
mask: torch.nn.Parameter | None,
|
|
171
|
+
sub_nodes: torch.Tensor,
|
|
172
|
+
num_nodes: int,
|
|
173
|
+
) -> torch.Tensor | None:
|
|
174
|
+
if mask is None:
|
|
175
|
+
return None
|
|
176
|
+
full = torch.zeros(num_nodes, dtype=torch.float32)
|
|
177
|
+
vals = torch.sigmoid(mask).detach().cpu()
|
|
178
|
+
full[sub_nodes.cpu()] = vals
|
|
179
|
+
return full
|
|
180
|
+
|
|
181
|
+
@staticmethod
|
|
182
|
+
def _scatter_edge(
|
|
183
|
+
mask: torch.nn.Parameter | None,
|
|
184
|
+
sub_edge_mask: torch.Tensor,
|
|
185
|
+
num_edges: int,
|
|
186
|
+
) -> torch.Tensor | None:
|
|
187
|
+
if mask is None:
|
|
188
|
+
return None
|
|
189
|
+
full = torch.zeros(num_edges, dtype=torch.float32)
|
|
190
|
+
idx = sub_edge_mask.nonzero(as_tuple=False).view(-1)
|
|
191
|
+
full[idx.cpu()] = torch.sigmoid(mask).detach().cpu()
|
|
192
|
+
return full
|
|
193
|
+
|
|
194
|
+
def _loss(
|
|
195
|
+
self,
|
|
196
|
+
pred: torch.Tensor,
|
|
197
|
+
node_idx: int | torch.Tensor,
|
|
198
|
+
sub_graph_level: bool,
|
|
199
|
+
target_class: int | None,
|
|
200
|
+
node_mask: torch.nn.Parameter | None,
|
|
201
|
+
edge_mask: torch.nn.Parameter | None,
|
|
202
|
+
) -> torch.Tensor:
|
|
203
|
+
if pred.dim() == 2:
|
|
204
|
+
log_logits = pred.log_softmax(dim=-1)
|
|
205
|
+
else:
|
|
206
|
+
log_logits = pred
|
|
207
|
+
|
|
208
|
+
if sub_graph_level:
|
|
209
|
+
idx = torch.zeros(1, dtype=torch.long, device=pred.device)
|
|
210
|
+
if target_class is None:
|
|
211
|
+
target_class = int(pred[0].argmax().item())
|
|
212
|
+
loss = F.nll_loss(
|
|
213
|
+
log_logits[0].unsqueeze(0),
|
|
214
|
+
torch.tensor([target_class], device=pred.device),
|
|
215
|
+
)
|
|
216
|
+
else:
|
|
217
|
+
if isinstance(node_idx, torch.Tensor) and node_idx.dim() == 0:
|
|
218
|
+
idx = node_idx.unsqueeze(0)
|
|
219
|
+
else:
|
|
220
|
+
idx = (
|
|
221
|
+
torch.as_tensor([node_idx], device=pred.device)
|
|
222
|
+
if not torch.is_tensor(node_idx)
|
|
223
|
+
else node_idx.reshape(-1)
|
|
224
|
+
)
|
|
225
|
+
if target_class is None:
|
|
226
|
+
target_class = int(pred[idx].argmax(dim=-1)[0].item())
|
|
227
|
+
loss = F.nll_loss(
|
|
228
|
+
log_logits[idx], torch.tensor([target_class], device=pred.device)
|
|
229
|
+
)
|
|
230
|
+
|
|
231
|
+
if edge_mask is not None and self.edge_entropy > 0:
|
|
232
|
+
loss += self.edge_entropy * self._entropy(torch.sigmoid(edge_mask))
|
|
233
|
+
if node_mask is not None and self.node_entropy > 0:
|
|
234
|
+
loss += self.node_entropy * self._entropy(torch.sigmoid(node_mask))
|
|
235
|
+
return loss
|
|
236
|
+
|
|
237
|
+
@staticmethod
|
|
238
|
+
def _entropy(p: torch.Tensor) -> torch.Tensor:
|
|
239
|
+
eps = 1e-8
|
|
240
|
+
return -(p * torch.log(p + eps) + (1 - p) * torch.log(1 - p + eps)).mean()
|
|
241
|
+
|
|
242
|
+
def _extract_subgraph(
|
|
243
|
+
self,
|
|
244
|
+
backend: Any,
|
|
245
|
+
data: Any,
|
|
246
|
+
index: int | torch.Tensor | None,
|
|
247
|
+
edge_index: torch.Tensor,
|
|
248
|
+
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
249
|
+
from torch_geometric.utils import k_hop_subgraph
|
|
250
|
+
|
|
251
|
+
if index is None:
|
|
252
|
+
n = edge_index.max().item() + 1
|
|
253
|
+
node_idx = torch.arange(n, device=edge_index.device)
|
|
254
|
+
return (
|
|
255
|
+
node_idx,
|
|
256
|
+
edge_index,
|
|
257
|
+
node_idx,
|
|
258
|
+
torch.ones(
|
|
259
|
+
edge_index.size(1), dtype=torch.bool, device=edge_index.device
|
|
260
|
+
),
|
|
261
|
+
)
|
|
262
|
+
node_idx = torch.as_tensor([index], device=edge_index.device).reshape(-1)
|
|
263
|
+
return k_hop_subgraph(
|
|
264
|
+
node_idx, num_hops=3, edge_index=edge_index, relabel_nodes=True
|
|
265
|
+
)
|
|
@@ -0,0 +1,136 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
import torch.nn.functional as F
|
|
7
|
+
|
|
8
|
+
from ...core.explanation import Explanation
|
|
9
|
+
from ...core.registry import register
|
|
10
|
+
from ..base import ExplanationAlgorithm
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@register("node_mask", "nodemask", "nm")
|
|
14
|
+
class NodeMask(ExplanationAlgorithm):
|
|
15
|
+
"""NodeMask: node mask learned by optimization.
|
|
16
|
+
|
|
17
|
+
Optimizes a (sigmoid) mask over the nodes of the target node's k-hop
|
|
18
|
+
subgraph so the model keeps its prediction, with an entropy regularizer to
|
|
19
|
+
force sparsity. The resulting node importance is re-projected onto the full
|
|
20
|
+
graph (0 outside the neighborhood).
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
epochs: int = 200,
|
|
26
|
+
lr: float = 0.05,
|
|
27
|
+
hops: int = 3,
|
|
28
|
+
suppress_ratio: float = 0.8,
|
|
29
|
+
entropy: float = 0.05,
|
|
30
|
+
**kwargs,
|
|
31
|
+
):
|
|
32
|
+
self.epochs = epochs
|
|
33
|
+
self.lr = lr
|
|
34
|
+
self.hops = hops
|
|
35
|
+
self.suppress_ratio = suppress_ratio
|
|
36
|
+
self.entropy = entropy
|
|
37
|
+
|
|
38
|
+
def explain(
|
|
39
|
+
self,
|
|
40
|
+
backend: Any,
|
|
41
|
+
model: Any,
|
|
42
|
+
data: Any,
|
|
43
|
+
index: int | torch.Tensor | None = None,
|
|
44
|
+
target_class: int | None = None,
|
|
45
|
+
**kwargs,
|
|
46
|
+
) -> Explanation:
|
|
47
|
+
model.eval()
|
|
48
|
+
x = backend.node_features(data)
|
|
49
|
+
edge_index = backend.edge_index(data)
|
|
50
|
+
num_nodes = backend.num_nodes(data)
|
|
51
|
+
|
|
52
|
+
nodes = self._to_node_ids(index, num_nodes)
|
|
53
|
+
root = int(nodes[0])
|
|
54
|
+
|
|
55
|
+
from torch_geometric.utils import k_hop_subgraph
|
|
56
|
+
|
|
57
|
+
device = x.device
|
|
58
|
+
sub_nodes, sub_edge_index, mapping, _ = k_hop_subgraph(
|
|
59
|
+
[root],
|
|
60
|
+
num_hops=self.hops,
|
|
61
|
+
edge_index=edge_index,
|
|
62
|
+
relabel_nodes=True,
|
|
63
|
+
)
|
|
64
|
+
x_sub = x[sub_nodes].to(device)
|
|
65
|
+
|
|
66
|
+
with torch.no_grad():
|
|
67
|
+
orig = backend.forward(model, x_sub, sub_edge_index)
|
|
68
|
+
ni = int(mapping.item() if torch.is_tensor(mapping) else mapping)
|
|
69
|
+
if target_class is None and orig.dim() == 2:
|
|
70
|
+
target_class = int(orig[ni].argmax().item())
|
|
71
|
+
target = target_class if target_class is not None else 0
|
|
72
|
+
|
|
73
|
+
if self.suppress_ratio > 0:
|
|
74
|
+
k = int(self.suppress_ratio * sub_nodes.numel())
|
|
75
|
+
k = max(0, k)
|
|
76
|
+
else:
|
|
77
|
+
k = max(0, sub_nodes.numel() - 1)
|
|
78
|
+
|
|
79
|
+
mask = torch.nn.Parameter(torch.zeros(sub_nodes.numel(), device=device))
|
|
80
|
+
optimizer = torch.optim.Adam([mask], lr=self.lr)
|
|
81
|
+
|
|
82
|
+
for _ in range(self.epochs):
|
|
83
|
+
optimizer.zero_grad()
|
|
84
|
+
node_mask = torch.sigmoid(mask)
|
|
85
|
+
pred = backend.forward(model, x_sub, sub_edge_index, node_mask=node_mask)
|
|
86
|
+
topk = node_mask.topk(max(k, 1)).values.min()
|
|
87
|
+
loss = self._loss(pred, ni, target, node_mask, topk)
|
|
88
|
+
loss.backward()
|
|
89
|
+
optimizer.step()
|
|
90
|
+
|
|
91
|
+
final = torch.sigmoid(mask).detach()
|
|
92
|
+
full = torch.zeros(num_nodes, dtype=torch.float32)
|
|
93
|
+
full[sub_nodes.cpu()] = final.cpu()
|
|
94
|
+
|
|
95
|
+
with torch.no_grad():
|
|
96
|
+
pred_masked = backend.forward(
|
|
97
|
+
model, x_sub, sub_edge_index, node_mask=final
|
|
98
|
+
)[ni]
|
|
99
|
+
|
|
100
|
+
return Explanation(
|
|
101
|
+
node_importance=full,
|
|
102
|
+
edge_importance=None,
|
|
103
|
+
feature_importance=None,
|
|
104
|
+
prediction_original=orig[ni].detach().cpu(),
|
|
105
|
+
prediction_explanation=pred_masked.detach().cpu(),
|
|
106
|
+
node_idx=root,
|
|
107
|
+
target_class=target,
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
def _loss(self, pred, node_idx, target_class, node_mask, topk) -> torch.Tensor:
|
|
111
|
+
if pred.dim() == 2:
|
|
112
|
+
log_logits = pred.log_softmax(dim=-1)
|
|
113
|
+
else:
|
|
114
|
+
log_logits = pred
|
|
115
|
+
idx = torch.as_tensor([node_idx], device=pred.device).reshape(-1)
|
|
116
|
+
loss = F.nll_loss(
|
|
117
|
+
log_logits[idx],
|
|
118
|
+
torch.tensor([target_class], device=pred.device),
|
|
119
|
+
)
|
|
120
|
+
loss += self.entropy * self._entropy(node_mask)
|
|
121
|
+
loss += (node_mask - topk.detach()).relu().mean()
|
|
122
|
+
return loss
|
|
123
|
+
|
|
124
|
+
@staticmethod
|
|
125
|
+
def _entropy(p: torch.Tensor) -> torch.Tensor:
|
|
126
|
+
eps = 1e-8
|
|
127
|
+
return -(p * torch.log(p + eps) + (1 - p) * torch.log(1 - p + eps)).mean()
|
|
128
|
+
|
|
129
|
+
@staticmethod
|
|
130
|
+
def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
|
|
131
|
+
if index is None:
|
|
132
|
+
return torch.zeros(1, dtype=torch.long)
|
|
133
|
+
if isinstance(index, int):
|
|
134
|
+
return torch.tensor([index], dtype=torch.long)
|
|
135
|
+
idx = torch.as_tensor(index, dtype=torch.long)
|
|
136
|
+
return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
|