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.
Files changed (103) hide show
  1. D4CMPP2/_Data/AGENTS.md +24 -0
  2. D4CMPP2/_Data/Aqsoldb.csv +9291 -0
  3. D4CMPP2/_Data/BradleyMP.csv +3042 -0
  4. D4CMPP2/_Data/Lipophilicity.csv +1131 -0
  5. D4CMPP2/_Data/README.md +8 -0
  6. D4CMPP2/_Data/__init__.py +26 -0
  7. D4CMPP2/_Data/optical.csv +20237 -0
  8. D4CMPP2/_Data/test.csv +190 -0
  9. D4CMPP2/__init__.py +16 -0
  10. D4CMPP2/__main__.py +5 -0
  11. D4CMPP2/_main.py +500 -0
  12. D4CMPP2/cli.py +7 -0
  13. D4CMPP2/exceptions.py +53 -0
  14. D4CMPP2/grid_search.py +259 -0
  15. D4CMPP2/network_refer.yaml +160 -0
  16. D4CMPP2/networks/AFP_model.py +72 -0
  17. D4CMPP2/networks/AFPwithSolv_model.py +72 -0
  18. D4CMPP2/networks/DMPNN_model.py +90 -0
  19. D4CMPP2/networks/DMPNNwithSolv_model.py +89 -0
  20. D4CMPP2/networks/GAT_model.py +48 -0
  21. D4CMPP2/networks/GATwithSolv_model.py +63 -0
  22. D4CMPP2/networks/GCN_model.py +113 -0
  23. D4CMPP2/networks/GCNwithSolv_model.py +103 -0
  24. D4CMPP2/networks/GC_model.py +122 -0
  25. D4CMPP2/networks/ISATPM_model.py +14 -0
  26. D4CMPP2/networks/ISATPN_model.py +199 -0
  27. D4CMPP2/networks/ISAT_model.py +90 -0
  28. D4CMPP2/networks/MPNN_model.py +56 -0
  29. D4CMPP2/networks/MPNNwithSolv_model.py +72 -0
  30. D4CMPP2/networks/__init__.py +25 -0
  31. D4CMPP2/networks/base.py +250 -0
  32. D4CMPP2/networks/registry.py +187 -0
  33. D4CMPP2/networks/src/AFP.py +118 -0
  34. D4CMPP2/networks/src/BiDropout.py +29 -0
  35. D4CMPP2/networks/src/DMPNN.py +35 -0
  36. D4CMPP2/networks/src/GAT.py +69 -0
  37. D4CMPP2/networks/src/GC.py +85 -0
  38. D4CMPP2/networks/src/GCN.py +71 -0
  39. D4CMPP2/networks/src/ISAT.py +153 -0
  40. D4CMPP2/networks/src/Linear.py +49 -0
  41. D4CMPP2/networks/src/MPNN.py +56 -0
  42. D4CMPP2/networks/src/SolventLayer.py +62 -0
  43. D4CMPP2/networks/src/__init__.py +0 -0
  44. D4CMPP2/networks/src/distGCN.py +21 -0
  45. D4CMPP2/networks/src/pyg_hetero.py +24 -0
  46. D4CMPP2/optimize.py +472 -0
  47. D4CMPP2/src/Analyzer/ISAAnalyzer.py +458 -0
  48. D4CMPP2/src/Analyzer/ISAPNAnalyzer.py +366 -0
  49. D4CMPP2/src/Analyzer/ISAwSAnalyzer.py +117 -0
  50. D4CMPP2/src/Analyzer/MolAnalyzer.py +319 -0
  51. D4CMPP2/src/Analyzer/__init__.py +54 -0
  52. D4CMPP2/src/Analyzer/core.py +480 -0
  53. D4CMPP2/src/Analyzer/factory.py +166 -0
  54. D4CMPP2/src/Analyzer/interpretation.py +232 -0
  55. D4CMPP2/src/Analyzer/results.py +101 -0
  56. D4CMPP2/src/DataManager/Dataset/GraphDataset.py +314 -0
  57. D4CMPP2/src/DataManager/Dataset/ISAGraphDataset.py +384 -0
  58. D4CMPP2/src/DataManager/Dataset/__init__.py +0 -0
  59. D4CMPP2/src/DataManager/GraphGenerator/ISAGraphGenerator.py +223 -0
  60. D4CMPP2/src/DataManager/GraphGenerator/MolGraphGenerator.py +73 -0
  61. D4CMPP2/src/DataManager/GraphGenerator/__init__.py +14 -0
  62. D4CMPP2/src/DataManager/ISADataManager.py +67 -0
  63. D4CMPP2/src/DataManager/MolDataManager.py +735 -0
  64. D4CMPP2/src/DataManager/__init__.py +14 -0
  65. D4CMPP2/src/DataManager/contracts.py +179 -0
  66. D4CMPP2/src/NetworkManager/ISANetworkManager.py +12 -0
  67. D4CMPP2/src/NetworkManager/NetworkManager.py +520 -0
  68. D4CMPP2/src/NetworkManager/__init__.py +14 -0
  69. D4CMPP2/src/PostProcessor.py +160 -0
  70. D4CMPP2/src/TrainManager/ISATrainManager.py +26 -0
  71. D4CMPP2/src/TrainManager/TrainManager.py +254 -0
  72. D4CMPP2/src/TrainManager/__init__.py +14 -0
  73. D4CMPP2/src/TrainManager/callbacks.py +119 -0
  74. D4CMPP2/src/__init__.py +0 -0
  75. D4CMPP2/src/utils/PATH.py +246 -0
  76. D4CMPP2/src/utils/__init__.py +0 -0
  77. D4CMPP2/src/utils/argparser.py +56 -0
  78. D4CMPP2/src/utils/checkpointing.py +90 -0
  79. D4CMPP2/src/utils/config_resolution.py +123 -0
  80. D4CMPP2/src/utils/config_validation.py +370 -0
  81. D4CMPP2/src/utils/csv_validation.py +105 -0
  82. D4CMPP2/src/utils/data_quality.py +181 -0
  83. D4CMPP2/src/utils/featureizer.py +202 -0
  84. D4CMPP2/src/utils/functional_group.csv +169 -0
  85. D4CMPP2/src/utils/graph_cache.py +213 -0
  86. D4CMPP2/src/utils/leaderboard.py +212 -0
  87. D4CMPP2/src/utils/metrics.py +31 -0
  88. D4CMPP2/src/utils/module_loader.py +147 -0
  89. D4CMPP2/src/utils/output.py +80 -0
  90. D4CMPP2/src/utils/reproducibility.py +70 -0
  91. D4CMPP2/src/utils/run_manifest.py +175 -0
  92. D4CMPP2/src/utils/scaler.py +40 -0
  93. D4CMPP2/src/utils/sculptor.py +713 -0
  94. D4CMPP2/src/utils/splitting.py +250 -0
  95. D4CMPP2/src/utils/supportfile_saver.py +94 -0
  96. D4CMPP2/src/utils/tools.py +156 -0
  97. D4CMPP2/src/utils/transfer_learning.py +111 -0
  98. d4cmpp2-0.4.0.dist-info/METADATA +420 -0
  99. d4cmpp2-0.4.0.dist-info/RECORD +103 -0
  100. d4cmpp2-0.4.0.dist-info/WHEEL +5 -0
  101. d4cmpp2-0.4.0.dist-info/entry_points.txt +2 -0
  102. d4cmpp2-0.4.0.dist-info/licenses/LICENSE +21 -0
  103. 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