D4CMPP2 0.4.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.
- D4CMPP2/_Data/AGENTS.md +24 -0
- D4CMPP2/_Data/Aqsoldb.csv +9291 -0
- D4CMPP2/_Data/BradleyMP.csv +3042 -0
- D4CMPP2/_Data/Lipophilicity.csv +1131 -0
- D4CMPP2/_Data/README.md +8 -0
- D4CMPP2/_Data/__init__.py +26 -0
- D4CMPP2/_Data/optical.csv +20237 -0
- D4CMPP2/_Data/test.csv +190 -0
- D4CMPP2/__init__.py +16 -0
- D4CMPP2/__main__.py +5 -0
- D4CMPP2/_main.py +500 -0
- D4CMPP2/cli.py +7 -0
- D4CMPP2/exceptions.py +53 -0
- D4CMPP2/grid_search.py +259 -0
- D4CMPP2/network_refer.yaml +160 -0
- D4CMPP2/networks/AFP_model.py +72 -0
- D4CMPP2/networks/AFPwithSolv_model.py +72 -0
- D4CMPP2/networks/DMPNN_model.py +90 -0
- D4CMPP2/networks/DMPNNwithSolv_model.py +89 -0
- D4CMPP2/networks/GAT_model.py +48 -0
- D4CMPP2/networks/GATwithSolv_model.py +63 -0
- D4CMPP2/networks/GCN_model.py +113 -0
- D4CMPP2/networks/GCNwithSolv_model.py +103 -0
- D4CMPP2/networks/GC_model.py +122 -0
- D4CMPP2/networks/ISATPM_model.py +14 -0
- D4CMPP2/networks/ISATPN_model.py +199 -0
- D4CMPP2/networks/ISAT_model.py +90 -0
- D4CMPP2/networks/MPNN_model.py +56 -0
- D4CMPP2/networks/MPNNwithSolv_model.py +72 -0
- D4CMPP2/networks/__init__.py +25 -0
- D4CMPP2/networks/base.py +250 -0
- D4CMPP2/networks/registry.py +187 -0
- D4CMPP2/networks/src/AFP.py +118 -0
- D4CMPP2/networks/src/BiDropout.py +29 -0
- D4CMPP2/networks/src/DMPNN.py +35 -0
- D4CMPP2/networks/src/GAT.py +69 -0
- D4CMPP2/networks/src/GC.py +85 -0
- D4CMPP2/networks/src/GCN.py +71 -0
- D4CMPP2/networks/src/ISAT.py +153 -0
- D4CMPP2/networks/src/Linear.py +49 -0
- D4CMPP2/networks/src/MPNN.py +56 -0
- D4CMPP2/networks/src/SolventLayer.py +62 -0
- D4CMPP2/networks/src/__init__.py +0 -0
- D4CMPP2/networks/src/distGCN.py +21 -0
- D4CMPP2/networks/src/pyg_hetero.py +24 -0
- D4CMPP2/optimize.py +472 -0
- D4CMPP2/src/Analyzer/ISAAnalyzer.py +458 -0
- D4CMPP2/src/Analyzer/ISAPNAnalyzer.py +366 -0
- D4CMPP2/src/Analyzer/ISAwSAnalyzer.py +117 -0
- D4CMPP2/src/Analyzer/MolAnalyzer.py +319 -0
- D4CMPP2/src/Analyzer/__init__.py +54 -0
- D4CMPP2/src/Analyzer/core.py +480 -0
- D4CMPP2/src/Analyzer/factory.py +166 -0
- D4CMPP2/src/Analyzer/interpretation.py +232 -0
- D4CMPP2/src/Analyzer/results.py +101 -0
- D4CMPP2/src/DataManager/Dataset/GraphDataset.py +314 -0
- D4CMPP2/src/DataManager/Dataset/ISAGraphDataset.py +384 -0
- D4CMPP2/src/DataManager/Dataset/__init__.py +0 -0
- D4CMPP2/src/DataManager/GraphGenerator/ISAGraphGenerator.py +223 -0
- D4CMPP2/src/DataManager/GraphGenerator/MolGraphGenerator.py +73 -0
- D4CMPP2/src/DataManager/GraphGenerator/__init__.py +14 -0
- D4CMPP2/src/DataManager/ISADataManager.py +67 -0
- D4CMPP2/src/DataManager/MolDataManager.py +735 -0
- D4CMPP2/src/DataManager/__init__.py +14 -0
- D4CMPP2/src/DataManager/contracts.py +179 -0
- D4CMPP2/src/NetworkManager/ISANetworkManager.py +12 -0
- D4CMPP2/src/NetworkManager/NetworkManager.py +520 -0
- D4CMPP2/src/NetworkManager/__init__.py +14 -0
- D4CMPP2/src/PostProcessor.py +160 -0
- D4CMPP2/src/TrainManager/ISATrainManager.py +26 -0
- D4CMPP2/src/TrainManager/TrainManager.py +254 -0
- D4CMPP2/src/TrainManager/__init__.py +14 -0
- D4CMPP2/src/TrainManager/callbacks.py +119 -0
- D4CMPP2/src/__init__.py +0 -0
- D4CMPP2/src/utils/PATH.py +246 -0
- D4CMPP2/src/utils/__init__.py +0 -0
- D4CMPP2/src/utils/argparser.py +56 -0
- D4CMPP2/src/utils/checkpointing.py +90 -0
- D4CMPP2/src/utils/config_resolution.py +123 -0
- D4CMPP2/src/utils/config_validation.py +370 -0
- D4CMPP2/src/utils/csv_validation.py +105 -0
- D4CMPP2/src/utils/data_quality.py +181 -0
- D4CMPP2/src/utils/featureizer.py +202 -0
- D4CMPP2/src/utils/functional_group.csv +169 -0
- D4CMPP2/src/utils/graph_cache.py +213 -0
- D4CMPP2/src/utils/leaderboard.py +212 -0
- D4CMPP2/src/utils/metrics.py +31 -0
- D4CMPP2/src/utils/module_loader.py +147 -0
- D4CMPP2/src/utils/output.py +80 -0
- D4CMPP2/src/utils/reproducibility.py +70 -0
- D4CMPP2/src/utils/run_manifest.py +175 -0
- D4CMPP2/src/utils/scaler.py +40 -0
- D4CMPP2/src/utils/sculptor.py +713 -0
- D4CMPP2/src/utils/splitting.py +250 -0
- D4CMPP2/src/utils/supportfile_saver.py +94 -0
- D4CMPP2/src/utils/tools.py +156 -0
- D4CMPP2/src/utils/transfer_learning.py +111 -0
- d4cmpp2-0.4.0.dist-info/METADATA +420 -0
- d4cmpp2-0.4.0.dist-info/RECORD +103 -0
- d4cmpp2-0.4.0.dist-info/WHEEL +5 -0
- d4cmpp2-0.4.0.dist-info/entry_points.txt +2 -0
- d4cmpp2-0.4.0.dist-info/licenses/LICENSE +21 -0
- d4cmpp2-0.4.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,384 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import warnings
|
|
3
|
+
from torch.utils.data import Dataset
|
|
4
|
+
from torch_geometric.data import Batch
|
|
5
|
+
import numpy as np
|
|
6
|
+
from ..contracts import ISA_BATCH_CONTRACT, LEGACY_ISA_BATCH_CONTRACT
|
|
7
|
+
|
|
8
|
+
class ISAGraphDataset_legacy(Dataset):
|
|
9
|
+
batch_contract = LEGACY_ISA_BATCH_CONTRACT
|
|
10
|
+
def __init__(self, graphs=None, target=None, smiles=None):
|
|
11
|
+
if graphs is None: return
|
|
12
|
+
self.graphs = graphs
|
|
13
|
+
self.r_node = [g['r_nd'].x for g in graphs]
|
|
14
|
+
self.r2r_edge = [g['r_nd', 'r2r', 'r_nd'].edge_attr for g in graphs]
|
|
15
|
+
self.i_node = [g['i_nd'].x for g in graphs]
|
|
16
|
+
self.d_node = [g['d_nd'].x for g in graphs]
|
|
17
|
+
self.d2d_edge = [g['d_nd', 'd2d', 'd_nd'].edge_attr for g in graphs]
|
|
18
|
+
self.target = torch.tensor(target).float()
|
|
19
|
+
self.smiles = smiles
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def __len__(self):
|
|
23
|
+
return len(self.graphs)
|
|
24
|
+
|
|
25
|
+
def __getitem__(self, idx):
|
|
26
|
+
return self.graphs[idx], self.r_node[idx], self.r2r_edge[idx], self.i_node[idx], self.d_node[idx], self.d2d_edge[idx], self.target[idx], self.smiles[idx]
|
|
27
|
+
|
|
28
|
+
def reload(self, data):
|
|
29
|
+
self.graphs, self.r_node, self.r2r_edge, self.i_node, self.d_node, self.d2d_edge, self.target, self.smiles = data
|
|
30
|
+
self.target = torch.stack(self.target)
|
|
31
|
+
if self.target.dim() == 1:
|
|
32
|
+
self.target = self.target.unsqueeze(-1)
|
|
33
|
+
|
|
34
|
+
# def subDataset(self, idx):
|
|
35
|
+
# self.graphs = [self.graphs[i] for i in idx]
|
|
36
|
+
# self.r_node = [self.r_node[i] for i in idx]
|
|
37
|
+
# self.r2r_edge = [self.r2r_edge[i] for i in idx]
|
|
38
|
+
# self.i_node = [self.i_node[i] for i in idx]
|
|
39
|
+
# self.d_node = [self.d_node[i] for i in idx]
|
|
40
|
+
# self.d2d_edge = [self.d2d_edge[i] for i in idx]
|
|
41
|
+
# self.target = self.target[np.array(idx, dtype=int)]
|
|
42
|
+
# self.smiles = [self.smiles[i] for i in idx]
|
|
43
|
+
|
|
44
|
+
def get_subDataset(self, idx):
|
|
45
|
+
graphs = [self.graphs[i] for i in idx]
|
|
46
|
+
r_node = [self.r_node[i] for i in idx]
|
|
47
|
+
r2r_edge = [self.r2r_edge[i] for i in idx]
|
|
48
|
+
i_node = [self.i_node[i] for i in idx]
|
|
49
|
+
d_node = [self.d_node[i] for i in idx]
|
|
50
|
+
d2d_edge = [self.d2d_edge[i] for i in idx]
|
|
51
|
+
target = [self.target[i] for i in idx]
|
|
52
|
+
smiles = [self.smiles[i] for i in idx]
|
|
53
|
+
|
|
54
|
+
dataset = ISAGraphDataset_legacy()
|
|
55
|
+
dataset.reload((graphs, r_node, r2r_edge, i_node, d_node, d2d_edge, target, smiles))
|
|
56
|
+
return dataset
|
|
57
|
+
|
|
58
|
+
@staticmethod
|
|
59
|
+
def collate(samples):
|
|
60
|
+
graphs, r_node, r2r_edge, i_node, d_node, d2d_edge, target, smiles = map(list, zip(*samples))
|
|
61
|
+
batched_graph = Batch.from_data_list(graphs)
|
|
62
|
+
return batched_graph, torch.concat(r_node,dim=0), torch.concat(r2r_edge,dim=0), torch.concat(i_node,dim=0), torch.concat(d_node,dim=0), torch.concat(d2d_edge,dim=0), torch.stack(target,dim=0), smiles
|
|
63
|
+
|
|
64
|
+
@staticmethod
|
|
65
|
+
def unwrapper(graph, r_node, r2r_edge, i_node, d_node, d2d_edge, target, smiles, device='cpu'):
|
|
66
|
+
graph = graph.to(device=device)
|
|
67
|
+
r_node = r_node.float().to(device=device)
|
|
68
|
+
r2r_edge = r2r_edge.float().to(device=device)
|
|
69
|
+
i_node = i_node.float().to(device=device)
|
|
70
|
+
d_node = d_node.float().to(device=device)
|
|
71
|
+
d2d_edge = d2d_edge.float().to(device=device)
|
|
72
|
+
target = target.float().to(device=device)
|
|
73
|
+
return {'graph':graph, 'r_node':r_node, 'r_edge':r2r_edge, 'i_node':i_node, 'd_node':d_node, 'd_edge':d2d_edge, 'target':target, 'smiles':smiles}
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class ISAGraphDataset(ISAGraphDataset_legacy):
|
|
77
|
+
batch_contract = ISA_BATCH_CONTRACT
|
|
78
|
+
def __init__(self, graphs=None, target=None, smiles=None, numeric_inputs=None, row_indices=None):
|
|
79
|
+
if graphs is None:
|
|
80
|
+
self.graphs = {}
|
|
81
|
+
self.target = None
|
|
82
|
+
self.smiles = {}
|
|
83
|
+
self.data_keys = []
|
|
84
|
+
return
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
for key in graphs:
|
|
88
|
+
setattr(self, key + '_graphs', graphs[key])
|
|
89
|
+
try:
|
|
90
|
+
setattr(self, key + '_r_node', [g['r_nd'].x for g in graphs[key]])
|
|
91
|
+
except KeyError:
|
|
92
|
+
warnings.warn(
|
|
93
|
+
f"Graph key {key!r} has no 'r_nd' node type; using zero features.",
|
|
94
|
+
RuntimeWarning,
|
|
95
|
+
stacklevel=2,
|
|
96
|
+
)
|
|
97
|
+
setattr(self, key + '_r_node', [torch.zeros((g['r_nd'].num_nodes, 1)) for g in graphs[key]])
|
|
98
|
+
|
|
99
|
+
try:
|
|
100
|
+
setattr(self, key + '_r2r_edge', [g['r_nd', 'r2r', 'r_nd'].edge_attr for g in graphs[key]])
|
|
101
|
+
except KeyError:
|
|
102
|
+
warnings.warn(
|
|
103
|
+
f"Graph key {key!r} has no 'r2r' edge type; using zero features.",
|
|
104
|
+
RuntimeWarning,
|
|
105
|
+
stacklevel=2,
|
|
106
|
+
)
|
|
107
|
+
setattr(self, key + '_r2r_edge', [torch.zeros((g['r_nd', 'r2r', 'r_nd'].num_edges, 1)) for g in graphs[key]])
|
|
108
|
+
# if 'i_nd' in graphs[key][0].nodes:
|
|
109
|
+
# setattr(self, key + '_i_node', [g.nodes['i_nd'].data['f'] for g in graphs[key]])
|
|
110
|
+
# else:
|
|
111
|
+
# setattr(self, key + '_i_node', [torch.zeros((g.num_nodes(), 0)) for g in graphs[key]])
|
|
112
|
+
try:
|
|
113
|
+
setattr(self, key + '_i2i_edge', [g['i_nd', 'i2i', 'i_nd'].edge_attr for g in graphs[key]])
|
|
114
|
+
except KeyError:
|
|
115
|
+
warnings.warn(
|
|
116
|
+
f"Graph key {key!r} has no 'i2i' edge type; using zero features.",
|
|
117
|
+
RuntimeWarning,
|
|
118
|
+
stacklevel=2,
|
|
119
|
+
)
|
|
120
|
+
setattr(self, key + '_i2i_edge', [torch.zeros((g['i_nd', 'i2i', 'i_nd'].num_edges, 1)) for g in graphs[key]])
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
# if 'i2i' in graphs[key][0].nodes:
|
|
125
|
+
# setattr(self, key + '_i_node', [g.nodes['i_nd'].data['f'] for g in graphs[key]])
|
|
126
|
+
# else:
|
|
127
|
+
# setattr(self, key + '_i_node', [torch.zeros((g.num_nodes(), 0)) for g in graphs[key]])
|
|
128
|
+
try:
|
|
129
|
+
setattr(self, key + '_i_node', [g['i_nd'].x for g in graphs[key]])
|
|
130
|
+
except KeyError:
|
|
131
|
+
warnings.warn(
|
|
132
|
+
f"Graph key {key!r} has no 'i_nd' node type; using zero features.",
|
|
133
|
+
RuntimeWarning,
|
|
134
|
+
stacklevel=2,
|
|
135
|
+
)
|
|
136
|
+
setattr(self, key + '_i_node', [torch.zeros((g['i_nd'].num_nodes, 1)) for g in graphs[key]])
|
|
137
|
+
|
|
138
|
+
try:
|
|
139
|
+
setattr(self, key + '_d_node', [g['d_nd'].x for g in graphs[key]])
|
|
140
|
+
except KeyError:
|
|
141
|
+
warnings.warn(
|
|
142
|
+
f"Graph key {key!r} has no 'd_nd' node type; using zero features.",
|
|
143
|
+
RuntimeWarning,
|
|
144
|
+
stacklevel=2,
|
|
145
|
+
)
|
|
146
|
+
setattr(self, key + '_d_node', [torch.zeros((g['d_nd'].num_nodes, 1)) for g in graphs[key]])
|
|
147
|
+
|
|
148
|
+
try:
|
|
149
|
+
setattr(self, key + '_d2d_edge', [g['d_nd', 'd2d', 'd_nd'].edge_attr for g in graphs[key]])
|
|
150
|
+
except KeyError:
|
|
151
|
+
warnings.warn(
|
|
152
|
+
f"Graph key {key!r} has no 'd2d' edge type; using zero features.",
|
|
153
|
+
RuntimeWarning,
|
|
154
|
+
stacklevel=2,
|
|
155
|
+
)
|
|
156
|
+
setattr(self, key + '_d2d_edge', [torch.zeros((g['d_nd', 'd2d', 'd_nd'].num_edges, 1)) for g in graphs[key]])
|
|
157
|
+
|
|
158
|
+
# if 'd_nd' in graphs[key][0].nodes:
|
|
159
|
+
# setattr(self, key + '_d_node', [g.nodes['d_nd'].data['f'] for g in graphs[key]])
|
|
160
|
+
# else:
|
|
161
|
+
# setattr(self, key + '_d_node', [torch.zeros((g.num_nodes(), 0)) for g in graphs[key]])
|
|
162
|
+
|
|
163
|
+
# if 'd2d' in graphs[key][0].edges:
|
|
164
|
+
# setattr(self, key + '_d2d_edge', [g.edges['d2d'].data['dist'] for g in graphs[key]])
|
|
165
|
+
# else:
|
|
166
|
+
# setattr(self, key + '_d2d_edge', [torch.zeros((g.num_edges(), 0)) for g in graphs[key]])
|
|
167
|
+
|
|
168
|
+
if numeric_inputs is not None:
|
|
169
|
+
for key in numeric_inputs:
|
|
170
|
+
setattr(self, key + '_var', numeric_inputs[key])
|
|
171
|
+
else:
|
|
172
|
+
numeric_inputs = {}
|
|
173
|
+
|
|
174
|
+
for key in smiles:
|
|
175
|
+
if key not in graphs:
|
|
176
|
+
raise ValueError(f"Key '{key}' in smiles is not found in graphs.")
|
|
177
|
+
setattr(self, key + '_smiles', smiles[key])
|
|
178
|
+
|
|
179
|
+
self.target = torch.tensor(target).float() if target is not None else None
|
|
180
|
+
self.original_row_index = (
|
|
181
|
+
torch.as_tensor(row_indices, dtype=torch.long) if row_indices is not None else None
|
|
182
|
+
)
|
|
183
|
+
self.data_keys = list(graphs.keys()) + list(numeric_inputs.keys()) + ['target']
|
|
184
|
+
|
|
185
|
+
def __len__(self):
|
|
186
|
+
if getattr(self, 'target', None) is not None:
|
|
187
|
+
return len(self.target)
|
|
188
|
+
else:
|
|
189
|
+
if len(self.data_keys) == 0:
|
|
190
|
+
return 0
|
|
191
|
+
key = self.data_keys[0]
|
|
192
|
+
if hasattr(self, key + '_graphs'):
|
|
193
|
+
return len(getattr(self, key + '_graphs'))
|
|
194
|
+
elif hasattr(self, key + '_var'):
|
|
195
|
+
return len(getattr(self, key + '_var'))
|
|
196
|
+
elif hasattr(self, key + '_smiles'):
|
|
197
|
+
return len(getattr(self, key + '_smiles'))
|
|
198
|
+
else:
|
|
199
|
+
return 0
|
|
200
|
+
|
|
201
|
+
def __getitem__(self, idx):
|
|
202
|
+
item = {}
|
|
203
|
+
for key in self.data_keys:
|
|
204
|
+
if hasattr(self, key + '_graphs'):
|
|
205
|
+
item[key + '_graphs'] = getattr(self, key + '_graphs')[idx]
|
|
206
|
+
if hasattr(self, key + '_smiles'):
|
|
207
|
+
item[key + '_smiles'] = getattr(self, key + '_smiles')[idx]
|
|
208
|
+
# if hasattr(self, key + '_node_feature'):
|
|
209
|
+
# nf = getattr(self, key + '_node_feature')[idx]
|
|
210
|
+
# if type(nf) is not torch.Tensor:
|
|
211
|
+
# nf = torch.tensor(nf, dtype=torch.float32)
|
|
212
|
+
# item[key + '_node_feature'] = nf
|
|
213
|
+
# if hasattr(self, key + '_edge_feature'):
|
|
214
|
+
# ef = getattr(self, key + '_edge_feature')[idx]
|
|
215
|
+
# if type(ef) is not torch.Tensor:
|
|
216
|
+
# ef = torch.tensor(ef, dtype=torch.float32)
|
|
217
|
+
# item[key + '_edge_feature'] = ef
|
|
218
|
+
if hasattr(self, key + '_r_node'):
|
|
219
|
+
f = getattr(self, key + '_r_node')[idx]
|
|
220
|
+
if type(f) is not torch.Tensor:
|
|
221
|
+
f = torch.tensor(f, dtype=torch.float32)
|
|
222
|
+
item[key + '_r_node'] = f
|
|
223
|
+
if hasattr(self, key + '_r2r_edge'):
|
|
224
|
+
f = getattr(self, key + '_r2r_edge')[idx]
|
|
225
|
+
if type(f) is not torch.Tensor:
|
|
226
|
+
f = torch.tensor(f, dtype=torch.float32)
|
|
227
|
+
item[key + '_r2r_edge'] = f
|
|
228
|
+
if hasattr(self, key + '_i_node'):
|
|
229
|
+
f = getattr(self, key + '_i_node')[idx]
|
|
230
|
+
if type(f) is not torch.Tensor:
|
|
231
|
+
f = torch.tensor(f, dtype=torch.float32)
|
|
232
|
+
item[key + '_i_node'] = f
|
|
233
|
+
if hasattr(self, key + '_i2i_edge'):
|
|
234
|
+
f = getattr(self, key + '_i2i_edge')[idx]
|
|
235
|
+
if type(f) is not torch.Tensor:
|
|
236
|
+
f = torch.tensor(f, dtype=torch.float32)
|
|
237
|
+
item[key + '_i2i_edge'] = f
|
|
238
|
+
if hasattr(self, key + '_d_node'):
|
|
239
|
+
f = getattr(self, key + '_d_node')[idx]
|
|
240
|
+
if type(f) is not torch.Tensor:
|
|
241
|
+
f = torch.tensor(f, dtype=torch.float32)
|
|
242
|
+
item[key + '_d_node'] = f
|
|
243
|
+
if hasattr(self, key + '_d2d_edge'):
|
|
244
|
+
f = getattr(self, key + '_d2d_edge')[idx]
|
|
245
|
+
if type(f) is not torch.Tensor:
|
|
246
|
+
f = torch.tensor(f, dtype=torch.float32)
|
|
247
|
+
item[key + '_d2d_edge'] = f
|
|
248
|
+
|
|
249
|
+
elif hasattr(self, key + '_var'):
|
|
250
|
+
var = getattr(self, key + '_var')[idx]
|
|
251
|
+
if torch.is_tensor(var):
|
|
252
|
+
var = var.float()
|
|
253
|
+
else:
|
|
254
|
+
var = torch.tensor(var, dtype=torch.float32)
|
|
255
|
+
item[key + '_var'] = var
|
|
256
|
+
elif hasattr(self, key + '_smiles'):
|
|
257
|
+
item[key + '_smiles'] = getattr(self, key + '_smiles')[idx]
|
|
258
|
+
if hasattr(self, 'target') and self.target is not None:
|
|
259
|
+
item['target'] = self.target[idx]
|
|
260
|
+
if self.original_row_index is not None:
|
|
261
|
+
item['original_row_index'] = self.original_row_index[idx]
|
|
262
|
+
return item
|
|
263
|
+
|
|
264
|
+
def reload(self, data):
|
|
265
|
+
if len(data) == 0:
|
|
266
|
+
self.target = None
|
|
267
|
+
self.data_keys = []
|
|
268
|
+
return
|
|
269
|
+
|
|
270
|
+
new_data_keys = []
|
|
271
|
+
for key in data[0]:
|
|
272
|
+
values = [d[key] for d in data]
|
|
273
|
+
if key == 'target':
|
|
274
|
+
if len(values) == 0:
|
|
275
|
+
setattr(self, key, None)
|
|
276
|
+
elif torch.is_tensor(values[0]):
|
|
277
|
+
target = torch.stack(values, dim=0).float()
|
|
278
|
+
if target.dim() == 1:
|
|
279
|
+
target = target.unsqueeze(-1)
|
|
280
|
+
setattr(self, key, target)
|
|
281
|
+
else:
|
|
282
|
+
target = torch.tensor(values, dtype=torch.float32)
|
|
283
|
+
if target.dim() == 1:
|
|
284
|
+
target = target.unsqueeze(-1)
|
|
285
|
+
setattr(self, key, target)
|
|
286
|
+
elif key == 'original_row_index':
|
|
287
|
+
setattr(self, key, torch.as_tensor(values, dtype=torch.long))
|
|
288
|
+
elif key.endswith('_var'):
|
|
289
|
+
if len(values) == 0:
|
|
290
|
+
setattr(self, key, torch.empty((0,), dtype=torch.float32))
|
|
291
|
+
elif torch.is_tensor(values[0]):
|
|
292
|
+
setattr(self, key, torch.stack(values, dim=0).float())
|
|
293
|
+
else:
|
|
294
|
+
setattr(self, key, torch.tensor(values, dtype=torch.float32))
|
|
295
|
+
else:
|
|
296
|
+
setattr(self, key, values)
|
|
297
|
+
if key.endswith('_graphs'):
|
|
298
|
+
new_data_keys.append(key[:-7]) # Remove '_graphs'
|
|
299
|
+
elif key.endswith('_var'):
|
|
300
|
+
new_data_keys.append(key[:-4]) # Remove '_var'
|
|
301
|
+
elif key == 'target':
|
|
302
|
+
new_data_keys.append(key)
|
|
303
|
+
|
|
304
|
+
self.data_keys = new_data_keys
|
|
305
|
+
|
|
306
|
+
|
|
307
|
+
def subDataset(self, idx):
|
|
308
|
+
for key in self.data_keys:
|
|
309
|
+
if hasattr(self, key + '_graphs'):
|
|
310
|
+
setattr(self, key + '_graphs', [getattr(self, key + '_graphs')[i] for i in idx])
|
|
311
|
+
# setattr(self, key + '_node_feature', [getattr(self, key + '_node_feature')[i] for i in idx])
|
|
312
|
+
# setattr(self, key + '_edge_feature', [getattr(self, key + '_edge_feature')[i] for i in idx])
|
|
313
|
+
if hasattr(self, key + '_r_node'):
|
|
314
|
+
setattr(self, key + '_r_node', [getattr(self, key + '_r_node')[i] for i in idx])
|
|
315
|
+
if hasattr(self, key + '_r2r_edge'):
|
|
316
|
+
setattr(self, key + '_r2r_edge', [getattr(self, key + '_r2r_edge')[i] for i in idx])
|
|
317
|
+
if hasattr(self, key + '_i_node'):
|
|
318
|
+
setattr(self, key + '_i_node', [getattr(self, key + '_i_node')[i] for i in idx])
|
|
319
|
+
if hasattr(self, key + '_i2i_edge'):
|
|
320
|
+
setattr(self, key + '_i2i_edge', [getattr(self, key + '_i2i_edge')[i] for i in idx])
|
|
321
|
+
if hasattr(self, key + '_d_node'):
|
|
322
|
+
setattr(self, key + '_d_node', [getattr(self, key + '_d_node')[i] for i in idx])
|
|
323
|
+
if hasattr(self, key + '_d2d_edge'):
|
|
324
|
+
setattr(self, key + '_d2d_edge', [getattr(self, key + '_d2d_edge')[i] for i in idx])
|
|
325
|
+
|
|
326
|
+
if hasattr(self, key + '_var'):
|
|
327
|
+
var = getattr(self, key + '_var')
|
|
328
|
+
if torch.is_tensor(var):
|
|
329
|
+
setattr(self, key + '_var', var[torch.tensor(idx, dtype=torch.long)])
|
|
330
|
+
else:
|
|
331
|
+
setattr(self, key + '_var', var[np.array(idx, dtype=int)])
|
|
332
|
+
if hasattr(self, key + '_smiles'):
|
|
333
|
+
setattr(self, key + '_smiles', [getattr(self, key + '_smiles')[i] for i in idx])
|
|
334
|
+
if self.target is not None:
|
|
335
|
+
self.target = self.target[np.array(idx, dtype=int)]
|
|
336
|
+
if self.original_row_index is not None:
|
|
337
|
+
self.original_row_index = self.original_row_index[torch.as_tensor(idx, dtype=torch.long)]
|
|
338
|
+
|
|
339
|
+
def get_subDataset(self, idx):
|
|
340
|
+
dataset = ISAGraphDataset()
|
|
341
|
+
dataset.reload([self[i] for i in idx])
|
|
342
|
+
return dataset
|
|
343
|
+
|
|
344
|
+
|
|
345
|
+
@staticmethod
|
|
346
|
+
def collate(samples):
|
|
347
|
+
batched_data = {}
|
|
348
|
+
for key in samples[0].keys():
|
|
349
|
+
if key.endswith('_graphs'):
|
|
350
|
+
batched_data[key] = Batch.from_data_list([s[key] for s in samples])
|
|
351
|
+
# elif key.endswith('_node_feature'):
|
|
352
|
+
# batched_data[key] = torch.concat([s[key] for s in samples], dim=0)
|
|
353
|
+
# elif key.endswith('_edge_feature'):
|
|
354
|
+
# batched_data[key] = torch.concat([s[key] for s in samples], dim=0)
|
|
355
|
+
elif key.endswith('_r_node') or key.endswith('_i_node') or key.endswith('_d_node'):
|
|
356
|
+
batched_data[key] = torch.concat([s[key] for s in samples], dim=0)
|
|
357
|
+
elif key.endswith('_r2r_edge') or key.endswith('_i2i_edge') or key.endswith('_d2d_edge'):
|
|
358
|
+
batched_data[key] = torch.concat([s[key] for s in samples], dim=0)
|
|
359
|
+
elif key.endswith('_var'):
|
|
360
|
+
batched_data[key] = torch.stack([s[key].reshape(-1) for s in samples], dim=0)
|
|
361
|
+
elif key.endswith('_smiles'):
|
|
362
|
+
batched_data[key] = [s[key] for s in samples]
|
|
363
|
+
elif key == 'target':
|
|
364
|
+
batched_data[key] = torch.stack([s[key].reshape(-1) for s in samples], dim=0)
|
|
365
|
+
elif key == 'original_row_index':
|
|
366
|
+
batched_data[key] = torch.stack([s[key].reshape(()) for s in samples], dim=0)
|
|
367
|
+
return batched_data
|
|
368
|
+
|
|
369
|
+
@staticmethod
|
|
370
|
+
def unwrapper(device='cpu', **batched_data):
|
|
371
|
+
for key in batched_data.keys():
|
|
372
|
+
if key.endswith('_graphs'):
|
|
373
|
+
batched_data[key] = batched_data[key].to(device=device)
|
|
374
|
+
elif key.endswith('_r_node') or key.endswith('_i_node') or key.endswith('_d_node'):
|
|
375
|
+
batched_data[key] = batched_data[key].float().to(device=device)
|
|
376
|
+
elif key.endswith('_r2r_edge') or key.endswith('_i2i_edge') or key.endswith('_d2d_edge'):
|
|
377
|
+
batched_data[key] = batched_data[key].float().to(device=device)
|
|
378
|
+
elif key.endswith('_var'):
|
|
379
|
+
batched_data[key] = batched_data[key].float().to(device=device)
|
|
380
|
+
elif key == 'target':
|
|
381
|
+
batched_data[key] = batched_data[key].float().to(device=device)
|
|
382
|
+
elif key == 'original_row_index':
|
|
383
|
+
batched_data[key] = batched_data[key].long().to(device=device)
|
|
384
|
+
return batched_data
|
|
File without changes
|
|
@@ -0,0 +1,223 @@
|
|
|
1
|
+
from rdkit import Chem
|
|
2
|
+
import torch
|
|
3
|
+
import numpy as np
|
|
4
|
+
import traceback
|
|
5
|
+
import warnings
|
|
6
|
+
from collections import deque
|
|
7
|
+
from torch_geometric.data import HeteroData
|
|
8
|
+
|
|
9
|
+
from .MolGraphGenerator import MolGraphGenerator
|
|
10
|
+
from D4CMPP2.src.utils.featureizer import InvalidAtomError
|
|
11
|
+
from D4CMPP2.src.utils.sculptor import SubgroupSplitter
|
|
12
|
+
|
|
13
|
+
class ISAGraphGenerator(MolGraphGenerator):
|
|
14
|
+
def __init__(self, frag_ref=None, sculptor_index=(6,2,0)):
|
|
15
|
+
self.sculptor = SubgroupSplitter(frag_ref,
|
|
16
|
+
get_index=True,
|
|
17
|
+
split_order=sculptor_index[0],
|
|
18
|
+
combine_rest_order=sculptor_index[1],
|
|
19
|
+
absorb_neighbor_order=sculptor_index[2],
|
|
20
|
+
overlapped_ring_combine=True
|
|
21
|
+
)
|
|
22
|
+
self.r_node_dim = None
|
|
23
|
+
self.i_node_dim = None
|
|
24
|
+
self.d_node_dim = None
|
|
25
|
+
self.r_edge_dim = None
|
|
26
|
+
self.i_edge_dim = None
|
|
27
|
+
self.d_edge_dim = None
|
|
28
|
+
self.node_dim= None
|
|
29
|
+
self.edge_dim = None
|
|
30
|
+
super().__init__()
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
self.set_feature_dim()
|
|
34
|
+
self.node_dim = self.r_node_dim
|
|
35
|
+
self.edge_dim = self.r_edge_dim
|
|
36
|
+
self.verbose = True
|
|
37
|
+
|
|
38
|
+
def set_feature_dim(self):
|
|
39
|
+
try:
|
|
40
|
+
get_graph = self.get_graph('FC1CCCCC1CCCOCCC')
|
|
41
|
+
except InvalidAtomError as e:
|
|
42
|
+
if self.verbose:
|
|
43
|
+
warnings.warn(
|
|
44
|
+
f"ISA feature-dimension probe encountered an invalid atom: {e}",
|
|
45
|
+
RuntimeWarning,
|
|
46
|
+
stacklevel=2,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
def get_graph(self, smi, **kwargs):
|
|
50
|
+
mol = Chem.MolFromSmiles(smi)
|
|
51
|
+
if mol is None:
|
|
52
|
+
raise Exception("Invalid SMILES: failed to generate mol object")
|
|
53
|
+
if kwargs.get("explicit_h",False):
|
|
54
|
+
mol = Chem.AddHs(mol)
|
|
55
|
+
g = self.generate_graph(mol, **kwargs)
|
|
56
|
+
atom_feature = self.af(mol)
|
|
57
|
+
bond_feature = self.bf(mol)
|
|
58
|
+
|
|
59
|
+
g['r_nd'].x = torch.tensor(atom_feature).float()
|
|
60
|
+
if self.r_node_dim is None:
|
|
61
|
+
self.r_node_dim = atom_feature.shape[1]
|
|
62
|
+
|
|
63
|
+
relation = ('r_nd', 'r2r', 'r_nd')
|
|
64
|
+
if mol.GetNumBonds() == 0:
|
|
65
|
+
edge_dim = self.r_edge_dim
|
|
66
|
+
if edge_dim is None:
|
|
67
|
+
edge_dim = self.bf(Chem.MolFromSmiles("C-C")).shape[1]
|
|
68
|
+
edge_count = g[relation].edge_index.shape[1]
|
|
69
|
+
edata = torch.zeros((edge_count, edge_dim), dtype=torch.float32)
|
|
70
|
+
else:
|
|
71
|
+
edata = torch.tensor(bond_feature).float()
|
|
72
|
+
edata = torch.cat([edata,edata],dim=0)
|
|
73
|
+
g[relation].edge_attr = edata
|
|
74
|
+
if self.r_edge_dim is None:
|
|
75
|
+
self.r_edge_dim = edata.shape[1]
|
|
76
|
+
|
|
77
|
+
g['i_nd'].x = torch.zeros((g['i_nd'].num_nodes, 1)).float()
|
|
78
|
+
if self.i_node_dim is None:
|
|
79
|
+
self.i_node_dim = 1
|
|
80
|
+
if self.i_edge_dim is None:
|
|
81
|
+
self.i_edge_dim = 0
|
|
82
|
+
|
|
83
|
+
g['d_nd'].x = torch.zeros((g['d_nd'].num_nodes, 1)).float()
|
|
84
|
+
if self.d_node_dim is None:
|
|
85
|
+
self.d_node_dim = 1
|
|
86
|
+
if self.d_edge_dim is None:
|
|
87
|
+
self.d_edge_dim = 0
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
return g
|
|
91
|
+
|
|
92
|
+
def generate_sub_graph(self, mol, frags):
|
|
93
|
+
src, dst = [], []
|
|
94
|
+
for bond in mol.GetBonds():
|
|
95
|
+
start, end = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
|
|
96
|
+
for frag in frags:
|
|
97
|
+
if start in frag:
|
|
98
|
+
if end in frag:
|
|
99
|
+
src.append(start)
|
|
100
|
+
dst.append(end)
|
|
101
|
+
break
|
|
102
|
+
else:
|
|
103
|
+
break
|
|
104
|
+
elif end in frag:
|
|
105
|
+
if start in frag:
|
|
106
|
+
src.append(start)
|
|
107
|
+
dst.append(end)
|
|
108
|
+
break
|
|
109
|
+
else:
|
|
110
|
+
break
|
|
111
|
+
|
|
112
|
+
return (src+dst, dst+src)
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def generate_dot_graph(self, mol, frags, max_dist = 4):
|
|
116
|
+
src, dst = [], []
|
|
117
|
+
for bond in mol.GetBonds():
|
|
118
|
+
start, end = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
|
|
119
|
+
for i,frag in enumerate(frags):
|
|
120
|
+
if start in frag:
|
|
121
|
+
start_frag = i
|
|
122
|
+
if end in frag:
|
|
123
|
+
end_frag = i
|
|
124
|
+
if start_frag != end_frag:
|
|
125
|
+
src.append(start_frag)
|
|
126
|
+
dst.append(end_frag)
|
|
127
|
+
if len(src) == 0:
|
|
128
|
+
return ([],[]), []
|
|
129
|
+
adjacency = [[] for _ in frags]
|
|
130
|
+
for start, end in zip(src, dst):
|
|
131
|
+
adjacency[start].append(end)
|
|
132
|
+
adjacency[end].append(start)
|
|
133
|
+
|
|
134
|
+
out_src, out_dst, distances = [], [], []
|
|
135
|
+
for root in range(len(frags)):
|
|
136
|
+
shortest = [-1] * len(frags)
|
|
137
|
+
shortest[root] = 0
|
|
138
|
+
queue = deque([root])
|
|
139
|
+
while queue:
|
|
140
|
+
current = queue.popleft()
|
|
141
|
+
if shortest[current] >= max_dist:
|
|
142
|
+
continue
|
|
143
|
+
for neighbor in adjacency[current]:
|
|
144
|
+
if shortest[neighbor] == -1:
|
|
145
|
+
shortest[neighbor] = shortest[current] + 1
|
|
146
|
+
queue.append(neighbor)
|
|
147
|
+
for target in range(root + 1, len(frags)):
|
|
148
|
+
distance = shortest[target]
|
|
149
|
+
if 0 < distance <= max_dist:
|
|
150
|
+
out_src.append(root)
|
|
151
|
+
out_dst.append(target)
|
|
152
|
+
distances.append(distance)
|
|
153
|
+
|
|
154
|
+
directed_distances = distances + distances
|
|
155
|
+
features = np.zeros((len(directed_distances), max_dist), dtype=np.float32)
|
|
156
|
+
for row, distance in enumerate(directed_distances):
|
|
157
|
+
features[row, distance - 1] = 1.0
|
|
158
|
+
return (out_src + out_dst, out_dst + out_src), features
|
|
159
|
+
|
|
160
|
+
def generate_graph(self,mol, **kwargs):
|
|
161
|
+
max_dist = kwargs.get('max_dist',4)
|
|
162
|
+
num_atoms = mol.GetNumAtoms()
|
|
163
|
+
mol_data = self.generate_mol_graph(mol)
|
|
164
|
+
|
|
165
|
+
frag = self.sculptor.fragmentation_with_condition(mol)
|
|
166
|
+
frag_data = self.generate_sub_graph(mol, frag)
|
|
167
|
+
|
|
168
|
+
i2d_src = []
|
|
169
|
+
i2d_dst = []
|
|
170
|
+
for i, f in enumerate(frag):
|
|
171
|
+
for a in f:
|
|
172
|
+
i2d_src.append(a)
|
|
173
|
+
i2d_dst.append(i)
|
|
174
|
+
|
|
175
|
+
if len(frag) ==1:
|
|
176
|
+
dot_data = ([0],[0])
|
|
177
|
+
dist = np.zeros((1,max_dist))
|
|
178
|
+
else:
|
|
179
|
+
dot_data, dist = self.generate_dot_graph(mol, frag, max_dist = max_dist)
|
|
180
|
+
|
|
181
|
+
graph_data = {
|
|
182
|
+
('r_nd', 'r2r', 'r_nd'): mol_data,
|
|
183
|
+
('r_nd', 'r2i', 'i_nd'): (list(range(num_atoms)), list(range(num_atoms))),
|
|
184
|
+
('i_nd', 'i2i', 'i_nd'): frag_data,
|
|
185
|
+
('i_nd', 'i2d', 'd_nd'): (i2d_src, i2d_dst),
|
|
186
|
+
('d_nd', 'd2d', 'd_nd'): dot_data,
|
|
187
|
+
('d_nd', 'd2r', 'r_nd'): (i2d_dst, i2d_src),
|
|
188
|
+
|
|
189
|
+
}
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
g = HeteroData()
|
|
193
|
+
g['r_nd'].num_nodes = num_atoms
|
|
194
|
+
g['i_nd'].num_nodes = num_atoms
|
|
195
|
+
g['d_nd'].num_nodes = len(frag)
|
|
196
|
+
for edge_type, (src, dst) in graph_data.items():
|
|
197
|
+
g[edge_type].edge_index = torch.tensor([src, dst], dtype=torch.long)
|
|
198
|
+
g['r_nd', 'r2r', 'r_nd'].edge_attr = torch.empty((len(mol_data[0]), 0))
|
|
199
|
+
g['i_nd', 'i2i', 'i_nd'].edge_attr = torch.empty((len(frag_data[0]), 0))
|
|
200
|
+
g['d_nd', 'd2d', 'd_nd'].edge_attr = torch.tensor(dist).float()
|
|
201
|
+
if self.d_edge_dim is None:
|
|
202
|
+
self.d_edge_dim = dist.shape[1]
|
|
203
|
+
|
|
204
|
+
return g
|
|
205
|
+
|
|
206
|
+
def get_empty_graph(self):
|
|
207
|
+
graph = HeteroData()
|
|
208
|
+
for node_type, dim in (("r_nd", self.r_node_dim), ("i_nd", 1), ("d_nd", 1)):
|
|
209
|
+
graph[node_type].x = torch.zeros((0, dim or 1), dtype=torch.float32)
|
|
210
|
+
relations = (
|
|
211
|
+
("r_nd", "r2r", "r_nd"),
|
|
212
|
+
("r_nd", "r2i", "i_nd"),
|
|
213
|
+
("i_nd", "i2i", "i_nd"),
|
|
214
|
+
("i_nd", "i2d", "d_nd"),
|
|
215
|
+
("d_nd", "d2d", "d_nd"),
|
|
216
|
+
("d_nd", "d2r", "r_nd"),
|
|
217
|
+
)
|
|
218
|
+
for relation in relations:
|
|
219
|
+
graph[relation].edge_index = torch.empty((2, 0), dtype=torch.long)
|
|
220
|
+
graph["r_nd", "r2r", "r_nd"].edge_attr = torch.zeros((0, self.r_edge_dim or 1))
|
|
221
|
+
graph["i_nd", "i2i", "i_nd"].edge_attr = torch.zeros((0, 0))
|
|
222
|
+
graph["d_nd", "d2d", "d_nd"].edge_attr = torch.zeros((0, self.d_edge_dim or 1))
|
|
223
|
+
return graph
|
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import rdkit.Chem as Chem
|
|
3
|
+
import numpy as np
|
|
4
|
+
import traceback
|
|
5
|
+
from torch_geometric.data import Data
|
|
6
|
+
from D4CMPP2.src.utils.featureizer import get_atom_features, get_bond_features, InvalidAtomError
|
|
7
|
+
|
|
8
|
+
class MolGraphGenerator:
|
|
9
|
+
def __init__(self):
|
|
10
|
+
self.af = get_atom_features
|
|
11
|
+
self.bf = get_bond_features
|
|
12
|
+
self.set_feature_dim()
|
|
13
|
+
|
|
14
|
+
# Set the feature dimensions by generating a dummy graph
|
|
15
|
+
def set_feature_dim(self):
|
|
16
|
+
self.node_dim = self.af(Chem.MolFromSmiles('C')).shape[1]
|
|
17
|
+
self.edge_dim = self.bf(Chem.MolFromSmiles('C-C')).shape[1]
|
|
18
|
+
|
|
19
|
+
# Get the graph from the SMILES
|
|
20
|
+
def get_graph(self,smi,**kwargs):
|
|
21
|
+
mol = Chem.MolFromSmiles(smi)
|
|
22
|
+
|
|
23
|
+
if kwargs.get("explicit_h",False) or mol.GetNumAtoms() == 1:
|
|
24
|
+
mol = Chem.AddHs(mol)
|
|
25
|
+
if mol is None:
|
|
26
|
+
raise Exception("Invalid SMILES: failed to generate mol object")
|
|
27
|
+
g = self.generate_graph(mol)
|
|
28
|
+
g = self.add_feature(g,mol)
|
|
29
|
+
|
|
30
|
+
return g
|
|
31
|
+
|
|
32
|
+
def get_empty_graph(self):
|
|
33
|
+
return Data(
|
|
34
|
+
x=torch.zeros((0, self.node_dim), dtype=torch.float32),
|
|
35
|
+
edge_index=torch.empty((2, 0), dtype=torch.long),
|
|
36
|
+
edge_attr=torch.zeros((0, self.edge_dim), dtype=torch.float32),
|
|
37
|
+
num_nodes=0,
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
# Add the features to the graph
|
|
41
|
+
def add_feature(self, g, mol):
|
|
42
|
+
atom_feature = self.af(mol)
|
|
43
|
+
g.x = torch.tensor(atom_feature).float()
|
|
44
|
+
|
|
45
|
+
if mol.GetNumBonds() == 0:
|
|
46
|
+
g.edge_attr = torch.zeros((0,self.edge_dim)).float()
|
|
47
|
+
else:
|
|
48
|
+
bond_feature = self.bf(mol)
|
|
49
|
+
edata = torch.tensor(bond_feature).float()
|
|
50
|
+
edata = torch.cat([edata,edata],dim=0)
|
|
51
|
+
g.edge_attr = edata
|
|
52
|
+
return g
|
|
53
|
+
|
|
54
|
+
# Generate the graph from the molecule object
|
|
55
|
+
def generate_mol_graph(self,mol):
|
|
56
|
+
num_atoms = mol.GetNumAtoms()
|
|
57
|
+
if num_atoms == 1:
|
|
58
|
+
return ([0],[0])
|
|
59
|
+
src, dst = [], []
|
|
60
|
+
for bond in mol.GetBonds():
|
|
61
|
+
start, end = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
|
|
62
|
+
src.append(start)
|
|
63
|
+
dst.append(end)
|
|
64
|
+
|
|
65
|
+
return (src+dst, dst+src)
|
|
66
|
+
|
|
67
|
+
def generate_graph(self,mol):
|
|
68
|
+
mol_data = self.generate_mol_graph(mol)
|
|
69
|
+
g = Data(
|
|
70
|
+
edge_index=torch.tensor(mol_data, dtype=torch.long),
|
|
71
|
+
num_nodes=mol.GetNumAtoms(),
|
|
72
|
+
)
|
|
73
|
+
return g
|