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,128 @@
1
+ from typing import List, Optional, Union
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+
7
+ class GroupAddRev(keras.layers.Layer):
8
+ r"""The Grouped Reversible GNN module from the `"Graph Neural Networks with
9
+ 1000 Layers" <https://arxiv.org/abs/2106.07476>`_ paper.
10
+
11
+ Args:
12
+ conv (keras.layers.Layer or List[keras.layers.Layer]): A seed GNN layer
13
+ or list of GNN layers.
14
+ split_dim (int, optional): The dimension across which to split groups.
15
+ (default: :obj:`-1`)
16
+ num_groups (int, optional): The number of groups. (default: :obj:`None`)
17
+
18
+ Example:
19
+ ```python
20
+ import numpy as np
21
+ from k3_node.layers import GCNConv
22
+ from k3_node.models import GroupAddRev
23
+
24
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
25
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
26
+
27
+ x = np.random.rand(10, 32).astype("float32")
28
+ model = GroupAddRev([GCNConv(16, 16), GCNConv(16, 16)]) # channels are split into 2 groups
29
+ out = model(x, edge_index=edge_index)
30
+ print(tuple(out.shape)) # (10, 32)
31
+ print(tuple(model.inverse(out, edge_index=edge_index).shape)) # (10, 32): reversible: recover the input
32
+ ```
33
+ """
34
+ def __init__(
35
+ self,
36
+ *args,
37
+ split_dim: int = -1,
38
+ num_groups: int = 2,
39
+ **kwargs,
40
+ ):
41
+ super().__init__(**kwargs)
42
+ self.split_dim = split_dim
43
+
44
+ if len(args) == 1 and isinstance(args[0], (list, tuple)):
45
+ self.convs = list(args[0])
46
+ elif len(args) > 1:
47
+ self.convs = list(args)
48
+ elif len(args) == 1:
49
+ conv = args[0]
50
+ assert num_groups is not None, "Please specify 'num_groups'"
51
+ self.convs = [conv]
52
+ for _ in range(num_groups - 1):
53
+ try:
54
+ cloned = conv.__class__.from_config(conv.get_config())
55
+ except Exception:
56
+ cloned = copy.deepcopy(conv)
57
+ self.convs.append(cloned)
58
+ else:
59
+ raise ValueError("GroupAddRev requires at least one layer argument")
60
+
61
+ if len(self.convs) < 2:
62
+ raise ValueError(f"The number of groups should not be smaller than '2' (got '{self.num_groups}')")
63
+
64
+ @property
65
+ def num_groups(self) -> int:
66
+ return len(self.convs)
67
+
68
+ def build(self, input_shape=None):
69
+ self.built = True
70
+
71
+ def reset_parameters(self):
72
+ for conv in self.convs:
73
+ if hasattr(conv, "reset_parameters"):
74
+ conv.reset_parameters()
75
+
76
+ def call(self, x, edge_index=None, *args):
77
+ if edge_index is None and isinstance(x, (tuple, list)):
78
+ if len(x) >= 2:
79
+ x, edge_index = x[0], x[1]
80
+ xs = ops.split(x, self.num_groups, axis=self.split_dim)
81
+ group_args = self._chunk_args(x, args)
82
+
83
+ ys = []
84
+ y_in = xs[1]
85
+ for item in xs[2:]:
86
+ y_in = y_in + item
87
+
88
+ for i in range(self.num_groups):
89
+ conv_out = self.convs[i](y_in, edge_index, *group_args[i])
90
+ y_in = xs[i] + conv_out
91
+ ys.append(y_in)
92
+
93
+ return ops.concatenate(ys, axis=self.split_dim)
94
+
95
+ def _chunk_args(self, x, args):
96
+ """As in PyG, extra tensor arguments shaped like ``x`` (e.g. a dropout mask) are split into
97
+ one chunk per group; other arguments are passed to every group unchanged."""
98
+ channels = x.shape[self.split_dim]
99
+ chunked = []
100
+ for arg in args:
101
+ if hasattr(arg, "shape") and len(arg.shape) == len(x.shape) and arg.shape[self.split_dim] == channels:
102
+ chunked.append(ops.split(arg, self.num_groups, axis=self.split_dim))
103
+ else:
104
+ chunked.append([arg] * self.num_groups)
105
+ return [[c[i] for c in chunked] for i in range(self.num_groups)]
106
+
107
+ def inverse(self, y, edge_index, *args):
108
+ ys = ops.split(y, self.num_groups, axis=self.split_dim)
109
+ group_args = self._chunk_args(y, args)
110
+
111
+ xs = []
112
+ for i in range(self.num_groups - 1, -1, -1):
113
+ if i != 0:
114
+ y_in = ys[i - 1]
115
+ else:
116
+ y_in = xs[0]
117
+ for item in xs[1:]:
118
+ y_in = y_in + item
119
+ conv_out = self.convs[i](y_in, edge_index, *group_args[i])
120
+ x_i = ys[i] - conv_out
121
+ xs.append(x_i)
122
+
123
+ return ops.concatenate(xs[::-1], axis=self.split_dim)
124
+
125
+ def __repr__(self) -> str:
126
+ return (f'{self.__class__.__name__}({self.convs[0]}, '
127
+ f'num_groups={self.num_groups})')
128
+
@@ -0,0 +1,484 @@
1
+ import numpy as np
2
+ import keras
3
+ from keras import ops
4
+ from typing import Optional, Callable
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+ from k3_node.layers.pool import radius_graph, global_add_pool, global_mean_pool
8
+ from k3_node.hub.hub_mixin import K3NodeHubMixin
9
+
10
+
11
+ DEFAULT_ATOMIC_MASSES = [
12
+ 0.0, 1.008, 4.0026, 6.94, 9.0122, 10.81, 12.011, 14.007, 15.999, 18.998,
13
+ 20.180, 22.990, 24.305, 26.982, 28.085, 30.974, 32.06, 35.45, 39.95,
14
+ 39.098, 40.078, 44.956, 47.867, 50.942, 51.996, 54.938, 55.845, 58.933,
15
+ 58.693, 63.546, 65.38, 69.723, 72.630, 74.922, 78.971, 79.904, 83.798,
16
+ 85.468, 87.62, 88.906, 91.224, 92.906, 95.95, 98.0, 101.07, 102.91,
17
+ 106.42, 107.87, 112.41, 114.82, 118.71, 121.76, 127.60, 126.90, 131.29,
18
+ 132.91, 137.33, 138.91, 140.12, 140.91, 144.24, 145.0, 150.36, 151.96,
19
+ 157.25, 158.93, 162.50, 164.93, 167.26, 168.93, 173.05, 174.97, 178.49,
20
+ 180.95, 183.84, 186.21, 190.23, 192.22, 195.08, 196.97, 200.59, 204.38,
21
+ 207.2, 208.98, 209.0, 210.0, 222.0, 223.0, 226.0, 227.0, 232.04,
22
+ 231.04, 238.03, 237.0, 244.0, 243.0, 247.0, 247.0, 251.0, 252.0,
23
+ ]
24
+
25
+
26
+ class ShiftedSoftplus(keras.layers.Layer):
27
+ r"""Shifted softplus activation function: :math:`\ln(1 + e^x) - \ln(2)`.
28
+
29
+ Example:
30
+ ```python
31
+ import numpy as np
32
+ from k3_node.models import ShiftedSoftplus
33
+
34
+ x = np.array([-1.0, 0.0, 1.0], dtype="float32")
35
+ print(tuple(ShiftedSoftplus()(x).shape)) # (3,): softplus(x) - log(2)
36
+ ```
37
+ """
38
+ def __init__(self, **kwargs):
39
+ super().__init__(**kwargs)
40
+ self.shift = float(np.log(2.0))
41
+
42
+ def call(self, x):
43
+ return ops.softplus(x) - self.shift
44
+
45
+
46
+ class GaussianSmearing(keras.layers.Layer):
47
+ r"""Smears interatomic distances using Gaussian basis functions.
48
+
49
+ Example:
50
+ ```python
51
+ import numpy as np
52
+ from k3_node.models import GaussianSmearing
53
+
54
+ dist = np.array([0.9, 1.5, 3.2], dtype="float32")
55
+ print(tuple(GaussianSmearing(start=0.0, stop=5.0, num_gaussians=10)(dist).shape)) # (3, 10)
56
+ ```
57
+ """
58
+ def __init__(
59
+ self,
60
+ start: float = 0.0,
61
+ stop: float = 5.0,
62
+ num_gaussians: int = 50,
63
+ **kwargs,
64
+ ):
65
+ super().__init__(**kwargs)
66
+ self.start = start
67
+ self.stop = stop
68
+ self.num_gaussians = num_gaussians
69
+
70
+ offset = np.linspace(start, stop, num_gaussians, dtype=np.float32)
71
+ diff = float(offset[1] - offset[0])
72
+ self.coeff = -0.5 / (diff ** 2)
73
+ self.offset = self.add_weight(
74
+ name="offset",
75
+ shape=(num_gaussians,),
76
+ initializer=keras.initializers.Constant(offset),
77
+ trainable=False,
78
+ dtype="float32",
79
+ )
80
+
81
+ def call(self, dist):
82
+ dist = ops.expand_dims(dist, -1) - ops.expand_dims(self.offset, 0)
83
+ return ops.exp(self.coeff * ops.power(dist, 2))
84
+
85
+
86
+ class RadiusInteractionGraph(keras.layers.Layer):
87
+ r"""Creates edges based on atom positions :obj:`pos` to all points within
88
+ the cutoff distance.
89
+
90
+ Example:
91
+ ```python
92
+ import numpy as np
93
+ from k3_node.models import RadiusInteractionGraph
94
+
95
+ z = np.array([6, 8, 1, 1, 1]) # atomic numbers of a small molecule
96
+ pos = np.random.rand(5, 3).astype("float32") * 2.0 # 3D coordinates (Angstrom)
97
+
98
+ graph = RadiusInteractionGraph(cutoff=1.5) # connect atoms closer than 1.5 Angstrom
99
+ edge_index, edge_weight = graph(pos) # edge_weight holds the distances
100
+ print(edge_index.shape[0], edge_weight.shape == (edge_index.shape[1],)) # 2 True
101
+ ```
102
+ """
103
+ def __init__(self, cutoff: float = 10.0, max_num_neighbors: int = 32, **kwargs):
104
+ super().__init__(**kwargs)
105
+ self.cutoff = cutoff
106
+ self.max_num_neighbors = max_num_neighbors
107
+
108
+ def call(self, pos, batch=None):
109
+ edge_index = radius_graph(
110
+ pos,
111
+ r=self.cutoff,
112
+ batch=batch,
113
+ max_num_neighbors=self.max_num_neighbors,
114
+ )
115
+ row = edge_index[0]
116
+ col = edge_index[1]
117
+ pos_row = ops.take(pos, row, axis=0)
118
+ pos_col = ops.take(pos, col, axis=0)
119
+ edge_weight = ops.sqrt(ops.sum(ops.power(pos_row - pos_col, 2), axis=-1))
120
+ return edge_index, edge_weight
121
+
122
+
123
+ class CFConv(MessagePassing):
124
+ r"""Continuous-filter convolution layer.
125
+
126
+ Example:
127
+ ```python
128
+ import numpy as np
129
+ import keras
130
+ from k3_node.models import CFConv
131
+
132
+ x = np.random.rand(4, 16).astype("float32") # atom embeddings
133
+ edge_index = np.array([[0, 1, 0, 2, 1, 3], [1, 0, 2, 0, 3, 1]])
134
+ dist = np.random.rand(6).astype("float32") * 3.0 # edge lengths
135
+ edge_attr = np.random.rand(6, 10).astype("float32") # expanded distances (e.g. GaussianSmearing)
136
+
137
+ filter_net = keras.Sequential([keras.layers.Dense(16, activation="softplus"), keras.layers.Dense(16)])
138
+ conv = CFConv(in_channels=16, out_channels=16, num_filters=16, nn=filter_net, cutoff=5.0)
139
+ print(tuple(conv(x, edge_index, dist, edge_attr).shape)) # (4, 16): continuous-filter convolution
140
+ ```
141
+ """
142
+ def __init__(
143
+ self,
144
+ in_channels: int,
145
+ out_channels: int,
146
+ num_filters: int,
147
+ nn: keras.layers.Layer,
148
+ cutoff: float,
149
+ **kwargs,
150
+ ):
151
+ super().__init__(aggr="add", **kwargs)
152
+ self.in_channels = in_channels
153
+ self.out_channels = out_channels
154
+ self.num_filters = num_filters
155
+ self.nn = nn
156
+ self.cutoff = cutoff
157
+ self.lin1 = keras.layers.Dense(num_filters, use_bias=False)
158
+ self.lin2 = keras.layers.Dense(out_channels, use_bias=True)
159
+
160
+ def build(self, input_shape=None):
161
+ self.lin1.build((None, self.in_channels))
162
+ self.lin2.build((None, self.num_filters))
163
+ super().build(input_shape)
164
+
165
+ def call(self, x, edge_index, edge_weight, edge_attr):
166
+ C = 0.5 * (ops.cos(edge_weight * np.pi / self.cutoff) + 1.0)
167
+ W = self.nn(edge_attr) * ops.expand_dims(C, -1)
168
+ x = self.lin1(x)
169
+ x = self.propagate(edge_index, x=x, W=W)
170
+ x = self.lin2(x)
171
+ return x
172
+
173
+ def message(self, x_j, W):
174
+ return x_j * W
175
+
176
+
177
+ class InteractionBlock(keras.layers.Layer):
178
+ r"""Interaction block used in SchNet.
179
+
180
+ Example:
181
+ ```python
182
+ import numpy as np
183
+ from k3_node.models import SchNetInteractionBlock
184
+
185
+ x = np.random.rand(4, 16).astype("float32")
186
+ edge_index = np.array([[0, 1, 0, 2, 1, 3], [1, 0, 2, 0, 3, 1]])
187
+ dist = np.random.rand(6).astype("float32") * 3.0 # edge lengths
188
+ edge_attr = np.random.rand(6, 10).astype("float32") # expanded distances (e.g. GaussianSmearing)
189
+
190
+ block = SchNetInteractionBlock(hidden_channels=16, num_gaussians=10, num_filters=16, cutoff=5.0)
191
+ print(tuple(block(x, edge_index, dist, edge_attr).shape)) # (4, 16)
192
+ ```
193
+ """
194
+ def __init__(
195
+ self,
196
+ hidden_channels: int,
197
+ num_gaussians: int,
198
+ num_filters: int,
199
+ cutoff: float,
200
+ **kwargs,
201
+ ):
202
+ super().__init__(**kwargs)
203
+ self.hidden_channels = hidden_channels
204
+ self.num_gaussians = num_gaussians
205
+ self.num_filters = num_filters
206
+ self.cutoff = cutoff
207
+
208
+ self.mlp = keras.Sequential([
209
+ keras.layers.Dense(num_filters),
210
+ ShiftedSoftplus(),
211
+ keras.layers.Dense(num_filters),
212
+ ])
213
+ self.conv = CFConv(hidden_channels, hidden_channels, num_filters, self.mlp, cutoff)
214
+ self.act = ShiftedSoftplus()
215
+ self.lin = keras.layers.Dense(hidden_channels)
216
+
217
+ def call(self, x, edge_index, edge_weight, edge_attr):
218
+ x = self.conv(x, edge_index, edge_weight, edge_attr)
219
+ x = self.act(x)
220
+ x = self.lin(x)
221
+ return x
222
+
223
+
224
+ class SchNet(K3NodeHubMixin, keras.Model):
225
+ r"""The continuous-filter convolutional neural network SchNet from the
226
+ `"SchNet: A Continuous-filter Convolutional Neural Network for Modeling
227
+ Quantum Interactions" <https://arxiv.org/abs/1706.08566>`_ paper.
228
+
229
+ Args:
230
+ hidden_channels (int, optional): Hidden embedding size. (default: 128)
231
+ num_filters (int, optional): The number of filters to use. (default: 128)
232
+ num_interactions (int, optional): The number of interaction blocks. (default: 6)
233
+ num_gaussians (int, optional): The number of gaussians. (default: 50)
234
+ cutoff (float, optional): Cutoff distance. (default: 10.0)
235
+ interaction_graph (callable, optional): Interaction graph builder. (default: None)
236
+ max_num_neighbors (int, optional): Maximum neighbors per atom. (default: 32)
237
+ readout (str, optional): Readout pooling (add, sum, mean). (default: "add")
238
+ dipole (bool, optional): Predict dipole moment magnitude. (default: False)
239
+ mean (float, optional): Mean of target property. (default: None)
240
+ std (float, optional): Standard deviation of target property. (default: None)
241
+ atomref (tensor, optional): Reference atomic values. (default: None)
242
+
243
+ Example:
244
+ ```python
245
+ import numpy as np
246
+ from k3_node.models import SchNet
247
+
248
+ z = np.array([6, 8, 1, 1, 1]) # atomic numbers of a small molecule
249
+ pos = np.random.rand(5, 3).astype("float32") * 2.0 # 3D coordinates (Angstrom)
250
+ batch = np.array([0, 0, 0, 1, 1]) # two molecules: atoms 0-2 and atoms 3-4
251
+
252
+ model = SchNet(hidden_channels=16, num_filters=16, num_interactions=2, num_gaussians=10, cutoff=5.0)
253
+ energy = model(z, pos, batch=batch)
254
+ print(tuple(energy.shape)) # (2, 1)
255
+ ```
256
+ """
257
+ def __init__(
258
+ self,
259
+ hidden_channels: int = 128,
260
+ num_filters: int = 128,
261
+ num_interactions: int = 6,
262
+ num_gaussians: int = 50,
263
+ cutoff: float = 10.0,
264
+ interaction_graph: Optional[Callable] = None,
265
+ max_num_neighbors: int = 32,
266
+ readout: str = "add",
267
+ dipole: bool = False,
268
+ mean: Optional[float] = None,
269
+ std: Optional[float] = None,
270
+ atomref: Optional[any] = None,
271
+ **kwargs,
272
+ ):
273
+ super().__init__(**kwargs)
274
+ self.hidden_channels = hidden_channels
275
+ self.num_filters = num_filters
276
+ self.num_interactions = num_interactions
277
+ self.num_gaussians = num_gaussians
278
+ self.cutoff = cutoff
279
+ self.readout = readout
280
+ self.dipole = dipole
281
+ self.mean = mean
282
+ self.std = std
283
+ self.scale = None
284
+
285
+ try:
286
+ import ase
287
+ masses = np.array(ase.data.atomic_masses, dtype=np.float32)
288
+ except ImportError:
289
+ masses = np.array(DEFAULT_ATOMIC_MASSES, dtype=np.float32)
290
+
291
+ self.atomic_mass = self.add_weight(
292
+ name="atomic_mass",
293
+ shape=(len(masses),),
294
+ initializer=keras.initializers.Constant(masses),
295
+ trainable=False,
296
+ dtype="float32",
297
+ )
298
+
299
+ self.embedding = keras.layers.Embedding(100, hidden_channels)
300
+
301
+ if interaction_graph is not None:
302
+ self.interaction_graph = interaction_graph
303
+ else:
304
+ self.interaction_graph = RadiusInteractionGraph(cutoff, max_num_neighbors)
305
+
306
+ self.distance_expansion = GaussianSmearing(0.0, cutoff, num_gaussians)
307
+
308
+ self.interactions = [
309
+ InteractionBlock(hidden_channels, num_gaussians, num_filters, cutoff)
310
+ for _ in range(num_interactions)
311
+ ]
312
+
313
+ self.lin1 = keras.layers.Dense(hidden_channels // 2)
314
+ self.act = ShiftedSoftplus()
315
+ self.lin2 = keras.layers.Dense(1)
316
+
317
+ self.has_atomref = atomref is not None
318
+ if atomref is not None:
319
+ self.atomref = keras.layers.Embedding(
320
+ 100,
321
+ 1,
322
+ embeddings_initializer=keras.initializers.Constant(atomref),
323
+ )
324
+ else:
325
+ self.atomref = None
326
+
327
+ def call(self, z, pos, batch=None, batch_size=None):
328
+ if batch is None:
329
+ batch = ops.zeros(ops.shape(z), dtype="int32")
330
+ batch_size = 1
331
+ else:
332
+ batch = ops.cast(batch, "int32")
333
+
334
+ z = ops.cast(z, "int32")
335
+ h = self.embedding(z)
336
+ edge_index, edge_weight = self.interaction_graph(pos, batch)
337
+ edge_attr = self.distance_expansion(edge_weight)
338
+
339
+ for interaction in self.interactions:
340
+ h = h + interaction(h, edge_index, edge_weight, edge_attr)
341
+
342
+ h = self.lin1(h)
343
+ h = self.act(h)
344
+ h = self.lin2(h)
345
+
346
+ if self.dipole:
347
+ mass = ops.take(self.atomic_mass, z, axis=0)
348
+ mass = ops.expand_dims(mass, -1)
349
+ M = global_add_pool(mass, batch, size=batch_size)
350
+ c = global_add_pool(mass * pos, batch, size=batch_size) / (M + 1e-8)
351
+ c_per_atom = ops.take(c, batch, axis=0)
352
+ h = h * (pos - c_per_atom)
353
+
354
+ if not self.dipole and self.mean is not None and self.std is not None:
355
+ h = h * self.std + self.mean
356
+
357
+ if not self.dipole and self.atomref is not None:
358
+ h = h + self.atomref(z)
359
+
360
+ if self.dipole or self.readout in ["add", "sum"]:
361
+ out = global_add_pool(h, batch, size=batch_size)
362
+ else:
363
+ out = global_mean_pool(h, batch, size=batch_size)
364
+
365
+ if self.dipole:
366
+ out = ops.sqrt(ops.sum(ops.power(out, 2), axis=-1, keepdims=True))
367
+
368
+ if self.scale is not None:
369
+ out = self.scale * out
370
+
371
+ return out
372
+
373
+
374
+ QM9_TARGETS = {0: 'dipole_moment', 1: 'isotropic_polarizability', 2: 'homo', 3: 'lumo', 4: 'gap',
375
+ 5: 'electronic_spatial_extent', 6: 'zpve', 7: 'energy_U0', 8: 'energy_U', 9: 'enthalpy_H',
376
+ 10: 'free_energy', 11: 'heat_capacity'}
377
+ _DEBYE, _BOHR = 0.20819433442462576, 0.5291772105638411 # ase.units.Debye and ase.units.Bohr
378
+
379
+
380
+ def _load_schnetpack_model(path):
381
+ """Reads a pickled schnetpack model without schnetpack: its classes are replaced by stand-ins
382
+ that keep the modules' parameters (``_parameters``, ``_buffers``, ``_modules``)."""
383
+ import pickle
384
+ import warnings
385
+
386
+ import torch
387
+
388
+ class Stub:
389
+ def __init__(self, *args, **kwargs):
390
+ pass
391
+
392
+ def __setstate__(self, state):
393
+ self.__dict__.update(state if isinstance(state, dict) else {"_state": state})
394
+
395
+ class StubUnpickler(pickle.Unpickler):
396
+ def find_class(self, module, name):
397
+ if module.startswith(("schnetpack", "ase")):
398
+ return type(name, (Stub,), {})
399
+ return super().find_class(module, name)
400
+
401
+ class PickleModule:
402
+ Unpickler = StubUnpickler
403
+ load = pickle.load
404
+
405
+ with warnings.catch_warnings():
406
+ warnings.simplefilter("ignore")
407
+ return torch.load(path, map_location="cpu", pickle_module=PickleModule, weights_only=False)
408
+
409
+
410
+ def _schnet_from_qm9_pretrained(cls, root: str, dataset, target: int):
411
+ r"""Returns a :class:`SchNet` pre-trained on QM9 target ``target`` (the official SchNetPack
412
+ models, as in PyG), and the train/validation/test split it was trained with. Reading the
413
+ checkpoint needs PyTorch (not SchNetPack)."""
414
+ import os
415
+ import os.path as osp
416
+
417
+ from k3_node.data.download import download_url
418
+ from k3_node.data.extract import extract_zip
419
+
420
+ assert 0 <= target <= 11
421
+ root = osp.expanduser(root)
422
+ if not osp.exists(osp.join(root, 'trained_schnet_models')):
423
+ path = download_url('http://www.quantum-machine.org/datasets/trained_schnet_models.zip', root)
424
+ extract_zip(path, root)
425
+ os.unlink(path)
426
+ folder = osp.join(root, 'trained_schnet_models', f'qm9_{QM9_TARGETS[target]}')
427
+
428
+ # Keep only the characterized molecules of the split, as positions in `dataset`
429
+ split = np.load(osp.join(folder, 'split.npz'))
430
+ idx = np.asarray(ops.convert_to_numpy(dataset._data.idx)).reshape(-1)
431
+ assoc = np.full(int(idx.max()) + 1, -1)
432
+ assoc[idx] = np.arange(len(idx))
433
+ subsets = [assoc[s[np.isin(s, idx)]] for s in (split['train_idx'], split['val_idx'], split['test_idx'])]
434
+
435
+ state = _load_schnetpack_model(osp.join(folder, 'best_model'))
436
+
437
+ def mod(obj, *path):
438
+ for name in path:
439
+ obj = obj._modules[name]
440
+ return obj
441
+
442
+ def param(obj, name):
443
+ value = obj._parameters.get(name, None)
444
+ value = obj._buffers[name] if value is None else value
445
+ return value.detach().cpu().numpy()
446
+
447
+ output = mod(state, 'output_modules', '0')
448
+ dipole = type(output).__name__ == 'DipoleMoment'
449
+ has_atomref = output._modules.get('atomref') is not None
450
+ atomref = param(mod(output, 'atomref'), 'weight') if has_atomref else None
451
+ net = cls(hidden_channels=128, num_filters=128, num_interactions=6, num_gaussians=50, cutoff=10.0,
452
+ dipole=dipole, atomref=atomref)
453
+ net(np.array([6, 1, 1, 1, 1]), np.array([[0, 0, 0], [0.6, 0.6, 0.6], [-0.6, -0.6, 0.6], [-0.6, 0.6, -0.6],
454
+ [0.6, -0.6, -0.6]], dtype="float32")) # create the weights
455
+
456
+ def dense(layer, source, bias=True): # torch stores (out, in); Keras (in, out)
457
+ layer.kernel.assign(param(source, 'weight').T)
458
+ if bias:
459
+ layer.bias.assign(param(source, 'bias'))
460
+
461
+ rep = mod(state, 'representation')
462
+ net.embedding.embeddings.assign(param(mod(rep, 'embedding'), 'weight'))
463
+ for i, block in enumerate(net.interactions):
464
+ src = mod(rep, 'interactions', str(i))
465
+ dense(block.mlp.layers[0], mod(src, 'filter_network', '0'))
466
+ dense(block.mlp.layers[2], mod(src, 'filter_network', '1'))
467
+ dense(block.lin, mod(src, 'dense'))
468
+ dense(block.conv.lin1, mod(src, 'cfconv', 'in2f'), bias=False)
469
+ dense(block.conv.lin2, mod(src, 'cfconv', 'f2out'))
470
+ out_net = mod(output, 'out_net', '1', 'out_net')
471
+ dense(net.lin1, mod(out_net, '0'))
472
+ dense(net.lin2, mod(out_net, '1'))
473
+ average = getattr(output._modules.get('atom_pool'), 'average', False)
474
+ net.readout = 'mean' if average is True else 'add'
475
+ standardize = mod(output, 'standardize')
476
+ net.mean = float(param(standardize, 'mean').reshape(-1)[0])
477
+ net.std = float(param(standardize, 'stddev').reshape(-1)[0])
478
+ units = [1.0] * 12
479
+ units[0], units[1], units[5] = _DEBYE, _BOHR ** 3, _BOHR ** 2
480
+ net.scale = 1.0 / units[target]
481
+ return net, tuple(dataset[s] for s in subsets)
482
+
483
+
484
+ SchNet.from_qm9_pretrained = classmethod(_schnet_from_qm9_pretrained)