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,162 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
import torch.nn.functional as F
|
|
7
|
+
from torch import nn
|
|
8
|
+
|
|
9
|
+
from ...core.explanation import Explanation
|
|
10
|
+
from ...core.model_utils import capture_node_embeddings, edge_embeddings
|
|
11
|
+
from ...core.registry import register
|
|
12
|
+
from ..base import ExplanationAlgorithm
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class _EdgeMaskMLP(nn.Module):
|
|
16
|
+
def __init__(self, in_dim: int, hidden: int = 64):
|
|
17
|
+
super().__init__()
|
|
18
|
+
self.mlp = nn.Sequential(
|
|
19
|
+
nn.Linear(in_dim, hidden),
|
|
20
|
+
nn.ReLU(),
|
|
21
|
+
nn.Linear(hidden, hidden),
|
|
22
|
+
nn.ReLU(),
|
|
23
|
+
nn.Linear(hidden, 1),
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
def forward(self, edge_emb: torch.Tensor) -> torch.Tensor:
|
|
27
|
+
return self.mlp(edge_emb).squeeze(-1)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@register("pg_explainer", "pgexplainer")
|
|
31
|
+
class PGExplainer(ExplanationAlgorithm):
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
epochs: int = 100,
|
|
35
|
+
lr: float = 0.01,
|
|
36
|
+
hidden: int = 64,
|
|
37
|
+
temp: float = 1.0,
|
|
38
|
+
loss_coeff: float = 0.5,
|
|
39
|
+
entropy_coeff: float = 0.005,
|
|
40
|
+
batch_nodes: int = 32,
|
|
41
|
+
**kwargs,
|
|
42
|
+
):
|
|
43
|
+
self.epochs = epochs
|
|
44
|
+
self.lr = lr
|
|
45
|
+
self.hidden = hidden
|
|
46
|
+
self.temp = temp
|
|
47
|
+
self.loss_coeff = loss_coeff
|
|
48
|
+
self.entropy_coeff = entropy_coeff
|
|
49
|
+
self.batch_nodes = batch_nodes
|
|
50
|
+
|
|
51
|
+
def explain(
|
|
52
|
+
self,
|
|
53
|
+
backend: Any,
|
|
54
|
+
model: Any,
|
|
55
|
+
data: Any,
|
|
56
|
+
index: int | torch.Tensor | None,
|
|
57
|
+
target_class: int | None = None,
|
|
58
|
+
**kwargs,
|
|
59
|
+
) -> Explanation:
|
|
60
|
+
model.eval()
|
|
61
|
+
x = backend.node_features(data)
|
|
62
|
+
edge_index = backend.edge_index(data)
|
|
63
|
+
num_nodes = backend.num_nodes(data)
|
|
64
|
+
|
|
65
|
+
embeddings = capture_node_embeddings(model, backend, data)
|
|
66
|
+
edge_emb = edge_embeddings(embeddings, edge_index).detach()
|
|
67
|
+
in_dim = edge_emb.size(-1)
|
|
68
|
+
|
|
69
|
+
with torch.no_grad():
|
|
70
|
+
base_logits = backend.forward(model, x, edge_index)
|
|
71
|
+
|
|
72
|
+
mlp = _EdgeMaskMLP(in_dim, hidden=self.hidden)
|
|
73
|
+
optimizer = torch.optim.Adam(mlp.parameters(), lr=self.lr)
|
|
74
|
+
|
|
75
|
+
train_idx = (
|
|
76
|
+
torch.arange(num_nodes, device=x.device)
|
|
77
|
+
if not hasattr(data, "train_mask")
|
|
78
|
+
else data.train_mask.nonzero(as_tuple=False).view(-1)
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
for epoch in range(self.epochs):
|
|
82
|
+
mlp.train()
|
|
83
|
+
optimizer.zero_grad()
|
|
84
|
+
idx = train_idx[
|
|
85
|
+
torch.randperm(len(train_idx), device=x.device)[: self.batch_nodes]
|
|
86
|
+
]
|
|
87
|
+
edge_weight = self._sample_mask(mlp, edge_emb, x.device)
|
|
88
|
+
pred = backend.forward(model, x, edge_index, edge_weight=edge_weight)
|
|
89
|
+
if pred.dim() == 2:
|
|
90
|
+
target = base_logits[idx].argmax(dim=-1)
|
|
91
|
+
ce = F.nll_loss(pred[idx].log_softmax(dim=-1), target)
|
|
92
|
+
else:
|
|
93
|
+
ce = -pred[idx].mean()
|
|
94
|
+
logit = mlp(edge_emb)
|
|
95
|
+
p = torch.sigmoid(logit)
|
|
96
|
+
entropy = (
|
|
97
|
+
-(p * torch.clamp(p, 1e-8, 1).log())
|
|
98
|
+
- ((1 - p) * torch.clamp(1 - p, 1e-8, 1).log())
|
|
99
|
+
).mean()
|
|
100
|
+
sparsity = edge_weight.mean()
|
|
101
|
+
loss = ce + self.loss_coeff * sparsity + self.entropy_coeff * entropy
|
|
102
|
+
loss.backward()
|
|
103
|
+
optimizer.step()
|
|
104
|
+
|
|
105
|
+
mlp.eval()
|
|
106
|
+
with torch.no_grad():
|
|
107
|
+
edge_mask = torch.sigmoid(mlp(edge_emb))
|
|
108
|
+
|
|
109
|
+
ni = (
|
|
110
|
+
0
|
|
111
|
+
if index is None
|
|
112
|
+
else (int(index[0]) if torch.is_tensor(index) else int(index))
|
|
113
|
+
)
|
|
114
|
+
pred_masked = backend.forward(model, x, edge_index, edge_weight=edge_mask)
|
|
115
|
+
pred_node = (
|
|
116
|
+
pred_masked[ni] if pred_masked.dim() == 2 else pred_masked.unsqueeze(0)
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
node_importance = self._node_importance(edge_mask, edge_index, num_nodes)
|
|
120
|
+
|
|
121
|
+
return Explanation(
|
|
122
|
+
node_importance=node_importance,
|
|
123
|
+
edge_importance=edge_mask.detach().cpu(),
|
|
124
|
+
feature_importance=None,
|
|
125
|
+
prediction_original=(
|
|
126
|
+
base_logits[ni].detach().reshape(1, -1).cpu()
|
|
127
|
+
if base_logits.dim() == 2
|
|
128
|
+
else base_logits.detach().reshape(1, -1).cpu()
|
|
129
|
+
),
|
|
130
|
+
prediction_explanation=pred_node.detach().reshape(1, -1).cpu(),
|
|
131
|
+
node_idx=None if index is None else ni,
|
|
132
|
+
target_class=target_class,
|
|
133
|
+
metadata={},
|
|
134
|
+
)
|
|
135
|
+
|
|
136
|
+
def _sample_mask(
|
|
137
|
+
self,
|
|
138
|
+
mlp: _EdgeMaskMLP,
|
|
139
|
+
edge_emb: torch.Tensor,
|
|
140
|
+
device: torch.device,
|
|
141
|
+
) -> torch.Tensor:
|
|
142
|
+
logits = mlp(edge_emb)
|
|
143
|
+
noise = torch.rand_like(logits)
|
|
144
|
+
gumbel_noise = torch.log(noise) - torch.log(1 - noise + 1e-8)
|
|
145
|
+
return torch.sigmoid((logits + gumbel_noise) / self.temp)
|
|
146
|
+
|
|
147
|
+
@staticmethod
|
|
148
|
+
def _node_importance(
|
|
149
|
+
edge_mask: torch.Tensor,
|
|
150
|
+
edge_index: torch.Tensor,
|
|
151
|
+
num_nodes: int,
|
|
152
|
+
) -> torch.Tensor:
|
|
153
|
+
w = edge_mask.detach()
|
|
154
|
+
src, dst = edge_index[0], edge_index[1]
|
|
155
|
+
node_imp = torch.zeros(num_nodes, dtype=torch.float32, device=edge_mask.device)
|
|
156
|
+
node_imp.scatter_add_(0, src, w)
|
|
157
|
+
node_imp.scatter_add_(0, dst, w)
|
|
158
|
+
degree = torch.zeros(num_nodes, dtype=torch.float32, device=edge_mask.device)
|
|
159
|
+
degree.scatter_add_(0, src, torch.ones_like(src, dtype=torch.float32))
|
|
160
|
+
degree.scatter_add_(0, dst, torch.ones_like(dst, dtype=torch.float32))
|
|
161
|
+
degree.clamp_(min=1)
|
|
162
|
+
return (node_imp / degree).cpu()
|
|
@@ -0,0 +1,393 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import torch
|
|
7
|
+
|
|
8
|
+
from ...core.explanation import Explanation
|
|
9
|
+
from ...core.registry import register
|
|
10
|
+
from ..base import ExplanationAlgorithm
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class MCTSNode:
|
|
14
|
+
__slots__ = ("children", "parent", "reward", "state", "value", "visits")
|
|
15
|
+
|
|
16
|
+
def __init__(self, state, parent=None):
|
|
17
|
+
self.state = frozenset(state)
|
|
18
|
+
self.parent = parent
|
|
19
|
+
self.children: dict[int, MCTSNode] = {}
|
|
20
|
+
self.visits = 0
|
|
21
|
+
self.value = 0.0
|
|
22
|
+
self.reward = -math.inf
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@register("subgraphx", "subgraph_x")
|
|
26
|
+
class SubgraphX(ExplanationAlgorithm):
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
num_hops: int = 3,
|
|
30
|
+
rollout: int = 30,
|
|
31
|
+
high_cpu: bool = False,
|
|
32
|
+
reward_method: str = "mi",
|
|
33
|
+
lambda_connect: float = 0.5,
|
|
34
|
+
lambda_size: float = 0.05,
|
|
35
|
+
max_nodes: int = 20,
|
|
36
|
+
prune: bool = True,
|
|
37
|
+
seed: int = 0,
|
|
38
|
+
**kwargs,
|
|
39
|
+
):
|
|
40
|
+
self.num_hops = num_hops
|
|
41
|
+
self.rollout = rollout
|
|
42
|
+
self.high_cpu = high_cpu
|
|
43
|
+
self.reward_method = reward_method
|
|
44
|
+
self.lambda_connect = lambda_connect
|
|
45
|
+
self.lambda_size = lambda_size
|
|
46
|
+
self.max_nodes = max_nodes
|
|
47
|
+
self.prune = prune
|
|
48
|
+
self.seed = seed
|
|
49
|
+
|
|
50
|
+
def explain(
|
|
51
|
+
self,
|
|
52
|
+
backend: Any,
|
|
53
|
+
model: Any,
|
|
54
|
+
data: Any,
|
|
55
|
+
index: int | torch.Tensor | None,
|
|
56
|
+
target_class: int | None = None,
|
|
57
|
+
**kwargs,
|
|
58
|
+
) -> Explanation:
|
|
59
|
+
if index is None:
|
|
60
|
+
raise ValueError("SubgraphX requiere un nodo de interés (index)")
|
|
61
|
+
model.eval()
|
|
62
|
+
x = backend.node_features(data)
|
|
63
|
+
edge_index = backend.edge_index(data)
|
|
64
|
+
num_nodes = backend.num_nodes(data)
|
|
65
|
+
|
|
66
|
+
ni = int(index[0]) if torch.is_tensor(index) else int(index)
|
|
67
|
+
device = x.device
|
|
68
|
+
baseline = x.mean(dim=0, keepdim=True)
|
|
69
|
+
|
|
70
|
+
logits = backend.forward(model, x, edge_index)
|
|
71
|
+
if target_class is None:
|
|
72
|
+
target_class = int(logits[ni].argmax().item()) if logits.dim() == 2 else 0
|
|
73
|
+
|
|
74
|
+
candidates = self._neighborhood(
|
|
75
|
+
edge_index, num_nodes, ni, self.num_hops, self.max_nodes
|
|
76
|
+
)
|
|
77
|
+
if not candidates:
|
|
78
|
+
candidates = list(range(num_nodes))
|
|
79
|
+
candidates.remove(ni)
|
|
80
|
+
|
|
81
|
+
best_state, _ = self._mcts(
|
|
82
|
+
backend, model, x, edge_index, ni, candidates, target_class, baseline
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
selected = set(best_state)
|
|
86
|
+
selected = self._grow_if_unfaithful(
|
|
87
|
+
backend,
|
|
88
|
+
model,
|
|
89
|
+
x,
|
|
90
|
+
edge_index,
|
|
91
|
+
ni,
|
|
92
|
+
target_class,
|
|
93
|
+
selected,
|
|
94
|
+
candidates,
|
|
95
|
+
baseline,
|
|
96
|
+
)
|
|
97
|
+
if self.prune:
|
|
98
|
+
selected = self._prune(
|
|
99
|
+
backend, model, x, edge_index, ni, target_class, selected, baseline
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
edge_mask = self._edge_mask(edge_index, selected, num_nodes, device)
|
|
103
|
+
node_importance = self._node_importance(edge_index, selected, num_nodes, device)
|
|
104
|
+
|
|
105
|
+
with torch.no_grad():
|
|
106
|
+
x_eff = self._masked_features(x, selected, baseline)
|
|
107
|
+
pred_masked = backend.forward(
|
|
108
|
+
model, x_eff, edge_index, edge_weight=edge_mask
|
|
109
|
+
)
|
|
110
|
+
pred_node = (
|
|
111
|
+
pred_masked[ni] if pred_masked.dim() == 2 else pred_masked.unsqueeze(0)
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
return Explanation(
|
|
115
|
+
node_importance=node_importance.cpu(),
|
|
116
|
+
edge_importance=edge_mask.cpu(),
|
|
117
|
+
feature_importance=None,
|
|
118
|
+
prediction_original=(
|
|
119
|
+
logits[ni].detach().reshape(1, -1).cpu()
|
|
120
|
+
if logits.dim() == 2
|
|
121
|
+
else logits.detach().reshape(1, -1).cpu()
|
|
122
|
+
),
|
|
123
|
+
prediction_explanation=pred_node.detach().reshape(1, -1).cpu(),
|
|
124
|
+
node_idx=ni,
|
|
125
|
+
target_class=target_class,
|
|
126
|
+
metadata={"selected": sorted(selected)},
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
def _neighborhood(self, edge_index, num_nodes, node, hops, max_nodes) -> list[int]:
|
|
130
|
+
visited = {node}
|
|
131
|
+
frontier = {node}
|
|
132
|
+
for _ in range(hops):
|
|
133
|
+
nxt: set[int] = set()
|
|
134
|
+
for u in frontier:
|
|
135
|
+
mask = (edge_index[0] == u) | (edge_index[1] == u)
|
|
136
|
+
nxt.update(edge_index[:, mask].flatten().tolist())
|
|
137
|
+
frontier = nxt - visited
|
|
138
|
+
visited |= frontier
|
|
139
|
+
ordered = sorted(visited)
|
|
140
|
+
if len(ordered) > max_nodes:
|
|
141
|
+
gated = set(ordered[:max_nodes])
|
|
142
|
+
gated.add(node)
|
|
143
|
+
ordered = sorted(gated)
|
|
144
|
+
return ordered
|
|
145
|
+
|
|
146
|
+
def _mcts(
|
|
147
|
+
self, backend, model, x, edge_index, node, candidates, target_class, baseline
|
|
148
|
+
):
|
|
149
|
+
root_state = (node,)
|
|
150
|
+
root = MCTSNode(root_state)
|
|
151
|
+
for _ in range(self.rollout):
|
|
152
|
+
leaf, path = self._select(root, node, candidates)
|
|
153
|
+
reward = self._evaluate(
|
|
154
|
+
backend, model, x, edge_index, leaf.state, node, target_class, baseline
|
|
155
|
+
)
|
|
156
|
+
self._backprop(path, reward)
|
|
157
|
+
best_leaf = self._best_child(root)
|
|
158
|
+
return best_leaf.state, best_leaf.reward
|
|
159
|
+
|
|
160
|
+
def _select(self, root, node, candidates):
|
|
161
|
+
current = root
|
|
162
|
+
path = [root]
|
|
163
|
+
while current.children:
|
|
164
|
+
unexplored = [
|
|
165
|
+
n
|
|
166
|
+
for n in candidates
|
|
167
|
+
if n not in current.state and n not in current.children
|
|
168
|
+
]
|
|
169
|
+
if unexplored:
|
|
170
|
+
nxt_node = unexplored[0]
|
|
171
|
+
new_state = tuple(sorted(current.state | {nxt_node}))
|
|
172
|
+
child = MCTSNode(new_state, parent=current)
|
|
173
|
+
current.children[nxt_node] = child
|
|
174
|
+
path.append(child)
|
|
175
|
+
return child, path
|
|
176
|
+
best_child, best_score = None, -math.inf
|
|
177
|
+
for child in current.children.values():
|
|
178
|
+
uct = (child.value / max(child.visits, 1)) + math.sqrt(
|
|
179
|
+
2 * math.log(max(root.visits, 1)) / (child.visits + 1)
|
|
180
|
+
)
|
|
181
|
+
if uct > best_score:
|
|
182
|
+
best_score, best_child = uct, child
|
|
183
|
+
if best_child is None or len(current.state) >= len(candidates):
|
|
184
|
+
return current, path
|
|
185
|
+
current = best_child
|
|
186
|
+
path.append(current)
|
|
187
|
+
return current, path
|
|
188
|
+
|
|
189
|
+
def _evaluate(
|
|
190
|
+
self, backend, model, x, edge_index, state, node, target_class, baseline
|
|
191
|
+
):
|
|
192
|
+
selected = set(state)
|
|
193
|
+
if len(selected) == 0:
|
|
194
|
+
return -1e6
|
|
195
|
+
weight = self._edge_mask(edge_index, selected, x.size(0), x.device)
|
|
196
|
+
x_eff = self._masked_features(x, selected, baseline)
|
|
197
|
+
with torch.no_grad():
|
|
198
|
+
pred = backend.forward(model, x_eff, edge_index, edge_weight=weight)
|
|
199
|
+
if pred.dim() == 2:
|
|
200
|
+
logits_n = pred[node]
|
|
201
|
+
logp = logits_n.log_softmax(-1)[target_class].item()
|
|
202
|
+
if int(logits_n.argmax().item()) != target_class:
|
|
203
|
+
logp -= 20.0
|
|
204
|
+
else:
|
|
205
|
+
logp = float(pred[node])
|
|
206
|
+
components = self._connectivity(edge_index, selected)
|
|
207
|
+
size_penalty = self.lambda_size * (len(selected) - 1)
|
|
208
|
+
return logp - self.lambda_connect * (components - 1) - size_penalty
|
|
209
|
+
|
|
210
|
+
def _grow_if_unfaithful(
|
|
211
|
+
self,
|
|
212
|
+
backend,
|
|
213
|
+
model,
|
|
214
|
+
x,
|
|
215
|
+
edge_index,
|
|
216
|
+
node,
|
|
217
|
+
target_class,
|
|
218
|
+
selected,
|
|
219
|
+
candidates,
|
|
220
|
+
baseline,
|
|
221
|
+
):
|
|
222
|
+
selected = set(selected)
|
|
223
|
+
if not self._preserves_class(
|
|
224
|
+
backend, model, x, edge_index, node, target_class, selected, baseline
|
|
225
|
+
):
|
|
226
|
+
remaining = [c for c in candidates if c not in selected]
|
|
227
|
+
added = set()
|
|
228
|
+
while remaining:
|
|
229
|
+
best, best_reward = None, -math.inf
|
|
230
|
+
for u in remaining:
|
|
231
|
+
trial = selected | {u}
|
|
232
|
+
weight = self._edge_mask(edge_index, trial, x.size(0), x.device)
|
|
233
|
+
x_eff = self._masked_features(x, trial, baseline)
|
|
234
|
+
with torch.no_grad():
|
|
235
|
+
pred = backend.forward(
|
|
236
|
+
model, x_eff, edge_index, edge_weight=weight
|
|
237
|
+
)
|
|
238
|
+
if pred.dim() != 2:
|
|
239
|
+
continue
|
|
240
|
+
r = pred[node].log_softmax(-1)[target_class].item()
|
|
241
|
+
if int(pred[node].argmax().item()) != target_class:
|
|
242
|
+
r -= 20.0
|
|
243
|
+
if r > best_reward:
|
|
244
|
+
best_reward, best = r, u
|
|
245
|
+
if best is None:
|
|
246
|
+
break
|
|
247
|
+
selected.add(best)
|
|
248
|
+
remaining.remove(best)
|
|
249
|
+
added.add(best)
|
|
250
|
+
if self._preserves_class(
|
|
251
|
+
backend,
|
|
252
|
+
model,
|
|
253
|
+
x,
|
|
254
|
+
edge_index,
|
|
255
|
+
node,
|
|
256
|
+
target_class,
|
|
257
|
+
selected,
|
|
258
|
+
baseline,
|
|
259
|
+
):
|
|
260
|
+
break
|
|
261
|
+
return selected
|
|
262
|
+
|
|
263
|
+
def _preserves_class(
|
|
264
|
+
self, backend, model, x, edge_index, node, target_class, selected, baseline
|
|
265
|
+
) -> bool:
|
|
266
|
+
if len(selected) == 0:
|
|
267
|
+
return False
|
|
268
|
+
weight = self._edge_mask(edge_index, selected, x.size(0), x.device)
|
|
269
|
+
x_eff = self._masked_features(x, selected, baseline)
|
|
270
|
+
with torch.no_grad():
|
|
271
|
+
pred = backend.forward(model, x_eff, edge_index, edge_weight=weight)
|
|
272
|
+
if pred.dim() != 2:
|
|
273
|
+
return True
|
|
274
|
+
return (
|
|
275
|
+
not (node >= pred.size(0))
|
|
276
|
+
and int(pred[node].argmax().item()) == target_class
|
|
277
|
+
)
|
|
278
|
+
|
|
279
|
+
@staticmethod
|
|
280
|
+
def _masked_features(x, selected, baseline):
|
|
281
|
+
sel = torch.zeros(x.size(0), dtype=torch.bool, device=x.device)
|
|
282
|
+
sel[list(selected)] = True
|
|
283
|
+
return baseline + (x - baseline) * sel.unsqueeze(-1)
|
|
284
|
+
|
|
285
|
+
def _prune(
|
|
286
|
+
self, backend, model, x, edge_index, node, target_class, selected, baseline
|
|
287
|
+
):
|
|
288
|
+
selected = set(selected)
|
|
289
|
+
weight = self._edge_mask(edge_index, selected, x.size(0), x.device)
|
|
290
|
+
x_eff = self._masked_features(x, selected, baseline)
|
|
291
|
+
with torch.no_grad():
|
|
292
|
+
pred = backend.forward(model, x_eff, edge_index, edge_weight=weight)
|
|
293
|
+
if pred.dim() != 2:
|
|
294
|
+
return selected
|
|
295
|
+
if int(pred[node].argmax().item()) != target_class:
|
|
296
|
+
return selected
|
|
297
|
+
keep = sorted(selected)
|
|
298
|
+
if len(keep) > 1:
|
|
299
|
+
ordered = self._removal_order(
|
|
300
|
+
backend, model, x, edge_index, node, target_class, keep, baseline
|
|
301
|
+
)
|
|
302
|
+
for u in ordered:
|
|
303
|
+
trial = set(keep) - {u}
|
|
304
|
+
if u == node or not trial:
|
|
305
|
+
continue
|
|
306
|
+
t_weight = self._edge_mask(edge_index, trial, x.size(0), x.device)
|
|
307
|
+
t_x = self._masked_features(x, trial, baseline)
|
|
308
|
+
with torch.no_grad():
|
|
309
|
+
t_pred = backend.forward(
|
|
310
|
+
model, t_x, edge_index, edge_weight=t_weight
|
|
311
|
+
)
|
|
312
|
+
if int(t_pred[node].argmax().item()) == target_class:
|
|
313
|
+
keep = list(trial)
|
|
314
|
+
return set(keep)
|
|
315
|
+
|
|
316
|
+
def _removal_order(
|
|
317
|
+
self, backend, model, x, edge_index, node, target_class, keep, baseline
|
|
318
|
+
):
|
|
319
|
+
scores = []
|
|
320
|
+
for u in keep:
|
|
321
|
+
if u == node:
|
|
322
|
+
continue
|
|
323
|
+
trial = set(keep) - {u}
|
|
324
|
+
t_weight = self._edge_mask(edge_index, trial, x.size(0), x.device)
|
|
325
|
+
t_x = self._masked_features(x, trial, baseline)
|
|
326
|
+
with torch.no_grad():
|
|
327
|
+
t_pred = backend.forward(model, t_x, edge_index, edge_weight=t_weight)
|
|
328
|
+
if int(t_pred[node].argmax().item()) != target_class:
|
|
329
|
+
scores.append((u, -math.inf))
|
|
330
|
+
else:
|
|
331
|
+
logp_drop = t_pred[node].log_softmax(-1)[target_class].item()
|
|
332
|
+
scores.append((u, logp_drop))
|
|
333
|
+
return [u for u, _ in sorted(scores, key=lambda t: t[1])]
|
|
334
|
+
|
|
335
|
+
def _backprop(self, path, reward):
|
|
336
|
+
for n in reversed(path):
|
|
337
|
+
n.visits += 1
|
|
338
|
+
n.value += reward
|
|
339
|
+
n.reward = max(n.reward, reward)
|
|
340
|
+
|
|
341
|
+
def _best_child(self, root):
|
|
342
|
+
best = root
|
|
343
|
+
queue = [root]
|
|
344
|
+
while queue:
|
|
345
|
+
cur = queue.pop(0)
|
|
346
|
+
if cur.reward > best.reward:
|
|
347
|
+
best = cur
|
|
348
|
+
queue.extend(cur.children.values())
|
|
349
|
+
return best
|
|
350
|
+
|
|
351
|
+
@staticmethod
|
|
352
|
+
def _edge_mask(edge_index, selected, num_nodes, device):
|
|
353
|
+
src = edge_index[0]
|
|
354
|
+
dst = edge_index[1]
|
|
355
|
+
both = torch.isin(
|
|
356
|
+
src, torch.as_tensor(list(selected), device=device)
|
|
357
|
+
) & torch.isin(dst, torch.as_tensor(list(selected), device=device))
|
|
358
|
+
return both.to(torch.float32)
|
|
359
|
+
|
|
360
|
+
@staticmethod
|
|
361
|
+
def _node_importance(edge_index, selected, num_nodes, device):
|
|
362
|
+
imp = torch.zeros(num_nodes, dtype=torch.float32, device=device)
|
|
363
|
+
sel = torch.as_tensor(list(selected), device=device)
|
|
364
|
+
for s in sel:
|
|
365
|
+
imp[s] = 1.0
|
|
366
|
+
return imp
|
|
367
|
+
|
|
368
|
+
@staticmethod
|
|
369
|
+
def _connectivity(edge_index, selected):
|
|
370
|
+
sel = set(selected)
|
|
371
|
+
seen = set()
|
|
372
|
+
comps = 0
|
|
373
|
+
adj = {}
|
|
374
|
+
src = edge_index[0].tolist()
|
|
375
|
+
dst = edge_index[1].tolist()
|
|
376
|
+
for u, v in zip(src, dst):
|
|
377
|
+
if u in sel and v in sel:
|
|
378
|
+
adj.setdefault(u, set()).add(v)
|
|
379
|
+
adj.setdefault(v, set()).add(u)
|
|
380
|
+
for s in sel:
|
|
381
|
+
if s in seen:
|
|
382
|
+
continue
|
|
383
|
+
comps += 1
|
|
384
|
+
stack = [s]
|
|
385
|
+
while stack:
|
|
386
|
+
cur = stack.pop()
|
|
387
|
+
if cur in seen:
|
|
388
|
+
continue
|
|
389
|
+
seen.add(cur)
|
|
390
|
+
for nb in adj.get(cur, ()):
|
|
391
|
+
if nb not in seen:
|
|
392
|
+
stack.append(nb)
|
|
393
|
+
return comps
|