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,351 @@
1
+ """Core layers and mathematical primitives for materials models."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from typing import Sequence, Callable, Optional, Union
7
+ import keras
8
+ from keras import layers, ops
9
+ import numpy as np
10
+ from k3_node.ops.segment import segment_sum
11
+
12
+
13
+ class SoftPlus2(layers.Layer):
14
+ """SoftPlus2 activation: log(exp(x) + 1) - log(2). Zero at the origin."""
15
+
16
+ def __init__(self, **kwargs):
17
+ super().__init__(**kwargs)
18
+ self.shift = float(math.log(2.0))
19
+
20
+ def call(self, x):
21
+ return ops.softplus(x) - self.shift
22
+
23
+
24
+ class SoftExponential(layers.Layer):
25
+ """Soft exponential activation with learnable alpha."""
26
+
27
+ def __init__(self, alpha: float = 0.0, **kwargs):
28
+ super().__init__(**kwargs)
29
+ self.init_alpha = float(alpha)
30
+ self._eps = 1e-6
31
+
32
+ def build(self, input_shape=None):
33
+ self.alpha = self.add_weight(
34
+ name="alpha",
35
+ shape=(),
36
+ initializer=keras.initializers.Constant(self.init_alpha),
37
+ trainable=True,
38
+ dtype="float32",
39
+ )
40
+ super().build(input_shape)
41
+
42
+ def call(self, x):
43
+ alpha = self.alpha
44
+ near_zero = ops.abs(alpha) < self._eps
45
+ safe_alpha = ops.where(near_zero, 1.0, alpha)
46
+
47
+ neg_log_arg = ops.where(alpha < 0.0, 1.0 - alpha * (x + alpha), ops.ones_like(x))
48
+ neg_log_arg = ops.maximum(neg_log_arg, self._eps)
49
+ neg = -ops.log(neg_log_arg) / safe_alpha
50
+ pos = ops.expm1(safe_alpha * x) / safe_alpha + safe_alpha
51
+
52
+ out = ops.where(alpha < 0.0, neg, pos)
53
+ return ops.where(near_zero, x, out)
54
+
55
+
56
+ def get_activation(act: Union[str, Callable, layers.Layer, None]):
57
+ """Resolve activation specification into a callable or Keras layer."""
58
+ if act is None:
59
+ return lambda x: x
60
+ if isinstance(act, str):
61
+ act_lower = act.lower()
62
+ if act_lower in ("swish", "silu"):
63
+ return ops.silu
64
+ elif act_lower == "softplus2":
65
+ return SoftPlus2()
66
+ elif act_lower == "softexp":
67
+ return SoftExponential()
68
+ elif act_lower == "softplus":
69
+ return ops.softplus
70
+ elif act_lower == "tanh":
71
+ return ops.tanh
72
+ elif act_lower == "sigmoid":
73
+ return ops.sigmoid
74
+ elif act_lower == "relu":
75
+ return ops.relu
76
+ elif act_lower in ("identity", "linear", "none"):
77
+ return lambda x: x
78
+ return keras.activations.get(act)
79
+ if isinstance(act, layers.Layer):
80
+ return act
81
+ return act
82
+
83
+
84
+ class MLP(layers.Layer):
85
+ """Multi-layer perceptron compatible with multi-backend Keras 3."""
86
+
87
+ def __init__(
88
+ self,
89
+ dims: Sequence[int],
90
+ activation: Union[str, Callable, layers.Layer, None] = "swish",
91
+ activate_last: bool = False,
92
+ bias_last: bool = True,
93
+ use_bias: bool = True,
94
+ **kwargs,
95
+ ):
96
+ super().__init__(**kwargs)
97
+ self.dims = list(dims)
98
+ self.activate_last = activate_last
99
+ self.bias_last = bias_last
100
+ self.use_bias = use_bias
101
+ self.act_fn = get_activation(activation)
102
+
103
+ self.dense_layers = []
104
+ for i in range(len(self.dims) - 1):
105
+ is_last = (i == len(self.dims) - 2)
106
+ bias = self.bias_last if is_last else self.use_bias
107
+ dense = layers.Dense(self.dims[i + 1], use_bias=bias)
108
+ if self.dims[i] is not None:
109
+ dense.build((None, self.dims[i]))
110
+ self.dense_layers.append(dense)
111
+
112
+ def call(self, x):
113
+ for i, layer in enumerate(self.dense_layers):
114
+ x = layer(x)
115
+ is_last = (i == len(self.dense_layers) - 1)
116
+ if not is_last or self.activate_last:
117
+ x = self.act_fn(x)
118
+ return x
119
+
120
+
121
+ class GatedMLP(layers.Layer):
122
+ """Gated multi-layer perceptron: layer(x) * sigmoid(gate(x))."""
123
+
124
+ def __init__(
125
+ self,
126
+ in_feats: int,
127
+ dims: Sequence[int],
128
+ activate_last: bool = True,
129
+ use_bias: bool = True,
130
+ **kwargs,
131
+ ):
132
+ super().__init__(**kwargs)
133
+ self.in_feats = in_feats
134
+ self.dims = [in_feats, *dims]
135
+ self.activate_last = activate_last
136
+ self.use_bias = use_bias
137
+
138
+ self.val_layers = []
139
+ self.gate_layers = []
140
+ for i in range(len(self.dims) - 1):
141
+ out_dim = self.dims[i + 1]
142
+ val = layers.Dense(out_dim, use_bias=use_bias)
143
+ gate = layers.Dense(out_dim, use_bias=use_bias)
144
+ if self.dims[i] is not None:
145
+ val.build((None, self.dims[i]))
146
+ gate.build((None, self.dims[i]))
147
+ self.val_layers.append(val)
148
+ self.gate_layers.append(gate)
149
+
150
+ def call(self, x):
151
+ h_val = x
152
+ h_gate = x
153
+ for i in range(len(self.val_layers)):
154
+ is_last = (i == len(self.val_layers) - 1)
155
+ h_val = self.val_layers[i](h_val)
156
+ h_gate = self.gate_layers[i](h_gate)
157
+ if not is_last:
158
+ h_val = ops.silu(h_val)
159
+ h_gate = ops.silu(h_gate)
160
+ else:
161
+ if self.activate_last:
162
+ h_val = ops.silu(h_val)
163
+ h_gate = ops.sigmoid(h_gate)
164
+ return h_val * h_gate
165
+
166
+
167
+ class EmbeddingBlock(layers.Layer):
168
+ """Embeddings for nodes (atoms), edges (bonds), and global states."""
169
+
170
+ def __init__(
171
+ self,
172
+ degree_rbf: int,
173
+ dim_node_embedding: int,
174
+ dim_edge_embedding: Optional[int] = None,
175
+ dim_state_embedding: Optional[int] = None,
176
+ dim_state_feats: Optional[int] = None,
177
+ ntypes_node: Optional[int] = None,
178
+ ntypes_state: Optional[int] = None,
179
+ include_state: bool = False,
180
+ activation: Union[str, Callable, None] = "swish",
181
+ **kwargs,
182
+ ):
183
+ super().__init__(**kwargs)
184
+ self.degree_rbf = degree_rbf
185
+ self.dim_node_embedding = dim_node_embedding
186
+ self.dim_edge_embedding = dim_edge_embedding
187
+ self.dim_state_embedding = dim_state_embedding
188
+ self.dim_state_feats = dim_state_feats
189
+ self.ntypes_node = ntypes_node
190
+ self.ntypes_state = ntypes_state
191
+ self.include_state = include_state
192
+
193
+ if ntypes_node is not None:
194
+ self.layer_node_embedding = layers.Embedding(ntypes_node, dim_node_embedding)
195
+ else:
196
+ self.layer_node_embedding = layers.Dense(dim_node_embedding, use_bias=False)
197
+
198
+ if dim_edge_embedding is not None:
199
+ self.layer_edge_embedding = MLP([degree_rbf, dim_edge_embedding], activation=activation, activate_last=True)
200
+ else:
201
+ self.layer_edge_embedding = None
202
+
203
+ if include_state:
204
+ if ntypes_state is not None and dim_state_embedding is not None:
205
+ self.layer_state_embedding = layers.Embedding(ntypes_state, dim_state_embedding)
206
+ elif dim_state_feats is not None:
207
+ self.layer_state_embedding = layers.Dense(dim_state_feats, use_bias=False)
208
+ elif dim_state_embedding is not None:
209
+ self.layer_state_embedding = layers.Dense(dim_state_embedding, use_bias=False)
210
+ else:
211
+ self.layer_state_embedding = None
212
+ else:
213
+ self.layer_state_embedding = None
214
+
215
+ def call(self, node_attr, edge_attr, state_attr=None):
216
+ if self.ntypes_node is not None:
217
+ node_idx = ops.clip(ops.cast(node_attr, "int32"), 0, self.ntypes_node - 1)
218
+ node_feat = self.layer_node_embedding(node_idx)
219
+ else:
220
+ node_feat = self.layer_node_embedding(node_attr)
221
+
222
+ edge_feat = self.layer_edge_embedding(edge_attr) if self.layer_edge_embedding is not None else edge_attr
223
+
224
+ state_feat = None
225
+ if self.include_state and state_attr is not None and self.layer_state_embedding is not None:
226
+ if self.ntypes_state is not None:
227
+ state_idx = ops.clip(ops.cast(state_attr, "int32"), 0, self.ntypes_state - 1)
228
+ state_feat = self.layer_state_embedding(state_idx)
229
+ else:
230
+ state_feat = self.layer_state_embedding(state_attr)
231
+ elif self.include_state and state_attr is not None:
232
+ state_feat = state_attr
233
+
234
+ return node_feat, edge_feat, state_feat
235
+
236
+
237
+ # --- Tensor math utilities for TensorNet ---
238
+
239
+ def vector_to_skewtensor(vector):
240
+ """Create skew-symmetric 3x3 tensor from a 3D vector.
241
+
242
+ ```
243
+ [0, -v_z, v_y]
244
+ [v_z, 0, -v_x]
245
+ [-v_y, v_x, 0]
246
+ ```
247
+ """
248
+ vector = ops.convert_to_tensor(vector)
249
+ vx = vector[..., 0]
250
+ vy = vector[..., 1]
251
+ vz = vector[..., 2]
252
+ zero = ops.zeros_like(vx)
253
+
254
+ row0 = ops.stack([zero, -vz, vy], axis=-1)
255
+ row1 = ops.stack([vz, zero, -vx], axis=-1)
256
+ row2 = ops.stack([-vy, vx, zero], axis=-1)
257
+ return ops.stack([row0, row1, row2], axis=-2)
258
+
259
+
260
+ def vector_to_symtensor(vector):
261
+ """Create symmetric traceless tensor from outer product of 3D vector with itself."""
262
+ vector = ops.convert_to_tensor(vector)
263
+ v_col = ops.expand_dims(vector, axis=-1)
264
+ v_row = ops.expand_dims(vector, axis=-2)
265
+ outer = ops.matmul(v_col, v_row)
266
+ # trace = outer[..., 0, 0] + outer[..., 1, 1] + outer[..., 2, 2]
267
+ trace = outer[..., 0, 0] + outer[..., 1, 1] + outer[..., 2, 2]
268
+ eye3 = ops.eye(3, dtype=outer.dtype)
269
+ scalars = ops.expand_dims(ops.expand_dims(trace / 3.0, axis=-1), axis=-1) * eye3
270
+ ndim = len(outer.shape)
271
+ axes = list(range(ndim - 2)) + [ndim - 1, ndim - 2]
272
+ sym = 0.5 * (outer + ops.transpose(outer, axes=axes))
273
+ return sym - scalars
274
+
275
+
276
+ def decompose_tensor(tensor):
277
+ """Decompose 3x3 Cartesian tensor into scalar (I), skew-symmetric (A), and symmetric traceless (S)."""
278
+ tensor = ops.convert_to_tensor(tensor)
279
+ trace = tensor[..., 0, 0] + tensor[..., 1, 1] + tensor[..., 2, 2]
280
+ eye3 = ops.eye(3, dtype=tensor.dtype)
281
+ scalars = ops.expand_dims(ops.expand_dims(trace / 3.0, axis=-1), axis=-1) * eye3
282
+ ndim = len(tensor.shape)
283
+ axes = list(range(ndim - 2)) + [ndim - 1, ndim - 2]
284
+ transposed = ops.transpose(tensor, axes=axes)
285
+ skew = 0.5 * (tensor - transposed)
286
+ sym = 0.5 * (tensor + transposed)
287
+ traceless = sym - scalars
288
+ return scalars, skew, traceless
289
+
290
+
291
+ def new_radial_tensor(scalars, skew, traceless, f_I, f_A, f_S):
292
+ """Multiply irreducible tensor components by radial invariant features."""
293
+ scalars_out = ops.expand_dims(ops.expand_dims(f_I, axis=-1), axis=-1) * scalars
294
+ skew_out = ops.expand_dims(ops.expand_dims(f_A, axis=-1), axis=-1) * skew
295
+ traceless_out = ops.expand_dims(ops.expand_dims(f_S, axis=-1), axis=-1) * traceless
296
+ return scalars_out, skew_out, traceless_out
297
+
298
+
299
+ def tensor_norm(tensor):
300
+ """Computes Frobenius norm squared across last two dimensions (3, 3)."""
301
+ return ops.sum(tensor ** 2, axis=(-2, -1))
302
+
303
+
304
+ def scatter_add(x, index, num_segments: int):
305
+ """Scatter sum x elements into segments indicated by index."""
306
+ index = ops.cast(index, "int32")
307
+ return segment_sum(x, index, num_segments=num_segments)
308
+
309
+
310
+ def scatter_mean(x, index, num_segments: int):
311
+ """Scatter mean x elements into segments indicated by index."""
312
+ index = ops.cast(index, "int32")
313
+ sums = segment_sum(x, index, num_segments=num_segments)
314
+ ones = ops.ones_like(x[..., :1])
315
+ counts = segment_sum(ones, index, num_segments=num_segments)
316
+ counts = ops.maximum(counts, 1.0)
317
+ return sums / counts
318
+
319
+
320
+ def infer_num_graphs(batch=None, num_graphs=None, state_attr=None):
321
+ """Infer the number of graphs in a batch across backends (JAX, PyTorch, TF)."""
322
+ if num_graphs is not None:
323
+ return num_graphs
324
+ if state_attr is not None:
325
+ shape = ops.shape(state_attr)
326
+ if len(shape) == 1:
327
+ return 1
328
+ return shape[0]
329
+ if batch is None:
330
+ return 1
331
+ shape = ops.shape(batch)
332
+ if len(shape) > 0 and shape[0] == 0:
333
+ return 0
334
+ try:
335
+ val = ops.convert_to_numpy(ops.max(batch))
336
+ return int(val) + 1
337
+ except Exception:
338
+ pass
339
+ try:
340
+ if hasattr(batch, "__getitem__"):
341
+ val = batch[-1]
342
+ if hasattr(val, "item"):
343
+ return int(val.item()) + 1
344
+ return int(val) + 1
345
+ except Exception:
346
+ pass
347
+ try:
348
+ return ops.cast(ops.max(batch) + 1, "int32")
349
+ except Exception:
350
+ return 1
351
+
@@ -0,0 +1,246 @@
1
+ """Multi-backend Keras 3 implementation of GRACE (Graph Atomic Cluster Expansion)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Sequence, Optional, Union, Tuple, Dict, Any, Literal
6
+ import keras
7
+ from keras import layers, ops
8
+ import numpy as np
9
+
10
+ from .core import MLP, scatter_add, infer_num_graphs
11
+ from .basis import (
12
+ ChebyshevRadialBasis,
13
+ compute_pair_vector_and_distance,
14
+ polynomial_cutoff,
15
+ )
16
+ from .so3net import RealSphericalHarmonics
17
+ from .readout import ReduceReadOut
18
+
19
+
20
+ class GraceSPBasis(layers.Layer):
21
+ """Single-particle ACE basis aggregation for atomic clusters.
22
+
23
+ Example:
24
+ ```python
25
+ import numpy as np
26
+ from k3_node.models import GraceSPBasis
27
+
28
+ edge_index = np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]])
29
+ node_type = np.array([6, 8, 1, 6])
30
+ rad_basis = np.random.rand(8, 6).astype("float32") # radial basis per bond
31
+ sh_basis = np.random.rand(8, 9).astype("float32") # spherical harmonics per bond (lmax=2)
32
+ layer = GraceSPBasis(n_rad_base=6, lmax=2, embedding_size=8)
33
+ print(tuple(layer(edge_index, node_type, rad_basis, sh_basis, num_nodes=4).shape)) # (4, 9, 8): atomic basis A_i
34
+ ```
35
+ """
36
+
37
+ def __init__(
38
+ self,
39
+ n_rad_base: int = 6,
40
+ lmax: int = 2,
41
+ embedding_size: int = 8,
42
+ ntypes_node: int = 95,
43
+ **kwargs,
44
+ ):
45
+ super().__init__(**kwargs)
46
+ self.n_rad_base = n_rad_base
47
+ self.lmax = lmax
48
+ self.num_sh = (lmax + 1) ** 2
49
+ self.elem_emb = layers.Embedding(ntypes_node, embedding_size)
50
+ self.radial_proj = layers.Dense(embedding_size, use_bias=False)
51
+
52
+ def call(self, edge_index, node_type, rad_basis, sh_basis, num_nodes: Optional[int] = None):
53
+ src = ops.cast(edge_index[0], "int32")
54
+ dst = ops.cast(edge_index[1], "int32")
55
+
56
+ # Neighbor chemical indicator
57
+ neigh_z = ops.take(node_type, dst, axis=0)
58
+ c_j = self.elem_emb(ops.cast(neigh_z, "int32"))
59
+
60
+ # Radial modulation
61
+ R_nl = self.radial_proj(rad_basis)
62
+ radial_atom = c_j * R_nl
63
+
64
+ # Outer product with spherical harmonics: [E, num_sh, emb_size]
65
+ sp_msg = ops.expand_dims(sh_basis, axis=-1) * ops.expand_dims(radial_atom, axis=1)
66
+
67
+ # Aggregate into central atom
68
+ A_i = scatter_add(sp_msg, src, num_segments=num_nodes)
69
+ return A_i
70
+
71
+
72
+ class GraceACEStack(layers.Layer):
73
+ """Multi-order Atomic Cluster Expansion stack accumulating rotational invariants.
74
+
75
+ Example:
76
+ ```python
77
+ import numpy as np
78
+ from k3_node.models import GraceACEStack
79
+
80
+ A_i = np.random.rand(4, 9, 8).astype("float32") # atomic basis from GraceSPBasis
81
+ print(tuple(GraceACEStack(embedding_size=8, max_order=2)(A_i).shape)) # (4, 16): body-order invariants
82
+ ```
83
+ """
84
+
85
+ def __init__(
86
+ self,
87
+ embedding_size: int = 8,
88
+ max_order: int = 3,
89
+ **kwargs,
90
+ ):
91
+ super().__init__(**kwargs)
92
+ self.embedding_size = embedding_size
93
+ self.max_order = max_order
94
+
95
+ def call(self, A_i):
96
+ A_00 = A_i[:, 0, :]
97
+ invariants = [A_00]
98
+
99
+ if self.max_order >= 2:
100
+ A2 = ops.sum(A_i ** 2, axis=1)
101
+ invariants.append(A2)
102
+
103
+ if self.max_order >= 3:
104
+ A3 = A_00 * ops.sum(A_i ** 2, axis=1)
105
+ invariants.append(A3)
106
+
107
+ return ops.concatenate(invariants, axis=-1)
108
+
109
+
110
+ class GRACE(keras.Model):
111
+ """Graph Atomic Cluster Expansion (GRACE) foundational interatomic potential.
112
+
113
+ Example:
114
+ ```python
115
+ import numpy as np
116
+ from k3_node.models import GRACE
117
+
118
+ # A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
119
+ structure = {
120
+ "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
121
+ "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
122
+ "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), # bond pairs forming angles
123
+ "node_type": np.array([6, 8, 1, 6]), # atomic numbers
124
+ "batch": np.zeros(4, dtype="int32"), # all atoms belong to structure 0
125
+ "state_attr": np.zeros((1, 2), dtype="float32"), # global state features
126
+ }
127
+
128
+ model = GRACE(cutoff=5.0, n_rad_base=6, lmax=2, embedding_size=8, max_order=2, nblocks=2,
129
+ readout_hidden=(16,))
130
+ energy = model(structure) # predicted property (e.g. energy) of the structure
131
+ print(tuple(energy.shape)) # (1,)
132
+ ```
133
+ """
134
+
135
+ def __init__(
136
+ self,
137
+ cutoff: float = 5.0,
138
+ n_rad_base: int = 6,
139
+ lmax: int = 2,
140
+ embedding_size: int = 8,
141
+ max_order: int = 3,
142
+ nblocks: int = 2,
143
+ readout_hidden: Sequence[int] = (32,),
144
+ ntypes_node: int = 95,
145
+ activation_type: str = "swish",
146
+ **kwargs,
147
+ ):
148
+ super().__init__(**kwargs)
149
+ self.cutoff = cutoff
150
+ self.nblocks = nblocks
151
+
152
+ self.chebyshev = ChebyshevRadialBasis(nfunc=n_rad_base, cutoff=cutoff, cutoff_exponent=5)
153
+ self.sh = RealSphericalHarmonics(lmax=lmax)
154
+
155
+ self.sp_bases = []
156
+ self.ace_stacks = []
157
+ self.readouts = []
158
+ inv_dim = embedding_size * min(max_order, 3)
159
+
160
+ for _ in range(nblocks):
161
+ self.sp_bases.append(
162
+ GraceSPBasis(
163
+ n_rad_base=n_rad_base,
164
+ lmax=lmax,
165
+ embedding_size=embedding_size,
166
+ ntypes_node=ntypes_node,
167
+ )
168
+ )
169
+ self.ace_stacks.append(
170
+ GraceACEStack(
171
+ embedding_size=embedding_size,
172
+ max_order=max_order,
173
+ )
174
+ )
175
+ self.readouts.append(
176
+ MLP([inv_dim, *readout_hidden, 1], activation=activation_type, activate_last=False)
177
+ )
178
+
179
+ self.graph_pool = ReduceReadOut(op="sum")
180
+
181
+ def _unpack_inputs(self, inputs):
182
+ if isinstance(inputs, dict):
183
+ pos = inputs.get("pos")
184
+ edge_index = inputs.get("edge_index")
185
+ node_type = inputs.get("node_type", inputs.get("z"))
186
+ pbc_offshift = inputs.get("pbc_offshift", None)
187
+ batch = inputs.get("batch", None)
188
+ num_graphs = inputs.get("num_graphs", None)
189
+ state_attr = inputs.get("state_attr", None)
190
+ return pos, edge_index, node_type, pbc_offshift, batch, num_graphs, state_attr
191
+ elif isinstance(inputs, (tuple, list)):
192
+ pos = inputs[0]
193
+ edge_index = inputs[1]
194
+ node_type = inputs[2]
195
+ pbc_offshift = inputs[3] if len(inputs) > 3 else None
196
+ batch = inputs[4] if len(inputs) > 4 else None
197
+ num_graphs = inputs[5] if len(inputs) > 5 else None
198
+ state_attr = inputs[6] if len(inputs) > 6 else None
199
+ return pos, edge_index, node_type, pbc_offshift, batch, num_graphs, state_attr
200
+ return inputs, None, None, None, None, None, None
201
+
202
+ def call(self, inputs, edge_index=None, node_type=None, pbc_offshift=None, batch=None, num_graphs=None, state_attr=None):
203
+ if edge_index is None:
204
+ (
205
+ pos,
206
+ edge_index,
207
+ node_type_in,
208
+ pbc_offshift_in,
209
+ batch_in,
210
+ num_graphs_in,
211
+ state_attr_in,
212
+ ) = self._unpack_inputs(inputs)
213
+ if node_type is None:
214
+ node_type = node_type_in
215
+ if pbc_offshift is None:
216
+ pbc_offshift = pbc_offshift_in
217
+ if batch is None:
218
+ batch = batch_in
219
+ if num_graphs is None:
220
+ num_graphs = num_graphs_in
221
+ if state_attr is None:
222
+ state_attr = state_attr_in
223
+ else:
224
+ pos = inputs
225
+
226
+ num_nodes = ops.shape(pos)[0]
227
+ if batch is None:
228
+ batch = ops.zeros((num_nodes,), dtype="int32")
229
+ else:
230
+ batch = ops.cast(batch, "int32")
231
+ n_graphs = infer_num_graphs(batch=batch, num_graphs=num_graphs, state_attr=state_attr)
232
+
233
+ vec, bond_dists = compute_pair_vector_and_distance(pos, edge_index, pbc_offshift)
234
+ rad_basis = self.chebyshev(bond_dists)
235
+ sh_basis = self.sh(vec)
236
+
237
+ total_atomic_energy = ops.zeros((num_nodes, 1), dtype=pos.dtype)
238
+ for i in range(self.nblocks):
239
+ A_i = self.sp_bases[i](edge_index, node_type, rad_basis, sh_basis, num_nodes=num_nodes)
240
+ invariants = self.ace_stacks[i](A_i)
241
+ e_block = self.readouts[i](invariants)
242
+ total_atomic_energy = total_atomic_energy + e_block
243
+
244
+ total_energy = self.graph_pool(total_atomic_energy, batch=batch, num_graphs=n_graphs)
245
+ return ops.squeeze(total_energy, axis=-1)
246
+