k3-node 1.0.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 (459) hide show
  1. k3_node/__init__.py +122 -0
  2. k3_node/applications/__init__.py +17 -0
  3. k3_node/applications/bio/__init__.py +21 -0
  4. k3_node/applications/chemistry/__init__.py +155 -0
  5. k3_node/applications/materials/__init__.py +127 -0
  6. k3_node/applications/materials/basis.py +449 -0
  7. k3_node/applications/materials/chgnet.py +360 -0
  8. k3_node/applications/materials/core.py +351 -0
  9. k3_node/applications/materials/grace.py +246 -0
  10. k3_node/applications/materials/io.py +230 -0
  11. k3_node/applications/materials/m3gnet.py +462 -0
  12. k3_node/applications/materials/megnet.py +395 -0
  13. k3_node/applications/materials/qet.py +220 -0
  14. k3_node/applications/materials/readout.py +235 -0
  15. k3_node/applications/materials/so3net.py +234 -0
  16. k3_node/applications/materials/tensornet.py +381 -0
  17. k3_node/applications/materials/test_materials.py +167 -0
  18. k3_node/applications/materials/wrappers.py +95 -0
  19. k3_node/data/__init__.py +47 -0
  20. k3_node/data/batch.py +102 -0
  21. k3_node/data/collate.py +282 -0
  22. k3_node/data/data.py +532 -0
  23. k3_node/data/database.py +154 -0
  24. k3_node/data/dataset.py +182 -0
  25. k3_node/data/download.py +49 -0
  26. k3_node/data/extract.py +45 -0
  27. k3_node/data/feature_store.py +70 -0
  28. k3_node/data/graph_store.py +92 -0
  29. k3_node/data/hetero_data.py +374 -0
  30. k3_node/data/hypergraph_data.py +59 -0
  31. k3_node/data/in_memory_dataset.py +177 -0
  32. k3_node/data/makedirs.py +7 -0
  33. k3_node/data/on_disk_dataset.py +77 -0
  34. k3_node/data/separate.py +115 -0
  35. k3_node/data/storage.py +593 -0
  36. k3_node/data/temporal.py +154 -0
  37. k3_node/data/test_batch.py +67 -0
  38. k3_node/data/test_data.py +68 -0
  39. k3_node/data/test_dataset_and_stores.py +111 -0
  40. k3_node/data/test_hetero_data.py +33 -0
  41. k3_node/data/test_temporal_and_hyper.py +32 -0
  42. k3_node/data/view.py +43 -0
  43. k3_node/datasets/__init__.py +88 -0
  44. k3_node/datasets/actor.py +101 -0
  45. k3_node/datasets/airports.py +84 -0
  46. k3_node/datasets/amazon.py +66 -0
  47. k3_node/datasets/ba2motif_dataset.py +73 -0
  48. k3_node/datasets/ba_shapes.py +81 -0
  49. k3_node/datasets/bitcoin_otc.py +77 -0
  50. k3_node/datasets/citation_full.py +81 -0
  51. k3_node/datasets/coauthor.py +66 -0
  52. k3_node/datasets/dblp.py +106 -0
  53. k3_node/datasets/digits.py +63 -0
  54. k3_node/datasets/email_eu_core.py +60 -0
  55. k3_node/datasets/entities.py +158 -0
  56. k3_node/datasets/explainer_dataset.py +101 -0
  57. k3_node/datasets/facebook.py +51 -0
  58. k3_node/datasets/fake.py +256 -0
  59. k3_node/datasets/freebase.py +90 -0
  60. k3_node/datasets/geometric_shapes.py +69 -0
  61. k3_node/datasets/github.py +51 -0
  62. k3_node/datasets/graph_generator/__init__.py +6 -0
  63. k3_node/datasets/graph_generator/ba_graph.py +20 -0
  64. k3_node/datasets/graph_generator/base.py +29 -0
  65. k3_node/datasets/graph_generator/er_graph.py +21 -0
  66. k3_node/datasets/icews.py +58 -0
  67. k3_node/datasets/imdb.py +96 -0
  68. k3_node/datasets/jodie.py +56 -0
  69. k3_node/datasets/karate.py +56 -0
  70. k3_node/datasets/lastfm_asia.py +51 -0
  71. k3_node/datasets/mesh_correspondence.py +50 -0
  72. k3_node/datasets/molecule_net.py +148 -0
  73. k3_node/datasets/motif_generator/__init__.py +7 -0
  74. k3_node/datasets/motif_generator/base.py +29 -0
  75. k3_node/datasets/motif_generator/custom.py +17 -0
  76. k3_node/datasets/motif_generator/cycle.py +25 -0
  77. k3_node/datasets/motif_generator/house.py +27 -0
  78. k3_node/datasets/movielens.py +55 -0
  79. k3_node/datasets/planetoid.py +137 -0
  80. k3_node/datasets/polblogs.py +63 -0
  81. k3_node/datasets/ppi.py +189 -0
  82. k3_node/datasets/qm7.py +65 -0
  83. k3_node/datasets/qm9.py +132 -0
  84. k3_node/datasets/reddit.py +121 -0
  85. k3_node/datasets/sbm_dataset.py +165 -0
  86. k3_node/datasets/seal.py +74 -0
  87. k3_node/datasets/shape_scenes.py +92 -0
  88. k3_node/datasets/test_datasets.py +322 -0
  89. k3_node/datasets/tu_dataset.py +131 -0
  90. k3_node/datasets/twitch.py +66 -0
  91. k3_node/datasets/webkb.py +102 -0
  92. k3_node/datasets/wikics.py +85 -0
  93. k3_node/datasets/word_net.py +184 -0
  94. k3_node/etl/__init__.py +37 -0
  95. k3_node/etl/encoders.py +248 -0
  96. k3_node/etl/graph_builders.py +270 -0
  97. k3_node/etl/relational_to_graph.py +201 -0
  98. k3_node/etl/table_to_graph.py +244 -0
  99. k3_node/etl/test_etl.py +318 -0
  100. k3_node/export/__init__.py +15 -0
  101. k3_node/export/cross_backend.py +172 -0
  102. k3_node/export/onnx_exporter.py +190 -0
  103. k3_node/export/runtime.py +254 -0
  104. k3_node/export/tensorrt_exporter.py +201 -0
  105. k3_node/export/test_export.py +337 -0
  106. k3_node/export/tflite_exporter.py +112 -0
  107. k3_node/hub/__init__.py +29 -0
  108. k3_node/hub/dataset_hub.py +242 -0
  109. k3_node/hub/hub_mixin.py +599 -0
  110. k3_node/hub/model_card.py +133 -0
  111. k3_node/hub/test_hub.py +419 -0
  112. k3_node/io/__init__.py +22 -0
  113. k3_node/io/fs.py +117 -0
  114. k3_node/io/npz.py +45 -0
  115. k3_node/io/off.py +29 -0
  116. k3_node/io/planetoid.py +98 -0
  117. k3_node/io/tu.py +137 -0
  118. k3_node/io/txt_array.py +58 -0
  119. k3_node/layers/__init__.py +14 -0
  120. k3_node/layers/aggr/__init__.py +70 -0
  121. k3_node/layers/aggr/attention.py +77 -0
  122. k3_node/layers/aggr/base.py +403 -0
  123. k3_node/layers/aggr/basic.py +412 -0
  124. k3_node/layers/aggr/deep_sets.py +65 -0
  125. k3_node/layers/aggr/deepsets.py +29 -0
  126. k3_node/layers/aggr/equilibrium.py +107 -0
  127. k3_node/layers/aggr/fused.py +43 -0
  128. k3_node/layers/aggr/gmt.py +89 -0
  129. k3_node/layers/aggr/gru.py +58 -0
  130. k3_node/layers/aggr/lcm.py +143 -0
  131. k3_node/layers/aggr/lstm.py +58 -0
  132. k3_node/layers/aggr/mlp.py +75 -0
  133. k3_node/layers/aggr/multi.py +154 -0
  134. k3_node/layers/aggr/patch_transformer.py +137 -0
  135. k3_node/layers/aggr/quantile.py +125 -0
  136. k3_node/layers/aggr/resolver.py +68 -0
  137. k3_node/layers/aggr/scaler.py +133 -0
  138. k3_node/layers/aggr/set2set.py +87 -0
  139. k3_node/layers/aggr/set_transformer.py +107 -0
  140. k3_node/layers/aggr/sort.py +68 -0
  141. k3_node/layers/aggr/test_aggr.py +337 -0
  142. k3_node/layers/aggr/utils.py +210 -0
  143. k3_node/layers/aggr/variance_preserving.py +54 -0
  144. k3_node/layers/attention/__init__.py +5 -0
  145. k3_node/layers/attention/pair_attention.py +448 -0
  146. k3_node/layers/attention/performer.py +187 -0
  147. k3_node/layers/attention/polynormer.py +160 -0
  148. k3_node/layers/attention/qformer.py +143 -0
  149. k3_node/layers/attention/sgformer.py +106 -0
  150. k3_node/layers/attention/test_attention.py +68 -0
  151. k3_node/layers/attention/test_pair_attention.py +91 -0
  152. k3_node/layers/conv/__init__.py +149 -0
  153. k3_node/layers/conv/agnn_conv.py +120 -0
  154. k3_node/layers/conv/antisymmetric_conv.py +94 -0
  155. k3_node/layers/conv/appnp.py +105 -0
  156. k3_node/layers/conv/appnp_conv.py +157 -0
  157. k3_node/layers/conv/arma_conv.py +231 -0
  158. k3_node/layers/conv/cg_conv.py +92 -0
  159. k3_node/layers/conv/cheb_conv.py +137 -0
  160. k3_node/layers/conv/cluster_gcn_conv.py +102 -0
  161. k3_node/layers/conv/conv.py +100 -0
  162. k3_node/layers/conv/crystal_conv.py +140 -0
  163. k3_node/layers/conv/cugraph.py +84 -0
  164. k3_node/layers/conv/diffusion_conv.py +144 -0
  165. k3_node/layers/conv/dir_gnn_conv.py +93 -0
  166. k3_node/layers/conv/dna_conv.py +192 -0
  167. k3_node/layers/conv/edge_conv.py +107 -0
  168. k3_node/layers/conv/eg_conv.py +155 -0
  169. k3_node/layers/conv/fa_conv.py +107 -0
  170. k3_node/layers/conv/feast_conv.py +126 -0
  171. k3_node/layers/conv/film_conv.py +143 -0
  172. k3_node/layers/conv/gat_conv.py +244 -0
  173. k3_node/layers/conv/gated_graph_conv.py +136 -0
  174. k3_node/layers/conv/gatv2_conv.py +205 -0
  175. k3_node/layers/conv/gcn.py +144 -0
  176. k3_node/layers/conv/gcn2_conv.py +126 -0
  177. k3_node/layers/conv/gcn_conv.py +135 -0
  178. k3_node/layers/conv/gen_conv.py +163 -0
  179. k3_node/layers/conv/general_conv.py +218 -0
  180. k3_node/layers/conv/gin_conv.py +218 -0
  181. k3_node/layers/conv/gmm_conv.py +172 -0
  182. k3_node/layers/conv/gps_conv.py +153 -0
  183. k3_node/layers/conv/graph_attention.py +262 -0
  184. k3_node/layers/conv/graph_conv.py +84 -0
  185. k3_node/layers/conv/gravnet_conv.py +93 -0
  186. k3_node/layers/conv/han_conv.py +175 -0
  187. k3_node/layers/conv/heat_conv.py +131 -0
  188. k3_node/layers/conv/hetero_conv.py +128 -0
  189. k3_node/layers/conv/hgt_conv.py +218 -0
  190. k3_node/layers/conv/hypergraph_conv.py +182 -0
  191. k3_node/layers/conv/le_conv.py +81 -0
  192. k3_node/layers/conv/lg_conv.py +58 -0
  193. k3_node/layers/conv/meshcnn_conv.py +84 -0
  194. k3_node/layers/conv/message_passing.py +451 -0
  195. k3_node/layers/conv/mf_conv.py +95 -0
  196. k3_node/layers/conv/mixhop_conv.py +108 -0
  197. k3_node/layers/conv/nn_conv.py +110 -0
  198. k3_node/layers/conv/pan_conv.py +100 -0
  199. k3_node/layers/conv/pdn_conv.py +109 -0
  200. k3_node/layers/conv/pna_conv.py +177 -0
  201. k3_node/layers/conv/point_conv.py +101 -0
  202. k3_node/layers/conv/point_gnn_conv.py +90 -0
  203. k3_node/layers/conv/point_transformer_conv.py +132 -0
  204. k3_node/layers/conv/ppf_conv.py +135 -0
  205. k3_node/layers/conv/ppnp.py +89 -0
  206. k3_node/layers/conv/res_gated_graph_conv.py +126 -0
  207. k3_node/layers/conv/rgat_conv.py +251 -0
  208. k3_node/layers/conv/rgcn_conv.py +321 -0
  209. k3_node/layers/conv/sage_conv.py +154 -0
  210. k3_node/layers/conv/sg_conv.py +96 -0
  211. k3_node/layers/conv/signed_conv.py +100 -0
  212. k3_node/layers/conv/simple_conv.py +75 -0
  213. k3_node/layers/conv/spline_conv.py +182 -0
  214. k3_node/layers/conv/ssg_conv.py +101 -0
  215. k3_node/layers/conv/supergat_conv.py +195 -0
  216. k3_node/layers/conv/tag_conv.py +98 -0
  217. k3_node/layers/conv/test_backend_consistency.py +164 -0
  218. k3_node/layers/conv/test_conv.py +176 -0
  219. k3_node/layers/conv/test_conv_pyg.py +566 -0
  220. k3_node/layers/conv/transformer_conv.py +168 -0
  221. k3_node/layers/conv/utils.py +403 -0
  222. k3_node/layers/conv/wl_conv.py +151 -0
  223. k3_node/layers/conv/x_conv.py +187 -0
  224. k3_node/layers/dense/__init__.py +40 -0
  225. k3_node/layers/dense/dense_gat_conv.py +149 -0
  226. k3_node/layers/dense/dense_gcn_conv.py +117 -0
  227. k3_node/layers/dense/dense_gin_conv.py +88 -0
  228. k3_node/layers/dense/dense_graph_conv.py +95 -0
  229. k3_node/layers/dense/dense_sage_conv.py +85 -0
  230. k3_node/layers/dense/diff_pool.py +76 -0
  231. k3_node/layers/dense/dmon_pool.py +223 -0
  232. k3_node/layers/dense/linear.py +327 -0
  233. k3_node/layers/dense/mincut_pool.py +92 -0
  234. k3_node/layers/dense/test_dense.py +377 -0
  235. k3_node/layers/functional/__init__.py +13 -0
  236. k3_node/layers/functional/bro.py +49 -0
  237. k3_node/layers/functional/edge_dropout.py +55 -0
  238. k3_node/layers/functional/gini.py +44 -0
  239. k3_node/layers/functional/test_functional.py +34 -0
  240. k3_node/layers/kge/__init__.py +17 -0
  241. k3_node/layers/kge/base.py +255 -0
  242. k3_node/layers/kge/complex.py +98 -0
  243. k3_node/layers/kge/distmult.py +79 -0
  244. k3_node/layers/kge/loader.py +50 -0
  245. k3_node/layers/kge/rotate.py +103 -0
  246. k3_node/layers/kge/test_kge.py +76 -0
  247. k3_node/layers/kge/transe.py +96 -0
  248. k3_node/layers/norm/__init__.py +23 -0
  249. k3_node/layers/norm/batch_norm.py +328 -0
  250. k3_node/layers/norm/diff_group_norm.py +141 -0
  251. k3_node/layers/norm/graph_norm.py +105 -0
  252. k3_node/layers/norm/graph_size_norm.py +57 -0
  253. k3_node/layers/norm/instance_norm.py +163 -0
  254. k3_node/layers/norm/layer_norm.py +245 -0
  255. k3_node/layers/norm/mean_subtraction_norm.py +57 -0
  256. k3_node/layers/norm/msg_norm.py +58 -0
  257. k3_node/layers/norm/pair_norm.py +94 -0
  258. k3_node/layers/norm/test_norm.py +275 -0
  259. k3_node/layers/pool/__init__.py +83 -0
  260. k3_node/layers/pool/approx_knn.py +101 -0
  261. k3_node/layers/pool/asap.py +173 -0
  262. k3_node/layers/pool/avg_pool.py +165 -0
  263. k3_node/layers/pool/cluster_pool.py +168 -0
  264. k3_node/layers/pool/connect/__init__.py +10 -0
  265. k3_node/layers/pool/connect/base.py +103 -0
  266. k3_node/layers/pool/connect/filter_edges.py +113 -0
  267. k3_node/layers/pool/consecutive.py +30 -0
  268. k3_node/layers/pool/decimation.py +48 -0
  269. k3_node/layers/pool/edge_pool.py +189 -0
  270. k3_node/layers/pool/glob.py +139 -0
  271. k3_node/layers/pool/graclus.py +66 -0
  272. k3_node/layers/pool/knn.py +253 -0
  273. k3_node/layers/pool/max_pool.py +159 -0
  274. k3_node/layers/pool/mem_pool.py +145 -0
  275. k3_node/layers/pool/pan_pool.py +144 -0
  276. k3_node/layers/pool/point_cloud.py +212 -0
  277. k3_node/layers/pool/pool.py +119 -0
  278. k3_node/layers/pool/sag_pool.py +174 -0
  279. k3_node/layers/pool/select/__init__.py +10 -0
  280. k3_node/layers/pool/select/base.py +112 -0
  281. k3_node/layers/pool/select/topk.py +206 -0
  282. k3_node/layers/pool/test_pool.py +456 -0
  283. k3_node/layers/pool/topk_pool.py +103 -0
  284. k3_node/layers/pool/voxel_grid.py +70 -0
  285. k3_node/layers/unpool/__init__.py +9 -0
  286. k3_node/layers/unpool/knn_interpolate.py +57 -0
  287. k3_node/layers/unpool/test_unpool.py +31 -0
  288. k3_node/loader/__init__.py +62 -0
  289. k3_node/loader/base.py +69 -0
  290. k3_node/loader/cache.py +68 -0
  291. k3_node/loader/cluster.py +127 -0
  292. k3_node/loader/data_list_loader.py +45 -0
  293. k3_node/loader/dataloader.py +117 -0
  294. k3_node/loader/dense_data_loader.py +62 -0
  295. k3_node/loader/dynamic_batch_sampler.py +93 -0
  296. k3_node/loader/graph_saint.py +188 -0
  297. k3_node/loader/hgt_loader.py +90 -0
  298. k3_node/loader/imbalanced_sampler.py +87 -0
  299. k3_node/loader/keras_dataset.py +334 -0
  300. k3_node/loader/link_loader.py +179 -0
  301. k3_node/loader/link_neighbor_loader.py +202 -0
  302. k3_node/loader/mixin.py +190 -0
  303. k3_node/loader/neighbor_loader.py +159 -0
  304. k3_node/loader/neighbor_sampler.py +167 -0
  305. k3_node/loader/node_loader.py +185 -0
  306. k3_node/loader/prefetch.py +115 -0
  307. k3_node/loader/random_node_loader.py +89 -0
  308. k3_node/loader/sampler_utils.py +499 -0
  309. k3_node/loader/shadow.py +115 -0
  310. k3_node/loader/temporal_dataloader.py +98 -0
  311. k3_node/loader/test_dataloader.py +113 -0
  312. k3_node/loader/test_keras_dataset.py +221 -0
  313. k3_node/loader/test_neighbor_loader.py +122 -0
  314. k3_node/loader/test_sampler_utils.py +82 -0
  315. k3_node/loader/test_samplers.py +96 -0
  316. k3_node/loader/test_subgraph_loaders.py +89 -0
  317. k3_node/loader/utils.py +232 -0
  318. k3_node/loader/zip_loader.py +88 -0
  319. k3_node/metrics.py +94 -0
  320. k3_node/models/__init__.py +424 -0
  321. k3_node/models/attentive_fp.py +232 -0
  322. k3_node/models/attract_repel.py +108 -0
  323. k3_node/models/autoencoder.py +318 -0
  324. k3_node/models/basic_gnn.py +443 -0
  325. k3_node/models/bio/__init__.py +4 -0
  326. k3_node/models/captum.py +52 -0
  327. k3_node/models/chemistry/__init__.py +4 -0
  328. k3_node/models/correct_and_smooth.py +146 -0
  329. k3_node/models/deep_graph_infomax.py +113 -0
  330. k3_node/models/deepgcn.py +121 -0
  331. k3_node/models/dimenet.py +737 -0
  332. k3_node/models/dimenet_utils.py +153 -0
  333. k3_node/models/gnnff.py +263 -0
  334. k3_node/models/gps_model.py +1122 -0
  335. k3_node/models/gpse.py +638 -0
  336. k3_node/models/graph_unet.py +199 -0
  337. k3_node/models/graphmae2.py +954 -0
  338. k3_node/models/graphormer.py +1258 -0
  339. k3_node/models/graphormer_3d.py +868 -0
  340. k3_node/models/grover.py +1066 -0
  341. k3_node/models/jumping_knowledge.py +200 -0
  342. k3_node/models/label_prop.py +110 -0
  343. k3_node/models/lightgcn.py +171 -0
  344. k3_node/models/linkx.py +181 -0
  345. k3_node/models/lpformer.py +404 -0
  346. k3_node/models/mask_label.py +114 -0
  347. k3_node/models/materials/__init__.py +33 -0
  348. k3_node/models/meta.py +133 -0
  349. k3_node/models/metapath2vec.py +234 -0
  350. k3_node/models/mlp.py +264 -0
  351. k3_node/models/mole_bert.py +379 -0
  352. k3_node/models/neural_fingerprint.py +95 -0
  353. k3_node/models/node2vec.py +213 -0
  354. k3_node/models/pmlp.py +157 -0
  355. k3_node/models/polynormer.py +229 -0
  356. k3_node/models/rect.py +93 -0
  357. k3_node/models/renet.py +221 -0
  358. k3_node/models/rev_gnn.py +128 -0
  359. k3_node/models/schnet.py +484 -0
  360. k3_node/models/sgformer.py +195 -0
  361. k3_node/models/signed_gcn.py +185 -0
  362. k3_node/models/test_attentive_fp.py +32 -0
  363. k3_node/models/test_attract_repel.py +33 -0
  364. k3_node/models/test_autoencoder.py +119 -0
  365. k3_node/models/test_basic_gnn.py +102 -0
  366. k3_node/models/test_correct_and_smooth.py +40 -0
  367. k3_node/models/test_deep_graph_infomax.py +68 -0
  368. k3_node/models/test_deepgcn.py +21 -0
  369. k3_node/models/test_dimenet.py +86 -0
  370. k3_node/models/test_domain_apis.py +138 -0
  371. k3_node/models/test_gnnff.py +24 -0
  372. k3_node/models/test_gps_model.py +271 -0
  373. k3_node/models/test_gpse.py +34 -0
  374. k3_node/models/test_graph_unet.py +26 -0
  375. k3_node/models/test_graphmae2.py +226 -0
  376. k3_node/models/test_graphormer.py +233 -0
  377. k3_node/models/test_graphormer3d.py +163 -0
  378. k3_node/models/test_grover.py +287 -0
  379. k3_node/models/test_jumping_knowledge.py +129 -0
  380. k3_node/models/test_label_prop.py +37 -0
  381. k3_node/models/test_lightgcn.py +38 -0
  382. k3_node/models/test_linkx.py +31 -0
  383. k3_node/models/test_lpformer.py +22 -0
  384. k3_node/models/test_mask_label.py +90 -0
  385. k3_node/models/test_meta.py +159 -0
  386. k3_node/models/test_metapath2vec.py +45 -0
  387. k3_node/models/test_mlp.py +62 -0
  388. k3_node/models/test_mole_bert.py +164 -0
  389. k3_node/models/test_neural_fingerprint.py +13 -0
  390. k3_node/models/test_node2vec.py +57 -0
  391. k3_node/models/test_pmlp.py +81 -0
  392. k3_node/models/test_polynormer.py +104 -0
  393. k3_node/models/test_rect.py +23 -0
  394. k3_node/models/test_renet.py +32 -0
  395. k3_node/models/test_rev_gnn.py +24 -0
  396. k3_node/models/test_schnet.py +43 -0
  397. k3_node/models/test_sgformer.py +48 -0
  398. k3_node/models/test_signed_gcn.py +28 -0
  399. k3_node/models/test_tgn.py +77 -0
  400. k3_node/models/test_unimol.py +179 -0
  401. k3_node/models/test_unimol2.py +114 -0
  402. k3_node/models/test_unimol_plus.py +131 -0
  403. k3_node/models/test_visnet.py +44 -0
  404. k3_node/models/tgn.py +382 -0
  405. k3_node/models/unimol.py +1156 -0
  406. k3_node/models/unimol2.py +616 -0
  407. k3_node/models/unimol_docking_v2.py +301 -0
  408. k3_node/models/unimol_plus.py +456 -0
  409. k3_node/models/utils.py +97 -0
  410. k3_node/models/visnet.py +759 -0
  411. k3_node/ops/__init__.py +4 -0
  412. k3_node/ops/conv.py +56 -0
  413. k3_node/ops/creation.py +43 -0
  414. k3_node/ops/graph.py +27 -0
  415. k3_node/ops/host.py +41 -0
  416. k3_node/ops/matmul.py +49 -0
  417. k3_node/ops/numpy.py +24 -0
  418. k3_node/ops/segment.py +54 -0
  419. k3_node/ops/sparse.py +51 -0
  420. k3_node/rag/__init__.py +49 -0
  421. k3_node/rag/encoders.py +312 -0
  422. k3_node/rag/pipeline.py +192 -0
  423. k3_node/rag/projector.py +184 -0
  424. k3_node/rag/subgraph.py +270 -0
  425. k3_node/rag/test_rag.py +347 -0
  426. k3_node/rag/verbalizer.py +162 -0
  427. k3_node/tasks/__init__.py +19 -0
  428. k3_node/tasks/backbone_resolver.py +125 -0
  429. k3_node/tasks/base.py +67 -0
  430. k3_node/tasks/graph_classification.py +270 -0
  431. k3_node/tasks/graph_regression.py +228 -0
  432. k3_node/tasks/link_prediction.py +306 -0
  433. k3_node/tasks/node_classification.py +194 -0
  434. k3_node/tasks/node_regression.py +138 -0
  435. k3_node/tasks/test_tasks.py +319 -0
  436. k3_node/test_docstring_examples.py +106 -0
  437. k3_node/test_training_forwarding.py +116 -0
  438. k3_node/training.py +115 -0
  439. k3_node/transforms/__init__.py +166 -0
  440. k3_node/transforms/base_transform.py +32 -0
  441. k3_node/transforms/compose.py +58 -0
  442. k3_node/transforms/general.py +676 -0
  443. k3_node/transforms/graph.py +1070 -0
  444. k3_node/transforms/spatial.py +797 -0
  445. k3_node/transforms/test_random_link_split.py +45 -0
  446. k3_node/transforms/test_spatial_transforms.py +65 -0
  447. k3_node/transforms/test_transforms.py +253 -0
  448. k3_node/transforms/utils.py +102 -0
  449. k3_node/utils/__init__.py +5 -0
  450. k3_node/utils/backend_import.py +12 -0
  451. k3_node/utils/graph.py +286 -0
  452. k3_node/utils/keras.py +94 -0
  453. k3_node/utils/random.py +103 -0
  454. k3_node/utils/smiles.py +235 -0
  455. k3_node-1.0.0.dist-info/METADATA +284 -0
  456. k3_node-1.0.0.dist-info/RECORD +459 -0
  457. k3_node-1.0.0.dist-info/WHEEL +5 -0
  458. k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
  459. k3_node-1.0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,137 @@
1
+ import os.path as osp
2
+ from typing import Callable, List, Optional
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ from k3_node.data import InMemoryDataset
7
+ from k3_node.io import fs, read_planetoid_data
8
+
9
+
10
+ class Planetoid(InMemoryDataset):
11
+ r"""The citation network datasets "Cora", "CiteSeer" and "PubMed" from the
12
+ "Revisiting Semi-Supervised Learning with Graph Embeddings" paper.
13
+ Nodes represent documents and edges represent citation links.
14
+ Training, validation and test splits are given by binary masks.
15
+
16
+ Args:
17
+ root (str): Root directory where the dataset should be saved.
18
+ name (str): The name of the dataset ("Cora", "CiteSeer", "PubMed").
19
+ split (str, optional): The type of dataset split ("public", "full", "geom-gcn", "random").
20
+ (default: "public")
21
+ num_train_per_class (int, optional): The number of training samples per class for "random" split.
22
+ (default: 20)
23
+ num_val (int, optional): The number of validation samples for "random" split. (default: 500)
24
+ num_test (int, optional): The number of test samples for "random" split. (default: 1000)
25
+ transform (callable, optional): A function/transform that takes in a Data object and returns a transformed version.
26
+ pre_transform (callable, optional): A function/transform that takes in a Data object and returns a transformed version.
27
+ force_reload (bool, optional): Whether to re-process the dataset. (default: False)
28
+ """
29
+
30
+ url = "https://github.com/kimiyoung/planetoid/raw/master/data"
31
+ geom_gcn_url = "https://raw.githubusercontent.com/graphdml-uiuc-jlu/geom-gcn/master"
32
+
33
+ def __init__(
34
+ self,
35
+ root: str,
36
+ name: str,
37
+ split: str = "public",
38
+ num_train_per_class: int = 20,
39
+ num_val: int = 500,
40
+ num_test: int = 1000,
41
+ transform: Optional[Callable] = None,
42
+ pre_transform: Optional[Callable] = None,
43
+ force_reload: bool = False,
44
+ ):
45
+ self.name = name
46
+ self.split = split.lower()
47
+ assert self.split in ["public", "full", "geom-gcn", "random"]
48
+
49
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
50
+ self.load(self.processed_paths[0])
51
+
52
+ if self.split == "full":
53
+ data = self.get(0)
54
+ val_m = ops.convert_to_numpy(data.val_mask)
55
+ test_m = ops.convert_to_numpy(data.test_mask)
56
+ train_mask = np.ones(data.num_nodes, dtype=bool)
57
+ train_mask[val_m | test_m] = False
58
+ data.train_mask = ops.convert_to_tensor(train_mask, dtype="bool")
59
+ self.data, self.slices = self.collate([data])
60
+
61
+ elif self.split == "random":
62
+ data = self.get(0)
63
+ num_nodes = data.num_nodes
64
+ y_np = ops.convert_to_numpy(data.y)
65
+ train_mask = np.zeros(num_nodes, dtype=bool)
66
+ for c in range(self.num_classes):
67
+ idx = np.where(y_np == c)[0]
68
+ perm = np.random.permutation(len(idx))
69
+ train_mask[idx[perm[:num_train_per_class]]] = True
70
+
71
+ remaining = np.where(~train_mask)[0]
72
+ remaining = remaining[np.random.permutation(len(remaining))]
73
+
74
+ val_mask = np.zeros(num_nodes, dtype=bool)
75
+ val_mask[remaining[:num_val]] = True
76
+
77
+ test_mask = np.zeros(num_nodes, dtype=bool)
78
+ test_mask[remaining[num_val : num_val + num_test]] = True
79
+
80
+ data.train_mask = ops.convert_to_tensor(train_mask, dtype="bool")
81
+ data.val_mask = ops.convert_to_tensor(val_mask, dtype="bool")
82
+ data.test_mask = ops.convert_to_tensor(test_mask, dtype="bool")
83
+
84
+ self.data, self.slices = self.collate([data])
85
+
86
+ @property
87
+ def raw_dir(self) -> str:
88
+ if self.split == "geom-gcn":
89
+ return osp.join(self.root, self.name, "geom-gcn", "raw")
90
+ return osp.join(self.root, self.name, "raw")
91
+
92
+ @property
93
+ def processed_dir(self) -> str:
94
+ if self.split == "geom-gcn":
95
+ return osp.join(self.root, self.name, "geom-gcn", "processed")
96
+ return osp.join(self.root, self.name, "processed")
97
+
98
+ @property
99
+ def raw_file_names(self) -> List[str]:
100
+ names = ["x", "tx", "allx", "y", "ty", "ally", "graph", "test.index"]
101
+ return [f"ind.{self.name.lower()}.{name}" for name in names]
102
+
103
+ @property
104
+ def processed_file_names(self) -> str:
105
+ return "data.pt"
106
+
107
+ def download(self):
108
+ for name in self.raw_file_names:
109
+ fs.cp(f"{self.url}/{name}", self.raw_dir)
110
+ if self.split == "geom-gcn":
111
+ for i in range(10):
112
+ url = f"{self.geom_gcn_url}/splits/{self.name.lower()}"
113
+ fs.cp(f"{url}_split_0.6_0.2_{i}.npz", self.raw_dir)
114
+
115
+ def process(self):
116
+ data = read_planetoid_data(self.raw_dir, self.name)
117
+
118
+ if self.split == "geom-gcn":
119
+ train_masks, val_masks, test_masks = [], [], []
120
+ for i in range(10):
121
+ name = f"{self.name.lower()}_split_0.6_0.2_{i}.npz"
122
+ splits = np.load(osp.join(self.raw_dir, name))
123
+ train_masks.append(splits["train_mask"])
124
+ val_masks.append(splits["val_mask"])
125
+ test_masks.append(splits["test_mask"])
126
+ data.train_mask = ops.convert_to_tensor(np.stack(train_masks, axis=1), dtype="bool")
127
+ data.val_mask = ops.convert_to_tensor(np.stack(val_masks, axis=1), dtype="bool")
128
+ data.test_mask = ops.convert_to_tensor(np.stack(test_masks, axis=1), dtype="bool")
129
+
130
+ if self.pre_transform is not None:
131
+ data = self.pre_transform(data)
132
+
133
+ self.save([data], self.processed_paths[0])
134
+
135
+ def __repr__(self) -> str:
136
+ return f"{self.name}()"
137
+
@@ -0,0 +1,63 @@
1
+ import os
2
+ import os.path as osp
3
+ from typing import Callable, List, Optional
4
+ import numpy as np
5
+ from keras import ops
6
+
7
+ from k3_node.data import Data, InMemoryDataset
8
+ from k3_node.io import fs
9
+
10
+
11
+ class PolBlogs(InMemoryDataset):
12
+ r"""The Political Blogs dataset containing 1,490 vertices and 19,025 edges.
13
+
14
+ Args:
15
+ root (str): Root directory where the dataset should be saved.
16
+ transform (callable, optional): Transform function.
17
+ pre_transform (callable, optional): Pre-transform function.
18
+ force_reload (bool, optional): Whether to re-process the dataset.
19
+ """
20
+
21
+ url = "https://netset.telecom-paris.fr/datasets/polblogs.tar.gz"
22
+
23
+ def __init__(
24
+ self,
25
+ root: str,
26
+ transform: Optional[Callable] = None,
27
+ pre_transform: Optional[Callable] = None,
28
+ force_reload: bool = False,
29
+ ):
30
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
31
+ self.load(self.processed_paths[0])
32
+
33
+ @property
34
+ def raw_file_names(self) -> List[str]:
35
+ return ["adjacency.tsv", "labels.tsv"]
36
+
37
+ @property
38
+ def processed_file_names(self) -> str:
39
+ return "data.pt"
40
+
41
+ def download(self):
42
+ tar_path = osp.join(self.raw_dir, "polblogs.tar.gz")
43
+ fs.cp(self.url, tar_path, extract=True)
44
+ if osp.exists(tar_path):
45
+ fs.rm(tar_path)
46
+
47
+ def process(self):
48
+ adj = np.loadtxt(self.raw_paths[0], delimiter="\t", usecols=(0, 1), dtype=np.int64)
49
+ edge_index = adj.T
50
+
51
+ y = np.loadtxt(self.raw_paths[1], delimiter="\t", usecols=(1,), dtype=np.int64)
52
+
53
+ data = Data(
54
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
55
+ y=ops.convert_to_tensor(y, dtype="int64"),
56
+ num_nodes=int(y.shape[0]),
57
+ )
58
+
59
+ if self.pre_transform is not None:
60
+ data = self.pre_transform(data)
61
+
62
+ self.save([data], self.processed_paths[0])
63
+
@@ -0,0 +1,189 @@
1
+ import json
2
+ import os
3
+ import os.path as osp
4
+ from itertools import product
5
+ from typing import Callable, List, Optional
6
+
7
+ import numpy as np
8
+ from keras import ops
9
+
10
+ from k3_node.data import (
11
+ Data,
12
+ InMemoryDataset,
13
+ download_url,
14
+ extract_zip,
15
+ )
16
+ from k3_node.layers.conv.utils import remove_self_loops
17
+
18
+
19
+ class PPI(InMemoryDataset):
20
+ r"""The protein-protein interaction networks from the `"Predicting
21
+ Multicellular Function through Multi-layer Tissue Networks"
22
+ <https://arxiv.org/abs/1707.04638>`_ paper, containing positional gene
23
+ sets, motif gene sets and immunological signatures as features (50 in
24
+ total) and gene ontology sets as labels (121 in total).
25
+
26
+ Args:
27
+ root (str): Root directory where the dataset should be saved.
28
+ split (str, optional): If :obj:`"train"`, loads the training dataset.
29
+ If :obj:`"val"`, loads the validation dataset.
30
+ If :obj:`"test"`, loads the test dataset. (default: :obj:`"train"`)
31
+ transform (callable, optional): A function/transform that takes in an
32
+ :obj:`k3_node.data.Data` object and returns a transformed
33
+ version. The data object will be transformed before every access.
34
+ (default: :obj:`None`)
35
+ pre_transform (callable, optional): A function/transform that takes in
36
+ an :obj:`k3_node.data.Data` object and returns a
37
+ transformed version. The data object will be transformed before
38
+ being saved to disk. (default: :obj:`None`)
39
+ pre_filter (callable, optional): A function that takes in an
40
+ :obj:`k3_node.data.Data` object and returns a boolean
41
+ value, indicating whether the data object should be included in the
42
+ final dataset. (default: :obj:`None`)
43
+ force_reload (bool, optional): Whether to re-process the dataset.
44
+ (default: :obj:`False`)
45
+
46
+ **STATS:**
47
+
48
+ .. list-table::
49
+ :widths: 10 10 10 10 10
50
+ :header-rows: 1
51
+
52
+ * - #graphs
53
+ - #nodes
54
+ - #edges
55
+ - #features
56
+ - #tasks
57
+ * - 20
58
+ - ~2,245.3
59
+ - ~61,318.4
60
+ - 50
61
+ - 121
62
+ """
63
+
64
+ url = "https://data.dgl.ai/dataset/ppi.zip"
65
+
66
+ def __init__(
67
+ self,
68
+ root: str,
69
+ split: str = "train",
70
+ transform: Optional[Callable] = None,
71
+ pre_transform: Optional[Callable] = None,
72
+ pre_filter: Optional[Callable] = None,
73
+ force_reload: bool = False,
74
+ ) -> None:
75
+ assert split.lower() in ["train", "val", "valid", "test"], f"Invalid split '{split}'"
76
+ self.split = "val" if split.lower() == "valid" else split.lower()
77
+
78
+ super().__init__(
79
+ root,
80
+ transform,
81
+ pre_transform,
82
+ pre_filter,
83
+ force_reload=force_reload,
84
+ )
85
+
86
+ if self.split == "train":
87
+ self.load(self.processed_paths[0])
88
+ elif self.split == "val":
89
+ self.load(self.processed_paths[1])
90
+ elif self.split == "test":
91
+ self.load(self.processed_paths[2])
92
+
93
+ @property
94
+ def raw_file_names(self) -> List[str]:
95
+ splits = ["train", "valid", "test"]
96
+ files = ["feats.npy", "graph_id.npy", "graph.json", "labels.npy"]
97
+ return [f"{split}_{name}" for split, name in product(splits, files)]
98
+
99
+ @property
100
+ def processed_file_names(self) -> List[str]:
101
+ return ["train.pt", "val.pt", "test.pt"]
102
+
103
+ def download(self) -> None:
104
+ path = download_url(self.url, self.root)
105
+ extract_zip(path, self.raw_dir)
106
+ if osp.exists(path):
107
+ os.unlink(path)
108
+
109
+ def process(self) -> None:
110
+ try:
111
+ import networkx as nx
112
+ from networkx.readwrite import json_graph
113
+ has_networkx = True
114
+ except ImportError:
115
+ has_networkx = False
116
+
117
+ for s, split in enumerate(["train", "valid", "test"]):
118
+ path = osp.join(self.raw_dir, f"{split}_graph.json")
119
+ with open(path, "r", encoding="utf-8") as f:
120
+ graph_json = json.load(f)
121
+
122
+ if has_networkx:
123
+ try:
124
+ G = nx.DiGraph(json_graph.node_link_graph(graph_json, edges="links"))
125
+ except TypeError:
126
+ G = nx.DiGraph(json_graph.node_link_graph(graph_json))
127
+ else:
128
+ links = graph_json.get("links", [])
129
+ src_all = np.array([link["source"] for link in links], dtype=np.int64)
130
+ dst_all = np.array([link["target"] for link in links], dtype=np.int64)
131
+
132
+ x_np = np.load(osp.join(self.raw_dir, f"{split}_feats.npy"))
133
+ y_np = np.load(osp.join(self.raw_dir, f"{split}_labels.npy"))
134
+
135
+ data_list = []
136
+ path = osp.join(self.raw_dir, f"{split}_graph_id.npy")
137
+ idx = np.load(path)
138
+ idx = idx - idx.min()
139
+
140
+ for i in range(int(idx.max()) + 1):
141
+ mask = (idx == i)
142
+ node_indices = np.where(mask)[0]
143
+
144
+ if has_networkx:
145
+ G_s = G.subgraph(node_indices.tolist())
146
+ edges = list(G_s.edges)
147
+ if len(edges) > 0:
148
+ edge_index = np.array(edges, dtype=np.int64).T
149
+ edge_index = edge_index - edge_index.min()
150
+ else:
151
+ edge_index = np.zeros((2, 0), dtype=np.int64)
152
+ else:
153
+ min_node, max_node = node_indices.min(), node_indices.max()
154
+ edge_mask = (
155
+ (src_all >= min_node)
156
+ & (src_all <= max_node)
157
+ & (dst_all >= min_node)
158
+ & (dst_all <= max_node)
159
+ )
160
+ sub_src = src_all[edge_mask] - min_node
161
+ sub_dst = dst_all[edge_mask] - min_node
162
+ edge_index = np.stack([sub_src, sub_dst], axis=0)
163
+
164
+ edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
165
+ edge_index, _ = remove_self_loops(edge_index)
166
+
167
+ x = ops.convert_to_tensor(x_np[mask], dtype="float32")
168
+ y = ops.convert_to_tensor(y_np[mask], dtype="float32")
169
+
170
+ data = Data(edge_index=edge_index, x=x, y=y)
171
+
172
+ if self.pre_filter is not None and not self.pre_filter(data):
173
+ continue
174
+
175
+ if self.pre_transform is not None:
176
+ data = self.pre_transform(data)
177
+
178
+ data_list.append(data)
179
+
180
+ self.save(data_list, self.processed_paths[s])
181
+
182
+ @property
183
+ def num_classes(self) -> int:
184
+ data = self.get(0)
185
+ y = getattr(data, "y", None)
186
+ if y is not None and len(y.shape) > 1:
187
+ return y.shape[-1]
188
+ return super().num_classes
189
+
@@ -0,0 +1,65 @@
1
+ from typing import Callable, Optional
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+ from k3_node.io import fs
7
+
8
+
9
+ class QM7b(InMemoryDataset):
10
+ r"""The QM7b dataset consisting of 7,211 molecules with 14 regression targets."""
11
+
12
+ url = "https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/qm7b.mat"
13
+
14
+ def __init__(
15
+ self,
16
+ root: str,
17
+ transform: Optional[Callable] = None,
18
+ pre_transform: Optional[Callable] = None,
19
+ pre_filter: Optional[Callable] = None,
20
+ force_reload: bool = False,
21
+ ):
22
+ super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
23
+ self.load(self.processed_paths[0])
24
+
25
+ @property
26
+ def raw_file_names(self) -> str:
27
+ return "qm7b.mat"
28
+
29
+ @property
30
+ def processed_file_names(self) -> str:
31
+ return "data.pt"
32
+
33
+ def download(self):
34
+ fs.cp(self.url, self.raw_dir)
35
+
36
+ def process(self):
37
+ from scipy.io import loadmat
38
+
39
+ data = loadmat(self.raw_paths[0])
40
+ coulomb_matrix = data["X"]
41
+ target = data["T"].astype(np.float32)
42
+
43
+ data_list = []
44
+ for i in range(target.shape[0]):
45
+ nz = np.nonzero(coulomb_matrix[i])
46
+ edge_index = np.stack([nz[0], nz[1]], axis=0).astype(np.int64)
47
+ edge_attr = coulomb_matrix[i, edge_index[0], edge_index[1]].astype(np.float32)
48
+ y = target[i].reshape(1, -1)
49
+ num_nodes = int(np.max(edge_index)) + 1 if edge_index.size > 0 else 0
50
+
51
+ d = Data(
52
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
53
+ edge_attr=ops.convert_to_tensor(edge_attr, dtype="float32"),
54
+ y=ops.convert_to_tensor(y, dtype="float32"),
55
+ num_nodes=num_nodes,
56
+ )
57
+ data_list.append(d)
58
+
59
+ if self.pre_filter is not None:
60
+ data_list = [d for d in data_list if self.pre_filter(d)]
61
+ if self.pre_transform is not None:
62
+ data_list = [self.pre_transform(d) for d in data_list]
63
+
64
+ self.save(data_list, self.processed_paths[0])
65
+
@@ -0,0 +1,132 @@
1
+ import os
2
+ import os.path as osp
3
+ from typing import Callable, List, Optional
4
+
5
+ import numpy as np
6
+
7
+ from k3_node.data import Data, InMemoryDataset
8
+ from k3_node.data.download import download_url
9
+ from k3_node.data.extract import extract_zip
10
+
11
+ HAR2EV = 27.211386246
12
+ KCALMOL2EV = 0.04336414
13
+ CONVERSION = np.array([1., 1., HAR2EV, HAR2EV, HAR2EV, 1., HAR2EV, HAR2EV, HAR2EV, HAR2EV, HAR2EV,
14
+ 1., KCALMOL2EV, KCALMOL2EV, KCALMOL2EV, KCALMOL2EV, 1., 1., 1.], dtype=np.float32)
15
+ ATOMREFS = {
16
+ 6: [0., 0., 0., 0., 0.],
17
+ 7: [-13.61312172, -1029.86312267, -1485.30251237, -2042.61123593, -2713.48485589],
18
+ 8: [-13.5745904, -1029.82456413, -1485.26398105, -2042.5727046, -2713.44632457],
19
+ 9: [-13.54887564, -1029.79887659, -1485.2382935, -2042.54701705, -2713.42063702],
20
+ 10: [-13.90303183, -1030.25891228, -1485.71166277, -2043.01812778, -2713.88796536],
21
+ 11: [0., 0., 0., 0., 0.],
22
+ }
23
+
24
+
25
+ class QM9(InMemoryDataset):
26
+ r"""The QM9 dataset: about 130,000 small organic molecules with their 3D structure and 19
27
+ regression targets (dipole moment, HOMO/LUMO energies, internal energy, ...), as in PyG.
28
+
29
+ Every molecule has 11 atom features ``x`` (one-hot H/C/N/O/F, atomic number, aromatic,
30
+ sp/sp2/sp3 hybridization, number of hydrogens), atomic numbers ``z``, positions ``pos``,
31
+ one-hot bond types ``edge_attr`` and the targets ``y`` of shape ``[1, 19]`` (energies in eV).
32
+ Requires RDKit to process the raw files.
33
+
34
+ Args:
35
+ root (str): Root directory where the dataset should be saved.
36
+ transform (callable, optional): A function applied to each graph when it is accessed.
37
+ pre_transform (callable, optional): A function applied to each graph before saving.
38
+ pre_filter (callable, optional): A function deciding which graphs to keep.
39
+ force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
40
+ """
41
+
42
+ raw_url = 'https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/molnet_publish/qm9.zip'
43
+ raw_url2 = 'https://ndownloader.figshare.com/files/3195404'
44
+
45
+ def __init__(self, root: str, transform: Optional[Callable] = None, pre_transform: Optional[Callable] = None,
46
+ pre_filter: Optional[Callable] = None, force_reload: bool = False):
47
+ super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
48
+ self.load(self.processed_paths[0])
49
+
50
+ def mean(self, target: int) -> float:
51
+ return float(np.asarray(self._data.y)[:, target].mean())
52
+
53
+ def std(self, target: int) -> float:
54
+ return float(np.asarray(self._data.y)[:, target].std(ddof=1))
55
+
56
+ def atomref(self, target: int) -> Optional[np.ndarray]:
57
+ r"""Per-element reference energies (a ``[100, 1]`` array indexed by atomic number), or
58
+ ``None`` for targets without them."""
59
+ if target not in ATOMREFS:
60
+ return None
61
+ out = np.zeros((100, 1), dtype=np.float32)
62
+ out[[1, 6, 7, 8, 9], 0] = ATOMREFS[target]
63
+ return out
64
+
65
+ @property
66
+ def raw_file_names(self) -> List[str]:
67
+ return ['gdb9.sdf', 'gdb9.sdf.csv', 'uncharacterized.txt']
68
+
69
+ @property
70
+ def processed_file_names(self) -> str:
71
+ return 'data_v3.pt'
72
+
73
+ def download(self):
74
+ path = download_url(self.raw_url, self.raw_dir)
75
+ extract_zip(path, self.raw_dir)
76
+ os.unlink(path)
77
+ download_url(self.raw_url2, self.raw_dir)
78
+ os.rename(osp.join(self.raw_dir, '3195404'), osp.join(self.raw_dir, 'uncharacterized.txt'))
79
+
80
+ def process(self):
81
+ from rdkit import Chem, RDLogger
82
+ from rdkit.Chem.rdchem import BondType as BT
83
+ from rdkit.Chem.rdchem import HybridizationType as HT
84
+
85
+ RDLogger.DisableLog('rdApp.*')
86
+ types = {'H': 0, 'C': 1, 'N': 2, 'O': 3, 'F': 4}
87
+ bonds = {BT.SINGLE: 0, BT.DOUBLE: 1, BT.TRIPLE: 2, BT.AROMATIC: 3}
88
+
89
+ with open(self.raw_paths[1]) as f:
90
+ target = np.array([[float(x) for x in line.split(',')[1:20]] for line in f.read().split('\n')[1:-1]],
91
+ dtype=np.float32)
92
+ target = np.concatenate([target[:, 3:], target[:, :3]], axis=-1) * CONVERSION
93
+ with open(self.raw_paths[2]) as f:
94
+ skip = {int(x.split()[0]) - 1 for x in f.read().split('\n')[9:-2]}
95
+
96
+ data_list = []
97
+ for i, mol in enumerate(Chem.SDMolSupplier(self.raw_paths[0], removeHs=False, sanitize=False)):
98
+ if i in skip:
99
+ continue
100
+ N = mol.GetNumAtoms()
101
+ pos = mol.GetConformer().GetPositions().astype(np.float32)
102
+ atoms = list(mol.GetAtoms())
103
+ z = np.array([a.GetAtomicNum() for a in atoms], dtype=np.int64)
104
+ rows, cols, edge_types = [], [], []
105
+ for bond in mol.GetBonds():
106
+ s, e = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
107
+ rows += [s, e]
108
+ cols += [e, s]
109
+ edge_types += 2 * [bonds[bond.GetBondType()]]
110
+ edge_index = np.array([rows, cols], dtype=np.int64).reshape(2, -1)
111
+ edge_type = np.array(edge_types, dtype=np.int64)
112
+ perm = np.argsort(edge_index[0] * N + edge_index[1], kind="stable")
113
+ edge_index, edge_type = edge_index[:, perm], edge_type[perm]
114
+
115
+ num_hs = np.zeros(N, dtype=np.float32)
116
+ np.add.at(num_hs, edge_index[1], (z == 1).astype(np.float32)[edge_index[0]])
117
+ hybrid = [a.GetHybridization() for a in atoms]
118
+ x = np.concatenate([
119
+ np.eye(len(types), dtype=np.float32)[[types[a.GetSymbol()] for a in atoms]],
120
+ np.stack([z, [a.GetIsAromatic() for a in atoms], [h == HT.SP for h in hybrid],
121
+ [h == HT.SP2 for h in hybrid], [h == HT.SP3 for h in hybrid], num_hs], axis=1).astype(np.float32),
122
+ ], axis=1)
123
+ data = Data(x=x, z=z, pos=pos, edge_index=edge_index,
124
+ edge_attr=np.eye(len(bonds), dtype=np.float32)[edge_type].reshape(-1, len(bonds)),
125
+ y=target[i][None], smiles=Chem.MolToSmiles(mol, isomericSmiles=True),
126
+ name=mol.GetProp('_Name'), idx=i)
127
+ if self.pre_filter is not None and not self.pre_filter(data):
128
+ continue
129
+ if self.pre_transform is not None:
130
+ data = self.pre_transform(data)
131
+ data_list.append(data)
132
+ self.save(data_list, self.processed_paths[0])